diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 432d51a..62543cd 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -5,6 +5,7 @@ import ( "fmt" "log/slog" "net" + "slices" "sort" "strconv" "strings" @@ -212,9 +213,74 @@ func (c *Compiler) compileBlrules(state *FirewallState) error { return nil } +// helperProtos is each kernel helper's default transport, used when an entry gives no proto. +var helperProtos = map[string]string{ + "amanda": "udp", "ftp": "tcp", "irc": "tcp", "netbios-ns": "udp", "pptp": "tcp", + "Q.931": "tcp", "RAS": "udp", "sane": "tcp", "sip": "udp", "snmp": "udp", "tftp": "udp", +} + +// compileHelper declares a ct helper object per (helper, proto) and assigns it +// after conntrack (-200) has created the entry; a raw-priority assignment is a no-op. +func (c *Compiler) compileHelper(state *FirewallState, tag string, ct config.ConntrackRule) error { + if ct.Source != "" || ct.Dest != "" { + slog.Warn("conntrack helper source/dest not supported, assigning globally", "rule", tag, "helper", ct.Helper) + } + chains := []string{"helper_prerouting", "helper_output"} + switch ct.Chain { + case config.ConntrackPrerouting: + chains = chains[:1] + case config.ConntrackOutput: + chains = chains[1:] + } + + proto := ct.Proto + if proto == "" { + if proto = helperProtos[ct.Helper]; proto == "" { + return fmt.Errorf("proto required for helper %q", ct.Helper) + } + } + for _, p := range strings.Split(proto, ",") { + p = strings.TrimSpace(p) + n, err := protoNumber(p) + if err != nil { + return err + } + name := ct.Helper + if p != helperProtos[ct.Helper] { + name += "-" + p + } + if !slices.ContainsFunc(state.Helpers, func(h Helper) bool { return h.Name == name }) { + state.Helpers = append(state.Helpers, Helper{Name: name, + Helper: expr.CtHelper{Name: ct.Helper, L3Proto: unix.NFPROTO_INET, L4Proto: n}}) + } + + matches, err := l4Matches(p, ct.DPort, ct.SPort) + if err != nil { + return err + } + for _, chain := range chains { + for _, m := range matches { + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: append(append([]expr.Any{}, m.exprs...), + &expr.Objref{Type: unix.NFT_OBJECT_CT_HELPER, Name: name}), + Tag: tag + ":" + chain, + }) + } + } + } + return nil +} + func (c *Compiler) compileConntrack(state *FirewallState) error { for i, ct := range c.cfg.Conntrack { tag := fmt.Sprintf("conntrack:%d", i) + if ct.Action == config.ConntrackHelper { + if err := c.compileHelper(state, tag, ct); err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + continue + } chains := []string{"prerouting"} switch ct.Chain { @@ -236,8 +302,6 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { switch ct.Action { case config.ConntrackNoTrack: exprs = append(exprs, &expr.Notrack{}) - case config.ConntrackHelper: - continue case config.ConntrackDrop: exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 40d7016..f970417 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -2417,3 +2417,63 @@ func TestCompile_OrigDestForwardRejected(t *testing.T) { t.Fatalf("Compile() error = %v, want forwarded ORIGDEST rejection", err) } } + +func TestCompile_ConntrackHelper(t *testing.T) { + compile := func(rules ...config.ConntrackRule) (*FirewallState, error) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}}, + Conntrack: rules, + PortGroups: map[string]config.PortGroup{}, + } + return NewCompiler(cfg).Compile() + } + + state, err := compile( + config.ConntrackRule{Action: config.ConntrackHelper, Helper: "ftp", Proto: "tcp", DPort: config.PortSpec{"21"}}, + config.ConntrackRule{Action: config.ConntrackHelper, Helper: "sip", Proto: "tcp", DPort: config.PortSpec{"5060"}, Chain: config.ConntrackPrerouting}, + config.ConntrackRule{Action: config.ConntrackHelper, Helper: "tftp", Chain: config.ConntrackOutput}, + ) + if err != nil { + t.Fatal(err) + } + + want := []Helper{ + {"ftp", expr.CtHelper{Name: "ftp", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_TCP}}, + {"sip-tcp", expr.CtHelper{Name: "sip", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_TCP}}, + {"tftp", expr.CtHelper{Name: "tftp", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_UDP}}, + } + if !reflect.DeepEqual(state.Helpers, want) { + t.Errorf("helpers = %+v, want %+v", state.Helpers, want) + } + + refs := func(chain string) []string { + var out []string + for _, r := range state.Rules[chain] { + ref, ok := r.Exprs[len(r.Exprs)-1].(*expr.Objref) + if !ok || ref.Type != unix.NFT_OBJECT_CT_HELPER { + t.Fatalf("%s: rule %s does not end in a ct helper objref", chain, r.Tag) + } + out = append(out, r.Tag+"="+ref.Name) + } + return out + } + if got := refs("helper_prerouting"); !reflect.DeepEqual(got, []string{ + "conntrack:0:helper_prerouting=ftp", "conntrack:1:helper_prerouting=sip-tcp"}) { + t.Errorf("helper_prerouting = %v", got) + } + if got := refs("helper_output"); !reflect.DeepEqual(got, []string{ + "conntrack:0:helper_output=ftp", "conntrack:2:helper_output=tftp"}) { + t.Errorf("helper_output = %v", got) + } + + ftp := state.Rules["helper_prerouting"][0].Exprs + l4, _ := l4Matches("tcp", config.PortSpec{"21"}, nil) + if !reflect.DeepEqual(ftp[:len(ftp)-1], l4[0].exprs) { + t.Errorf("ftp rule does not match tcp dport 21: %#v", ftp) + } + + if _, err := compile(config.ConntrackRule{Action: config.ConntrackHelper, Helper: "nope"}); err == nil { + t.Error("expected error for unknown helper without proto") + } +} diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index e0a8167..81ee80e 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -20,15 +20,25 @@ type ManagedRule struct { type FirewallState struct { Rules map[string][]ManagedRule + // Helpers are the ct helper objects in kernel (insertion) order. + Helpers []Helper +} + +// Helper is a named ct helper object. +type Helper struct { + Name string `json:"name"` + Helper expr.CtHelper `json:"helper"` } type ChangeSet struct { - Add []ManagedRule - Remove []ManagedRule + Add []ManagedRule + Remove []ManagedRule + AddHelpers []Helper + RemoveHelpers []string } func (cs *ChangeSet) Empty() bool { - return len(cs.Add) == 0 && len(cs.Remove) == 0 + return len(cs.Add) == 0 && len(cs.Remove) == 0 && len(cs.AddHelpers) == 0 && len(cs.RemoveHelpers) == 0 } func (cs *ChangeSet) Summary() string { @@ -49,6 +59,12 @@ func (cs *ChangeSet) Summary() string { fmt.Fprintf(&b, " - [%s] %s (handle %d)\n", r.Chain, r.Tag, r.Handle) } } + for _, h := range cs.AddHelpers { + fmt.Fprintf(&b, " + ct helper %q\n", h.Name) + } + for _, n := range cs.RemoveHelpers { + fmt.Fprintf(&b, " - ct helper %q\n", n) + } return b.String() } @@ -56,7 +72,7 @@ func (cs *ChangeSet) Summary() string { // each chain, replace the middle, and insert the new rules before the first kept // suffix rule (or append when there is none). func computeDiff(current, desired *FirewallState) *ChangeSet { - cs := &ChangeSet{} + cs := diffHelpers(current, desired) chains := make([]string, 0, len(current.Rules)+len(desired.Rules)) for c := range current.Rules { @@ -109,7 +125,10 @@ func ruleEqual(a, b ManagedRule) bool { // restoreChangeSet replaces every managed rule in current with the snapshot's, // in snapshot order, so a restore cannot reorder rules. func restoreChangeSet(current, snap *FirewallState) *ChangeSet { - cs := &ChangeSet{} + cs := &ChangeSet{AddHelpers: snap.Helpers} + for _, h := range current.Helpers { + cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name) + } for _, rules := range current.Rules { for _, r := range rules { if r.Tag != "" { @@ -131,3 +150,29 @@ func restoreChangeSet(current, snap *FirewallState) *ChangeSet { } return cs } + +// diffHelpers replaces any ct helper object that is missing or differs. +// L3Proto is ignored: the kernel narrows inet to ip/ip6 for single-family helpers such as pptp. +func diffHelpers(current, desired *FirewallState) *ChangeSet { + cs := &ChangeSet{} + same := func(a, b expr.CtHelper) bool { return a.Name == b.Name && a.L4Proto == b.L4Proto } + find := func(hs []Helper, name string) (expr.CtHelper, bool) { + for _, h := range hs { + if h.Name == name { + return h.Helper, true + } + } + return expr.CtHelper{}, false + } + for _, h := range current.Helpers { + if want, ok := find(desired.Helpers, h.Name); !ok || !same(want, h.Helper) { + cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name) + } + } + for _, h := range desired.Helpers { + if have, ok := find(current.Helpers, h.Name); !ok || !same(have, h.Helper) { + cs.AddHelpers = append(cs.AddHelpers, h) + } + } + return cs +} diff --git a/internal/nftables/diff_test.go b/internal/nftables/diff_test.go index f0eff4c..482fda0 100644 --- a/internal/nftables/diff_test.go +++ b/internal/nftables/diff_test.go @@ -5,6 +5,7 @@ import ( "testing" "github.com/google/nftables/expr" + "golang.org/x/sys/unix" ) func TestRestoreChangeSet(t *testing.T) { @@ -64,3 +65,42 @@ func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) { t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add)) } } + +func TestDiffHelpers(t *testing.T) { + h := func(name, typ string, l3 uint16, l4 uint8) Helper { + return Helper{Name: name, Helper: expr.CtHelper{Name: typ, L3Proto: l3, L4Proto: l4}} + } + ftp := h("ftp", "ftp", unix.NFPROTO_INET, unix.IPPROTO_TCP) + tftp := h("tftp", "tftp", unix.NFPROTO_INET, unix.IPPROTO_UDP) + sipUDP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_UDP) + sipTCP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_TCP) + + current := &FirewallState{Helpers: []Helper{tftp, ftp, sipUDP}} + desired := &FirewallState{Helpers: []Helper{ftp, sipTCP}} + + cs := computeDiff(current, desired) + if !reflect.DeepEqual(cs.RemoveHelpers, []string{"tftp", "sip"}) { + t.Errorf("remove = %v", cs.RemoveHelpers) + } + if !reflect.DeepEqual(cs.AddHelpers, []Helper{sipTCP}) { + t.Errorf("add = %v", cs.AddHelpers) + } + + // Restore recreates every helper so kernel listing order matches the snapshot. + cs = restoreChangeSet(desired, current) + if !reflect.DeepEqual(cs.RemoveHelpers, []string{"ftp", "sip"}) || !reflect.DeepEqual(cs.AddHelpers, current.Helpers) { + t.Errorf("restore = -%v +%v", cs.RemoveHelpers, cs.AddHelpers) + } + + live := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_IPV4, unix.IPPROTO_TCP)}} + want := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_INET, unix.IPPROTO_TCP)}} + if cs := computeDiff(live, want); !cs.Empty() { + t.Errorf("kernel-narrowed l3proto should not diff: %+v", cs) + } + if cs := computeDiff(desired, desired); !cs.Empty() { + t.Errorf("identical helpers should be empty: %+v", cs) + } + if cs := computeDiff(&FirewallState{}, desired); cs.Empty() { + t.Error("missing helpers should not be empty") + } +} diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 77a4afe..b783e6a 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -4,6 +4,7 @@ import ( "fmt" "github.com/google/nftables" + "github.com/google/nftables/expr" "git.unkin.net/unkin/tomswall/internal/config" ) @@ -69,6 +70,20 @@ func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftable 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 { @@ -112,6 +127,13 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol }) } + 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 { @@ -172,6 +194,18 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { 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) @@ -205,6 +239,7 @@ type Snapshot struct { 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. @@ -241,6 +276,7 @@ func (e *Engine) Snapshot() (*Snapshot, error) { if err != nil { return nil, err } + snap.Helpers = state.Helpers return snap, nil } @@ -257,6 +293,7 @@ func (e *Engine) Restore(s *Snapshot) error { if err != nil { return err } + want.Helpers = s.Helpers current, err := e.readCurrentState() if err != nil { return err @@ -264,6 +301,10 @@ func (e *Engine) Restore(s *Snapshot) error { 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 }