package nftables import ( "fmt" "github.com/google/nftables" "github.com/google/nftables/expr" "github.com/mdlayher/netlink" "golang.org/x/sys/unix" "git.unkin.net/unkin/tomswall/internal/config" ) type Engine struct { cfg *config.Config conn *nftables.Conn } func NewEngine(cfg *config.Config) (*Engine, error) { conn, err := nftables.New(nftables.WithSockOptions(largeBuffers)) if err != nil { return nil, fmt.Errorf("connecting to nftables: %w", err) } return &Engine{cfg: cfg, conn: conn}, nil } // batchBufSize bounds one batch: the kernel rejects a batch larger than the // send buffer (EMSGSIZE) and drops ACKs beyond the receive buffer (ENOBUFS) // after committing it. // ponytail: fixed cap of tens of thousands of rules; size per batch if exceeded. const batchBufSize = 64 << 20 // largeBuffers raises both socket buffers, ignoring rmem_max/wmem_max when // CAP_NET_ADMIN allows it and falling back to the capped sizes otherwise. func largeBuffers(c *netlink.Conn) error { rc, err := c.SyscallConn() if err != nil { return err } var serr error err = rc.Control(func(fd uintptr) { for _, o := range [][2]int{{unix.SO_SNDBUFFORCE, unix.SO_SNDBUF}, {unix.SO_RCVBUFFORCE, unix.SO_RCVBUF}} { if unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[0], batchBufSize) == nil { continue } if serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[1], batchBufSize); serr != nil { return } } }) if err != nil { return err } return serr } func (e *Engine) ensureTable() *nftables.Table { return e.conn.AddTable(&nftables.Table{ Family: nftables.TableFamilyINet, Name: e.cfg.Settings.TableName, }) } // ensureChains declares the base chains; policies overrides their default policy. func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain { chains := map[string]*nftables.Chain{ "input": { Name: "input", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookInput, Priority: nftables.ChainPriorityFilter, Policy: policyPtr(nftables.ChainPolicyDrop), }, "forward": { Name: "forward", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookForward, Priority: nftables.ChainPriorityFilter, Policy: policyPtr(nftables.ChainPolicyDrop), }, "output": { Name: "output", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookOutput, Priority: nftables.ChainPriorityFilter, Policy: policyPtr(nftables.ChainPolicyAccept), }, "postrouting": { Name: "postrouting", Table: table, Type: nftables.ChainTypeNAT, Hooknum: nftables.ChainHookPostrouting, Priority: nftables.ChainPriorityNATSource, }, "prerouting": { Name: "prerouting", Table: table, Type: nftables.ChainTypeNAT, Hooknum: nftables.ChainHookPrerouting, Priority: nftables.ChainPriorityNATDest, }, "helper_prerouting": { Name: "helper_prerouting", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookPrerouting, Priority: nftables.ChainPriorityMangle, }, "helper_output": { Name: "helper_output", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookOutput, Priority: nftables.ChainPriorityMangle, }, "raw_prerouting": { Name: "raw_prerouting", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookPrerouting, Priority: nftables.ChainPriorityRaw, Policy: policyPtr(nftables.ChainPolicyAccept), }, "raw_output": { Name: "raw_output", Table: table, Type: nftables.ChainTypeFilter, Hooknum: nftables.ChainHookOutput, Priority: nftables.ChainPriorityRaw, Policy: policyPtr(nftables.ChainPolicyAccept), }, } for name, chain := range chains { if p, ok := policies[name]; ok { chain.Policy = policyPtr(p) } chains[name] = e.conn.AddChain(chain) } return chains } func (e *Engine) Plan() (*ChangeSet, error) { compiler := NewCompiler(e.cfg) desired, err := compiler.Compile() if err != nil { return nil, fmt.Errorf("compiling config: %w", err) } current, err := e.readCurrentState() if err != nil { return nil, fmt.Errorf("reading current state: %w", err) } return computeDiff(current, desired), nil } func (e *Engine) Apply(changes *ChangeSet) error { return e.apply(changes, nil) } func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { table := e.ensureTable() chains := e.ensureChains(table, policies) for _, r := range changes.Remove { e.conn.DelRule(&nftables.Rule{ Table: table, Chain: chains[r.Chain], Handle: r.Handle, }) } for _, n := range changes.RemoveHelpers { e.conn.DeleteObject(helperObj(table, n, expr.CtHelper{})) } for _, h := range changes.AddHelpers { e.conn.AddObj(helperObj(table, h.Name, h.Helper)) } for _, r := range changes.Add { chain, ok := chains[r.Chain] if !ok { return fmt.Errorf("unknown chain %q", r.Chain) } rule := &nftables.Rule{ Table: table, Chain: chain, Exprs: r.Exprs, UserData: []byte(r.Tag), } if r.Before != 0 { rule.Position = r.Before e.conn.InsertRule(rule) } else { e.conn.AddRule(rule) } } return e.conn.Flush() } func (e *Engine) Flush() error { tables, err := e.conn.ListTables() if err != nil { return fmt.Errorf("listing tables: %w", err) } for _, t := range tables { if t.Name == e.cfg.Settings.TableName { e.conn.DelTable(t) return e.conn.Flush() } } return nil } func (e *Engine) findTable() (*nftables.Table, error) { tables, err := e.conn.ListTables() if err != nil { return nil, fmt.Errorf("listing tables: %w", err) } for _, t := range tables { if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { return t, nil } } return nil, nil } func (e *Engine) readCurrentState() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } ourTable, err := e.findTable() if err != nil || ourTable == nil { return state, err } objs, err := e.conn.GetNamedObjects(ourTable) if err != nil { return nil, fmt.Errorf("listing objects: %w", err) } for _, o := range objs { if no, ok := o.(*nftables.NamedObj); ok && no.Type == nftables.ObjTypeCtHelper { if h, ok := no.Obj.(*expr.CtHelper); ok { state.Helpers = append(state.Helpers, Helper{Name: no.Name, Helper: *h}) } } } chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) if err != nil { return nil, fmt.Errorf("listing chains: %w", err) } for _, chain := range chains { if chain.Table.Name != e.cfg.Settings.TableName { continue } rules, err := e.conn.GetRules(ourTable, chain) if err != nil { return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err) } for _, rule := range rules { state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{ Chain: chain.Name, Handle: rule.Handle, Exprs: rule.Exprs, Tag: string(rule.UserData), }) } } return state, nil } // Snapshot is the tomswall table as captured live, serialisable so a revert // survives the process that took it. type Snapshot struct { Table string `json:"table"` Present bool `json:"present"` Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` Rules map[string][]SnapshotRule `json:"rules,omitempty"` Helpers []Helper `json:"helpers,omitempty"` } // SnapshotRule is a managed rule with its expressions in netlink wire format. type SnapshotRule struct { Tag string `json:"tag"` Exprs [][]byte `json:"exprs"` } // Snapshot captures the live tomswall table so Restore can roll back to it. func (e *Engine) Snapshot() (*Snapshot, error) { snap := &Snapshot{Table: e.cfg.Settings.TableName} t, err := e.findTable() if err != nil || t == nil { return snap, err } snap.Present = true chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) if err != nil { return nil, fmt.Errorf("listing chains: %w", err) } snap.Policies = make(map[string]nftables.ChainPolicy) for _, c := range chains { if c.Table.Name == snap.Table && c.Policy != nil { snap.Policies[c.Name] = *c.Policy } } state, err := e.readCurrentState() if err != nil { return nil, err } snap.Rules, err = encodeState(state) if err != nil { return nil, err } snap.Helpers = state.Helpers return snap, nil } // Restore atomically returns the tomswall table to the snapshot: rule order // and chain policies included, or removed if it was absent. func (e *Engine) Restore(s *Snapshot) error { if s.Table != e.cfg.Settings.TableName { return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) } if !s.Present { return e.Flush() } want, err := decodeState(s.Rules) if err != nil { return err } want.Helpers = s.Helpers current, err := e.readCurrentState() if err != nil { return err } return e.apply(restoreChangeSet(current, want), s.Policies) } func helperObj(table *nftables.Table, name string, h expr.CtHelper) *nftables.NamedObj { return &nftables.NamedObj{Table: table, Name: name, Type: nftables.ObjTypeCtHelper, Obj: &h} } func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { return &p }