From 42a4dab6a3d24a72085108b42cc58c29f6872074 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Fri, 9 Oct 2026 23:39:23 +1100 Subject: [PATCH] Emit one family guard per rule, none in ip/ip6 tables --- internal/nftables/compiler.go | 96 +++++++++++++++++++++++++++--- internal/nftables/compiler_test.go | 78 ++++++++++++------------ internal/nftables/guards_test.go | 92 ++++++++++++++++++++++++++++ 3 files changed, 218 insertions(+), 48 deletions(-) create mode 100644 internal/nftables/guards_test.go diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 53d5b55..0af760c 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -67,10 +67,70 @@ func (c *Compiler) Compile() (*FirewallState, error) { } c.compileMSSClamp(state) limitLogs(state, c.cfg.Settings.LogLimit) + familyGuards(state, c.cfg.Settings.AddressFamily) return state, nil } +// familyGuards leaves each rule at most one meta nfproto guard, ahead of its first network-header +// payload as nft emits it, and drops rules whose guards contradict; nft list cannot decode +// conflicting guards. An ip/ip6 table is its own guard: guards go, other-family rules are dropped. +func familyGuards(state *FirewallState, family config.AddressFamily) { + table := map[config.AddressFamily]byte{config.FamilyIP: unix.NFPROTO_IPV4, config.FamilyIP6: unix.NFPROTO_IPV6}[family] + for chain, rules := range state.Rules { + var out []ManagedRule + for _, r := range rules { + if exprs, ok := normalizeGuards(r.Exprs, table); ok { + r.Exprs = exprs + out = append(out, r) + } + } + state.Rules[chain] = out + } +} + +func normalizeGuards(in []expr.Any, table byte) ([]expr.Any, bool) { + var fam byte + at := -1 + out := make([]expr.Any, 0, len(in)) + for i := 0; i < len(in); i++ { + if p := nfprotoGuard(in, i); p != 0 { + if (fam != 0 && p != fam) || (table != 0 && p != table) { + return nil, false + } + if fam == 0 { + fam, at = p, len(out) + } + i++ + continue + } + out = append(out, in[i]) + } + if fam == 0 || table != 0 { + return out, true + } + if l3 := slices.IndexFunc(out, func(e expr.Any) bool { + p, ok := e.(*expr.Payload) + return ok && p.Base == expr.PayloadBaseNetworkHeader + }); l3 >= 0 && l3 < at { + at = l3 + } + return slices.Insert(out, at, matchNFProto(fam)...), true +} + +// nfprotoGuard is the family a meta nfproto == match at in[i] guards for, else 0. +func nfprotoGuard(in []expr.Any, i int) byte { + m, ok := in[i].(*expr.Meta) + if !ok || m.Key != expr.MetaKeyNFPROTO || i+1 >= len(in) { + return 0 + } + c, ok := in[i+1].(*expr.Cmp) + if !ok || c.Op != expr.CmpOpEq || c.Register != m.Register || len(c.Data) != 1 { + return 0 + } + return c.Data[0] +} + // limitLogs puts a limit in front of every log expression. A limit stops the // whole rule, so like shorewall's separate LOG rule, a log followed by an action // splits into a limited log-only rule and the same rule without the log. A @@ -1697,6 +1757,7 @@ func matchTCPFlags(flags, mask byte) []expr.Any { func matchSmurfDrop(iface string) []expr.Any { var exprs []expr.Any exprs = append(exprs, matchIfaceName(true, iface)...) + exprs = append(exprs, matchNFProto(unix.NFPROTO_IPV4)...) exprs = append(exprs, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, &expr.Bitwise{ @@ -1857,32 +1918,49 @@ func parseSPortOrRange(s string) ([]expr.Any, error) { } func matchSourceCIDR(cidr string) ([]expr.Any, error) { - return matchAddrCIDR(cidr, true) + return matchGuardedCIDR(cidr, true) } func matchDestCIDR(cidr string) ([]expr.Any, error) { - return matchAddrCIDR(cidr, false) + return matchGuardedCIDR(cidr, false) } -// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table. -func matchOrigDest(addr string) ([]expr.Any, error) { +// matchGuardedCIDR guards an address match with its family's nfproto, unless the list mixes families. +func matchGuardedCIDR(cidr string, isSrc bool) ([]expr.Any, error) { + e, err := matchAddrCIDR(cidr, isSrc) + if err != nil { + return nil, err + } + if p := addrFamily(cidr); p != 0 { + return append(matchNFProto(p), e...), nil + } + return e, nil +} + +// addrFamily is the NFPROTO shared by every address in a (negated) comma list, or 0 if they mix. +func addrFamily(list string) byte { var proto byte - for i, a := range strings.Split(strings.TrimPrefix(addr, "!"), ",") { + for i, a := range strings.Split(strings.TrimPrefix(list, "!"), ",") { a, _, _ = strings.Cut(a, "/") p := byte(unix.NFPROTO_IPV6) if ip := net.ParseIP(a); ip != nil && ip.To4() != nil { p = unix.NFPROTO_IPV4 } if i > 0 && p != proto { - return nil, fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr) + return 0 } proto = p } + return proto +} + +// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table. +func matchOrigDest(addr string) ([]expr.Any, error) { dst, err := matchDestCIDR(addr) - if err != nil { - return nil, err + if err == nil && addrFamily(addr) == 0 { + err = fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr) } - return append(matchNFProto(proto), dst...), nil + return dst, err } func matchNFProto(proto byte) []expr.Any { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index dccbc3a..1bed22b 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -210,7 +210,7 @@ func TestMatchSourceCIDR_IPv6(t *testing.T) { } for _, tt := range tests { - exprs, err := matchSourceCIDR(tt.input) + exprs, err := matchAddrCIDR(tt.input, true) if tt.wantErr { if err == nil { t.Errorf("matchSourceCIDR(%q) should fail", tt.input) @@ -241,7 +241,7 @@ func TestMatchDestCIDR_IPv6(t *testing.T) { } for _, tt := range tests { - exprs, err := matchDestCIDR(tt.input) + exprs, err := matchAddrCIDR(tt.input, false) if tt.wantErr { if err == nil { t.Errorf("matchDestCIDR(%q) should fail", tt.input) @@ -961,7 +961,7 @@ func TestCompile_RateLimit(t *testing.T) { } func TestNegatedAddress(t *testing.T) { - exprs, err := matchSourceCIDR("!192.168.1.0/24") + exprs, err := matchAddrCIDR("!192.168.1.0/24", true) if err != nil { t.Fatalf("matchSourceCIDR(!192.168.1.0/24) error: %v", err) } @@ -973,7 +973,7 @@ func TestNegatedAddress(t *testing.T) { t.Errorf("negated address should use CmpOpNeq, got %v", cmp.Op) } - exprs, err = matchDestCIDR("!10.0.0.1") + exprs, err = matchAddrCIDR("!10.0.0.1", false) if err != nil { t.Fatalf("matchDestCIDR(!10.0.0.1) error: %v", err) } @@ -1968,57 +1968,57 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "address list after colon belongs to one zone", rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"}, - want: map[string][]string{"forward": {"iif=eth1 oif=eth0 daddr=192.0.2.1", "iif=eth1 oif=eth0 daddr=198.51.100.1"}}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth0 ip4 daddr=192.0.2.1", "iif=eth1 oif=eth0 ip4 daddr=198.51.100.1"}}, }, { name: "zone:address inside a list", rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"}, - want: map[string][]string{"input": {"iif=eth1", "iif=eth2 saddr=203.0.113.7"}}, + want: map[string][]string{"input": {"iif=eth1", "iif=eth2 ip4 saddr=203.0.113.7"}}, }, { name: "dnat source list", rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}, - "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10"}}, + "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat source list skips the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "net,svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, - want: map[string][]string{"prerouting": {"iif=eth0"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10"}}, + want: map[string][]string{"prerouting": {"iif=eth0"}, "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat lone source zone may equal the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, - want: map[string][]string{"prerouting": {"iif=eth2"}, "forward": {"iif=eth2 oif=eth2 daddr=192.0.2.10"}}, + want: map[string][]string{"prerouting": {"iif=eth2"}, "forward": {"iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat exclusion source skips the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all!fw,anycast", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth1", "iif=eth0"}, - "forward": {"iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, + "forward": {"iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat intrazone exclusion source keeps the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all+!fw,anycast,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth2"}, - "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, + "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat all source expands per zone and skips fw and the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0"}, - "forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, + "forward": {"iif=eth3 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat any+ source keeps the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "any+", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0", "iif=eth2"}, - "forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, + "forward": {"iif=eth3 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}}, }, { name: "dnat to fw accepts in input", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.1", Proto: "tcp", DPort: config.PortSpec{"80"}}, - want: map[string][]string{"prerouting": {"iif=eth0"}, "input": {"iif=eth0 daddr=192.0.2.1"}}, + want: map[string][]string{"prerouting": {"iif=eth0"}, "input": {"iif=eth0 ip4 daddr=192.0.2.1"}}, }, { name: "redirect accepts in input without daddr", @@ -2028,13 +2028,13 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "dnat source address list", rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, - want: map[string][]string{"prerouting": {"iif=eth0 saddr=192.0.2.5", "iif=eth0 saddr=198.51.100.5"}, - "forward": {"iif=eth0 oif=eth2 saddr=192.0.2.5 daddr=192.0.2.10", "iif=eth0 oif=eth2 saddr=198.51.100.5 daddr=192.0.2.10"}}, + want: map[string][]string{"prerouting": {"iif=eth0 ip4 saddr=192.0.2.5", "iif=eth0 ip4 saddr=198.51.100.5"}, + "forward": {"iif=eth0 oif=eth2 ip4 saddr=192.0.2.5 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 saddr=198.51.100.5 daddr=192.0.2.10"}}, }, { name: "negated address list stays one AND-ed rule", rule: config.Rule{Action: config.RuleAccept, Source: "net:!192.0.2.5,198.51.100.5", Dest: "fw"}, - want: map[string][]string{"input": {"iif=eth0 !saddr=192.0.2.5 !saddr=198.51.100.5"}}, + want: map[string][]string{"input": {"iif=eth0 ip4 !saddr=192.0.2.5 !saddr=198.51.100.5"}}, }, { name: "zone named like all/any keyword is a plain zone", @@ -2049,7 +2049,7 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "interface-less zone kept when address narrows it", rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn:192.0.2.1"}, - want: map[string][]string{"forward": {"iif=eth1 daddr=192.0.2.1"}}, + want: map[string][]string{"forward": {"iif=eth1 ip4 daddr=192.0.2.1"}}, }, { name: "fw source matches dest zone oif", @@ -2059,25 +2059,25 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "fw to all has no oif", rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "all:192.0.2.1"}, - want: map[string][]string{"output": {"daddr=192.0.2.1"}}, + want: map[string][]string{"output": {"ip4 daddr=192.0.2.1"}}, }, { name: "dnat origdest", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5"}, - "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, + "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.17"}}, }, { name: "dnat origdest list", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5,203.0.113.6"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5", "iif=eth0 ip4 daddr=203.0.113.6"}, - "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, + "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.17"}}, }, { name: "dnat negated origdest list", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,203.0.113.6"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 !daddr=203.0.113.5 !daddr=203.0.113.6"}, - "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, + "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.17"}}, }, { name: "accept origdest", @@ -2487,7 +2487,7 @@ func TestCompile_RejectPerProto(t *testing.T) { } func TestNegatedAddressList(t *testing.T) { - exprs, err := matchDestCIDR("!192.0.2.1,198.51.100.1") + exprs, err := matchAddrCIDR("!192.0.2.1,198.51.100.1", false) if err != nil { t.Fatalf("matchDestCIDR error: %v", err) } @@ -2713,7 +2713,7 @@ func TestCompile_ConntrackZones(t *testing.T) { { name: "source and dest addresses", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:192.0.2.1,198.51.100.1", Dest: "fw:203.0.113.1"}, - want: map[string][]string{"raw_prerouting": {"iif=eth0 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 saddr=198.51.100.1 daddr=203.0.113.1"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 ip4 saddr=198.51.100.1 daddr=203.0.113.1"}}, }, { name: "fw source goes to raw_output with dest oif", @@ -2723,7 +2723,7 @@ func TestCompile_ConntrackZones(t *testing.T) { { name: "all matches no interface", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Dest: "fw:192.0.2.53"}, - want: map[string][]string{"raw_prerouting": {"daddr=192.0.2.53"}}, + want: map[string][]string{"raw_prerouting": {"ip4 daddr=192.0.2.53"}}, }, { name: "interface-less zone fails closed", @@ -2783,7 +2783,7 @@ func TestCompile_ConntrackZones(t *testing.T) { { name: "negated addresses stay one AND-ed match", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"}, - want: map[string][]string{"raw_prerouting": {"iif=eth0 !saddr=192.0.2.1 !saddr=198.51.100.1"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 !saddr=192.0.2.1 !saddr=198.51.100.1"}}, }, { name: "sport", @@ -2803,7 +2803,7 @@ func TestCompile_ConntrackZones(t *testing.T) { { name: "fw dest zone with address matches daddr", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw:192.0.2.1"}, - want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=192.0.2.1"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 daddr=192.0.2.1"}}, }, { name: "unknown zone fails closed", @@ -2858,7 +2858,7 @@ func TestCompile_ConntrackZones(t *testing.T) { { name: "dest zone with address matches daddr", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"}, - want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 daddr=203.0.113.10"}}, }, } for _, tt := range tests { @@ -2957,7 +2957,7 @@ func TestCompile_ConntrackHelperZones(t *testing.T) { { 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"}}, + want: map[string][]string{"helper_prerouting": {"iif=eth0 ip4 daddr=203.0.113.10"}}, }, { name: "dest zone without address is rejected in prerouting", @@ -3023,7 +3023,7 @@ func TestCompile_AllIncludesFirewallMatches(t *testing.T) { { name: "all address kept on added fw rules", rule: config.Rule{Action: config.RuleAccept, Source: "all:192.0.2.5", Dest: "all"}, - want: map[string][]string{"input": {"saddr=192.0.2.5"}, "output": {"saddr=192.0.2.5"}, "forward": {"saddr=192.0.2.5"}}, + want: map[string][]string{"input": {"ip4 saddr=192.0.2.5"}, "output": {"ip4 saddr=192.0.2.5"}, "forward": {"ip4 saddr=192.0.2.5"}}, }, { name: "dnat with all source skips fw", @@ -3607,17 +3607,17 @@ func TestCompile_DestPlusOverridesIntraZone(t *testing.T) { func TestCompile_HostsIntraZone(t *testing.T) { state := mustCompile(t, hostsCfg(nil)) want := []string{ - "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24", - "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24", + "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 daddr=198.51.100.0/24", + "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 daddr=192.0.2.0/24", } if got := describeTagged(state, "forward", "intra:lan"); !reflect.DeepEqual(got, want) { t.Errorf("lan intra = %q, want %q", got, want) } want = []string{ - "iif=wlo1 ip4 !saddr=192.0.2.0/24 oif=enp2s0 ip4 !daddr=198.51.100.0/24", - "iif=wlo1 ip6 oif=enp2s0 ip6", - "iif=enp2s0 ip4 !saddr=198.51.100.0/24 oif=wlo1 ip4 !daddr=192.0.2.0/24", - "iif=enp2s0 ip6 oif=wlo1 ip6", + "iif=wlo1 ip4 !saddr=192.0.2.0/24 oif=enp2s0 !daddr=198.51.100.0/24", + "iif=wlo1 ip6 oif=enp2s0", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24 oif=wlo1 !daddr=192.0.2.0/24", + "iif=enp2s0 ip6 oif=wlo1", } if got := describeTagged(state, "forward", "intra:net"); !reflect.DeepEqual(got, want) { t.Errorf("net intra = %q, want %q", got, want) @@ -3630,8 +3630,8 @@ func TestCompile_HostsIntraZone(t *testing.T) { t.Errorf("explicit lan lan policy must replace implicit accept, got %q", got) } want = []string{ - "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24", - "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24", + "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 daddr=198.51.100.0/24", + "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 daddr=192.0.2.0/24", } if got := describeTagged(state, "forward", "policy:0"); !reflect.DeepEqual(got, want) { t.Errorf("lan lan policy = %q, want %q", got, want) @@ -3650,7 +3650,7 @@ func TestCompile_HostsRouteBack(t *testing.T) { if got := describeTagged(mustCompile(t, hostsCfg(sameIface(false))), "forward", "intra:lan"); len(got) != 0 { t.Errorf("same-interface hosts without routeback = %q, want none", got) } - want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24"} + want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=wlo1 daddr=192.0.2.0/24"} if got := describeTagged(mustCompile(t, hostsCfg(sameIface(true))), "forward", "intra:lan"); !reflect.DeepEqual(got, want) { t.Errorf("routeback hosts intra = %q, want %q", got, want) } diff --git a/internal/nftables/guards_test.go b/internal/nftables/guards_test.go new file mode 100644 index 0000000..670b364 --- /dev/null +++ b/internal/nftables/guards_test.go @@ -0,0 +1,92 @@ +package nftables + +import ( + "os" + "os/exec" + "testing" + + "github.com/google/nftables/expr" + + "git.unkin.net/unkin/tomswall/internal/config" +) + +func guardCfg(af config.AddressFamily) *config.Config { + return hostsCfg(func(c *config.Config) { + c.Settings.AddressFamily = af + c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64") + c.Hosts[0].Exclusions = []string{"192.0.2.9", "2001:db8::9"} + c.Interfaces[0].Options.NoSmurfs = true + c.Rules = append(c.Rules, + config.Rule{Action: config.RuleDrop, Source: "net:203.0.113.7", Dest: "lan:192.0.2.5"}, + config.Rule{Action: config.RuleDrop, Source: "net:2001:db8:1::7", Dest: "lan:2001:db8::5"}) + }) +} + +func TestCompile_FamilyGuardsDecodable(t *testing.T) { + addrLen := map[config.AddressFamily]uint32{config.FamilyIP: 4, config.FamilyIP6: 16} + for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} { + t.Run(string(af), func(t *testing.T) { + for chain, rules := range mustCompile(t, guardCfg(af)).Rules { + for _, r := range rules { + guards, l3 := 0, false + for i, e := range r.Exprs { + if nfprotoGuard(r.Exprs, i) != 0 { + guards++ + if l3 { + t.Errorf("%s %s: guard after a network payload: %s", chain, r.Tag, describeRule(r)) + } + } + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader { + l3 = true + if n := addrLen[af]; n != 0 && (p.Len == 4 || p.Len == 16) && p.Len != n { + t.Errorf("%s %s: other-family address in %s table: %s", chain, r.Tag, af, describeRule(r)) + } + } + } + if max := map[bool]int{true: 1, false: 0}[af == config.FamilyINET]; guards > max { + t.Errorf("%s %s: %d family guards in %s table: %s", chain, r.Tag, guards, af, describeRule(r)) + } + } + } + }) + } +} + +// TestNetnsNftListDecodes applies each family in a fresh user+net namespace and requires +// nft(8) to list the ruleset and a second plan to be empty. Needs unshare and nft. +func TestNetnsNftListDecodes(t *testing.T) { + if af := os.Getenv("TOMSWALL_NETNS_CHILD"); af != "" { + netnsChild(t, config.AddressFamily(af)) + return + } + if os.Getenv("TOMSWALL_NETNS_TEST") == "" { + t.Skip("set TOMSWALL_NETNS_TEST=1 to run (needs unshare and nft)") + } + for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} { + cmd := exec.Command("unshare", "-rn", os.Args[0], "-test.run=^TestNetnsNftListDecodes$", "-test.v") + cmd.Env = append(os.Environ(), "TOMSWALL_NETNS_CHILD="+string(af)) + if out, err := cmd.CombinedOutput(); err != nil { + t.Errorf("%s: %v\n%s", af, err, out) + } + } +} + +func netnsChild(t *testing.T, af config.AddressFamily) { + e, err := NewEngine(guardCfg(af)) + if err != nil { + t.Fatal(err) + } + cs, err := e.Plan() + if err != nil { + t.Fatal(err) + } + if err := e.Apply(cs); err != nil { + t.Fatal(err) + } + if out, err := exec.Command("nft", "list", "ruleset").CombinedOutput(); err != nil { + t.Fatalf("nft list ruleset: %v\n%s", err, out) + } + if cs, err = e.Plan(); err != nil || len(cs.Add)+len(cs.Remove) != 0 { + t.Fatalf("second plan not empty: %d add, %d remove, err %v", len(cs.Add), len(cs.Remove), err) + } +}