package nftables import ( "fmt" "github.com/google/nftables" "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() if err != nil { return nil, fmt.Errorf("connecting to nftables: %w", err) } return &Engine{cfg: cfg, conn: conn}, nil } 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, }, "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 _, 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 } 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"` } // 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 } 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 } current, err := e.readCurrentState() if err != nil { return err } return e.apply(restoreChangeSet(current, want), s.Policies) } func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { return &p }