diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 4f3f430..a86c5b5 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,6 +213,46 @@ 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", +} + +func helperObjName(helper, proto string) string { + if proto != helperProtos[helper] { + return helper + "-" + proto + } + return helper +} + +// expandHelper declares a ct helper object per (helper, proto) and returns one entry per proto. +func expandHelper(state *FirewallState, ct config.ConntrackRule) ([]config.ConntrackRule, error) { + proto := ct.Proto + if proto == "" { + if proto = helperProtos[ct.Helper]; proto == "" { + return nil, fmt.Errorf("proto required for helper %q", ct.Helper) + } + } + var out []config.ConntrackRule + for _, p := range strings.Split(proto, ",") { + p = strings.TrimSpace(p) + n, err := protoNumber(p) + if err != nil { + return nil, err + } + name := helperObjName(ct.Helper, 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}}) + } + pct := ct + pct.Proto = p + out = append(out, pct) + } + return out, nil +} + func (c *Compiler) compileConntrack(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for i, ct := range c.cfg.Conntrack { @@ -219,6 +260,16 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { if config.HasZoneExclusion(ct.Source) || config.HasZoneExclusion(ct.Dest) { return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i) } + cts := []config.ConntrackRule{ct} + if ct.Action == config.ConntrackHelper { + if ct.Chain == "" { + ct.Chain = config.ConntrackBoth + } + var err error + if cts, err = expandHelper(state, ct); err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + } srcs, dsts := c.zoneSpecs(ct.Source), c.zoneSpecs(ct.Dest) if len(srcs) == 0 { srcs = []config.ZoneSpec{{}} @@ -246,8 +297,10 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { for _, dst := range dsts { for _, dstAddr := range splitAddrs(dst.Addr) { for _, chain := range chains { - if err := c.compileConntrackPair(state, tag, chain, ct, src.Zone, srcAddr, dst.Zone, dstAddr); err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) + for _, pct := range cts { + if err := c.compileConntrackPair(state, tag, chain, pct, src.Zone, srcAddr, dst.Zone, dstAddr); err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } } } } @@ -259,11 +312,10 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { } // compileConntrackPair matches iif of the source zone in raw_prerouting and oif of the dest zone in raw_output. +// Helpers are assigned after conntrack (-200) has created the entry, so they go to the mangle-priority +// helper_* chain instead; a raw-priority assignment is a no-op. func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule, srcZone, srcAddr, dstZone, dstAddr string) error { - if ct.Action == config.ConntrackHelper { - return nil - } if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" && (dstAddr == "" || strings.HasPrefix(dstAddr, "!")) { return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone) @@ -275,6 +327,10 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, if chain == "raw_output" { srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr) } + out := chain + if ct.Action == config.ConntrackHelper { + out = "helper_" + strings.TrimPrefix(chain, "raw_") + } for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { matches, err := c.buildMatchExprs(srcIface, dstIface, chain, ct.Proto, ct.DPort, ct.SPort, srcAddr, dstAddr) @@ -288,11 +344,13 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, exprs = append(exprs, &expr.Notrack{}) case config.ConntrackDrop: exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) + case config.ConntrackHelper: + exprs = append(exprs, &expr.Objref{Type: unix.NFT_OBJECT_CT_HELPER, Name: helperObjName(ct.Helper, ct.Proto)}) } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, + state.Rules[out] = append(state.Rules[out], ManagedRule{ + Chain: out, Exprs: exprs, - Tag: tag + ":" + chain, + Tag: tag + ":" + out, }) } } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index d0eae9b..7be75de 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -2424,6 +2424,66 @@ func TestCompile_OrigDestForwardRejected(t *testing.T) { } } +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") + } +} + func TestCompile_ConntrackZones(t *testing.T) { tests := []struct { name string @@ -2657,3 +2717,85 @@ func TestCompile_RuleZoneExclusionIntraZoneSymmetric(t *testing.T) { } } } + +func TestCompile_ConntrackHelperZones(t *testing.T) { + tests := []struct { + name string + ct config.ConntrackRule + want map[string][]string + wantErr string + }{ + { + name: "non-fw source is prerouting only", + ct: config.ConntrackRule{Source: "net", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {"iif=eth0"}}, + }, + { + name: "fw source is output only with dest oif", + ct: config.ConntrackRule{Source: "fw", Dest: "lan", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_output": {"oif=eth1"}}, + }, + { + name: "any source is global in both chains", + ct: config.ConntrackRule{Source: "any", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {""}, "helper_output": {""}}, + }, + { + name: "dest zone with address in prerouting", + ct: config.ConntrackRule{Source: "net", Dest: "lan:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, + }, + { + name: "dest zone without address is rejected in prerouting", + ct: config.ConntrackRule{Source: "net", Dest: "lan"}, + wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, + }, + { + name: "chain output with non-fw source is rejected", + ct: config.ConntrackRule{Source: "net", Chain: config.ConntrackOutput}, + wantErr: "chain output needs SOURCE fw", + }, + { + name: "interface-less zone fails closed", + ct: config.ConntrackRule{Source: "dmz"}, + want: map[string][]string{}, + }, + { + name: "unknown zone fails closed", + ct: config.ConntrackRule{Source: "nte"}, + want: map[string][]string{}, + }, + { + name: "exclusion rejected", + ct: config.ConntrackRule{Source: "all!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.ct.Action, tt.ct.Helper = config.ConntrackHelper, "ftp" + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Zones["dmz"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Conntrack = []config.ConntrackRule{tt.ct} + }) + if tt.wantErr != "" { + if _, err := NewCompiler(cfg).Compile(); err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("Compile() error = %v, want %q", err, tt.wantErr) + } + return + } + state := mustCompile(t, cfg) + got := map[string][]string{} + for _, chain := range []string{"helper_prerouting", "helper_output"} { + for _, r := range taggedRules(state, chain, "conntrack:0:"+chain) { + got[chain] = append(got[chain], describeRule(r)) + } + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %v, want %v", got, tt.want) + } + }) + } +} 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 5444970..c6781c6 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, + }, "raw_prerouting": { Name: "raw_prerouting", Table: table, @@ -128,6 +143,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 { @@ -188,6 +210,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) @@ -221,6 +255,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. @@ -257,6 +292,7 @@ func (e *Engine) Snapshot() (*Snapshot, error) { if err != nil { return nil, err } + snap.Helpers = state.Helpers return snap, nil } @@ -273,6 +309,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 @@ -280,6 +317,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 } diff --git a/tomswall.example.yaml b/tomswall.example.yaml index bc3e3b8..b48acf4 100644 --- a/tomswall.example.yaml +++ b/tomswall.example.yaml @@ -197,7 +197,6 @@ snat: # comment: "Skip conntrack for DNS" # - action: helper # source: loc -# dest: net # proto: tcp # dport: [21] # helper: ftp