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 } var tableFamilies = map[config.AddressFamily]nftables.TableFamily{ config.FamilyINET: nftables.TableFamilyINet, config.FamilyIP: nftables.TableFamilyIPv4, config.FamilyIP6: nftables.TableFamilyIPv6, } func (e *Engine) family() nftables.TableFamily { if f, ok := tableFamilies[e.cfg.Settings.AddressFamily]; ok { return f } return nftables.TableFamilyINet } func addressFamily(tf nftables.TableFamily) config.AddressFamily { for f, t := range tableFamilies { if t == tf { return f } } return config.FamilyINET } // withFamily is the engine for the same table name in another address family. func (e *Engine) withFamily(f config.AddressFamily) *Engine { cfg := *e.cfg cfg.Settings.AddressFamily = f return &Engine{cfg: &cfg, conn: e.conn} } // overlaps reports whether tables of families a and b filter the same traffic: // inet covers both ip and ip6, which do not overlap each other. func overlaps(a, b nftables.TableFamily) bool { if a == b { return false } return (a == nftables.TableFamilyINet && (b == nftables.TableFamilyIPv4 || b == nftables.TableFamilyIPv6)) || (b == nftables.TableFamilyINet && (a == nftables.TableFamilyIPv4 || a == nftables.TableFamilyIPv6)) } // staleTables are our-named tables left by a different address_family that // would still filter the traffic this family now owns. func (e *Engine) staleTables() ([]*nftables.Table, error) { tables, err := e.conn.ListTables() if err != nil { return nil, fmt.Errorf("listing tables: %w", err) } var stale []*nftables.Table for _, t := range tables { if t.Name == e.cfg.Settings.TableName && overlaps(e.family(), t.Family) { stale = append(stale, t) } } return stale, 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: e.family(), 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 { stale, err := e.staleTables() if err != nil { return err } for _, t := range stale { e.conn.DelTable(t) } 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.staleTables() if err != nil { return err } own, err := e.findTable() if err != nil { return err } if own != nil { tables = append(tables, own) } if len(tables) == 0 { return nil } for _, t := range tables { e.conn.DelTable(t) } return e.conn.Flush() } 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 == e.family() { 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(e.family()) 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"` Family config.AddressFamily `json:"family,omitempty"` 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, // falling back to the overlapping table of another family that apply replaces. func (e *Engine) Snapshot() (*Snapshot, error) { t, err := e.findTable() if err != nil { return nil, err } if t == nil { stale, err := e.staleTables() if err != nil { return nil, err } // ponytail: captures one stale table; an inet config replacing both ip and ip6 restores only the first. if len(stale) > 0 { return e.withFamily(addressFamily(stale[0].Family)).Snapshot() } } snap := &Snapshot{Table: e.cfg.Settings.TableName, Family: addressFamily(e.family())} if t == nil { return snap, nil } snap.Present = true chains, err := e.conn.ListChainsOfTableFamily(e.family()) 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, byte(e.family())) 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) } // Snapshots predating the family field are of the inet table. if f := addressFamily(tableFamilies[s.Family]); f != addressFamily(e.family()) { return e.withFamily(f).Restore(s) } if !s.Present { return e.Flush() } want, err := decodeState(s.Rules, byte(e.family())) 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 }