Files
tomswall/internal/nftables/engine.go
T
unkin-agent c557c4b78a
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
Attach conntrack helpers via ct helper objects
2026-10-03 23:50:31 +10:00

311 lines
7.8 KiB
Go

package nftables
import (
"fmt"
"github.com/google/nftables"
"github.com/google/nftables/expr"
"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,
},
"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,
},
}
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
}