diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 77d2e81..900eecb 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -1212,9 +1212,12 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { // zoneMatch classifies a packet into a zone: an interface (empty: any) and, for a hosts entry, one host // address; excl carves out hosts exclusions and the hosts of sub-zones, which shorewall matches first. +// fam (an NFPROTO; addr's family when set, else 0: any) guards addr and excl so IPv4 offsets are +// never compared against IPv6 bytes. type zoneMatch struct { iface, addr string excl []string + fam byte } // resolveZone returns nil (fail closed) for an unknown zone, or one with neither interfaces nor hosts @@ -1236,7 +1239,13 @@ func (c *Compiler) resolveZone(zone, addr string) []zoneMatch { hasHosts := false for _, iface := range c.cfg.ZoneInterfaces(zone) { sub, back := c.subZoneHosts(zone, iface) - out = append(out, zoneMatch{iface: iface, excl: sub}) + if len(sub) == 0 { + out = append(out, zoneMatch{iface: iface}) + } else { + v4, v6 := splitFamily(sub) + out = append(out, zoneMatch{iface: iface, excl: v4, fam: unix.NFPROTO_IPV4}, + zoneMatch{iface: iface, excl: v6, fam: unix.NFPROTO_IPV6}) + } for _, b := range back { if addrsOverlap(b, addr) { out = append(out, zoneMatch{iface: iface, addr: b}) @@ -1249,9 +1258,14 @@ func (c *Compiler) resolveZone(zone, addr string) []zoneMatch { } hasHosts = true sub, _ := c.subZoneHosts(zone, h.Interface) + v4, v6 := splitFamily(slices.Concat(h.Exclusions, sub)) for _, a := range h.Addresses { if addrsOverlap(a, addr) { - out = append(out, zoneMatch{iface: h.Interface, addr: a, excl: slices.Concat(h.Exclusions, sub)}) + m := zoneMatch{iface: h.Interface, addr: a, excl: v6} + if p, err := parsePrefix(a); err == nil && p.Addr().Is4() { + m.excl = v4 + } + out = append(out, m) } } } @@ -1300,6 +1314,18 @@ func addrsOverlap(host, rule string) bool { return false } +// splitFamily partitions addresses by family; unparsable ones go to v6 so zoneMatchExprs still rejects them. +func splitFamily(addrs []string) (v4, v6 []string) { + for _, a := range addrs { + if p, err := parsePrefix(a); err == nil && p.Addr().Is4() { + v4 = append(v4, a) + } else { + v6 = append(v6, a) + } + } + return v4, v6 +} + func parsePrefix(s string) (netip.Prefix, error) { if a, err := netip.ParseAddr(s); err == nil { return netip.PrefixFrom(a, a.BitLen()), nil @@ -1307,29 +1333,31 @@ func parsePrefix(s string) (netip.Prefix, error) { return netip.ParsePrefix(s) } -// zoneMatchExprs matches a zone on the in (src) or out interface plus its host address, guarded by -// the address family so an IPv4 host never matches IPv6 bytes in an inet table. -// ponytail: exclusions are unguarded, so an IPv6 packet whose bytes hit an IPv4 exclusion skips the zone. +// zoneMatchExprs matches a zone on the in (src) or out interface plus its host address and exclusions, +// all guarded by m.fam so an IPv4 address never matches IPv6 bytes in an inet table. func zoneMatchExprs(m zoneMatch, src bool) ([]expr.Any, error) { var out []expr.Any if m.iface != "" { out = matchIfaceName(src, m.iface) } + var addr []expr.Any if m.addr != "" { p, err := parsePrefix(m.addr) if err != nil { return nil, fmt.Errorf("invalid host address %q", m.addr) } - e, err := matchAddrCIDR(m.addr, src) - if err != nil { + if addr, err = matchAddrCIDR(m.addr, src); err != nil { return nil, err } - fam := byte(unix.NFPROTO_IPV6) + m.fam = unix.NFPROTO_IPV6 if p.Addr().Is4() { - fam = unix.NFPROTO_IPV4 + m.fam = unix.NFPROTO_IPV4 } - out = append(append(out, matchNFProto(fam)...), e...) } + if m.fam != 0 { + out = append(out, matchNFProto(m.fam)...) + } + out = append(out, addr...) if len(m.excl) > 0 { e, err := matchAddrCIDR("!"+strings.Join(m.excl, ","), src) if err != nil { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 108ca25..dad2976 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1924,6 +1924,8 @@ func describeRule(r ManagedRule) string { parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00")) case expr.MetaKeyOIFNAME: parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) + case expr.MetaKeyNFPROTO: + parts = append(parts, map[byte]string{unix.NFPROTO_IPV4: "ip4", unix.NFPROTO_IPV6: "ip6"}[cmp.Data[0]]) } case *expr.Payload: if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) { @@ -2062,25 +2064,25 @@ func TestCompile_CommaZoneLists(t *testing.T) { { 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 daddr=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"}}, }, { 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 daddr=203.0.113.5", "iif=eth0 daddr=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"}}, }, { 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 !daddr=203.0.113.5 !daddr=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"}}, }, { name: "accept origdest", rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"}, - want: map[string][]string{"input": {"iif=eth0 daddr=203.0.113.5"}}, + want: map[string][]string{"input": {"iif=eth0 ip4 daddr=203.0.113.5"}}, }, { name: "origdest does not scope interface-less zone", @@ -2090,7 +2092,7 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "accept ipv6 origdest", rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"}, - want: map[string][]string{"input": {"iif=eth0 daddr=2001:db8::5"}}, + want: map[string][]string{"input": {"iif=eth0 ip6 daddr=2001:db8::5"}}, }, { name: "blrule zone list", @@ -3379,7 +3381,7 @@ func TestCompile_HostsZoneRules(t *testing.T) { defer slog.SetDefault(prev) state := mustCompile(t, hostsCfg(nil)) - want := []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24", "iif=tun0"} + want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24", "iif=tun0"} if got := describeTagged(state, "input", "rule:0"); !reflect.DeepEqual(got, want) { t.Errorf("input rule:0 = %q, want %q", got, want) } @@ -3391,13 +3393,13 @@ func TestCompile_HostsZoneRules(t *testing.T) { func TestCompile_HostsSubZoneBeforeParent(t *testing.T) { state := mustCompile(t, hostsCfg(nil)) want := []string{ - "iif=wlo1 !saddr=192.0.2.0/24", - "iif=enp2s0 !saddr=198.51.100.0/24", + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", } if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) { t.Errorf("net->fw policy = %q, want %q", got, want) } - want = []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24"} + want = []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"} if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) { t.Errorf("lan->fw policy = %q, want %q", got, want) } @@ -3421,33 +3423,33 @@ func TestCompile_HostsZoneMatches(t *testing.T) { }{ {"forward to hosts zone", func(c *config.Config) { c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "vpn", Dest: "lan", Proto: "tcp"}} - }, "forward", "rule:0", []string{"iif=tun0 oif=wlo1 daddr=192.0.2.0/24", "iif=tun0 oif=enp2s0 daddr=198.51.100.0/24"}}, + }, "forward", "rule:0", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24", "iif=tun0 oif=enp2s0 ip4 daddr=198.51.100.0/24"}}, {"rule address prunes non-overlapping hosts", func(c *config.Config) { c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "lan:192.0.2.5", Dest: "fw", Proto: "tcp"}} - }, "input", "rule:0", []string{"iif=wlo1 saddr=192.0.2.0/24 saddr=192.0.2.5"}}, + }, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 saddr=192.0.2.5"}}, {"host exclusions and sub-zone exclusions fall back to the parent", func(c *config.Config) { c.Hosts[1].Exclusions = []string{"198.51.100.7"} c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp"}} }, "input", "rule:0", []string{ - "iif=wlo1 !saddr=192.0.2.0/24", - "iif=enp2s0 !saddr=198.51.100.0/24", - "iif=enp2s0 saddr=198.51.100.7", + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", + "iif=enp2s0 ip4 saddr=198.51.100.7", }}, {"host exclusion", func(c *config.Config) { c.Hosts[1].Exclusions = []string{"198.51.100.7"} - }, "input", "policy:1", []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24 !saddr=198.51.100.7"}}, + }, "input", "policy:1", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24 !saddr=198.51.100.7"}}, {"DNAT from hosts zone", func(c *config.Config) { c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "lan", Dest: "vpn:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"80"}}} - }, "prerouting", "rule:0", []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24"}}, + }, "prerouting", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, {"conntrack from hosts zone", func(c *config.Config) { c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Source: "lan", Proto: "udp"}} - }, "raw_prerouting", "conntrack:0:raw_prerouting", []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24"}}, + }, "raw_prerouting", "conntrack:0:raw_prerouting", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, {"blrule from hosts zone", func(c *config.Config) { c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "lan", Dest: "all"}} - }, "input", "blrule:0", []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24"}}, + }, "input", "blrule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, {"all expansion includes hosts zone", func(c *config.Config) { c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "all!net,vpn", Dest: "fw", Proto: "tcp"}} - }, "input", "rule:0", []string{"iif=wlo1 saddr=192.0.2.0/24", "iif=enp2s0 saddr=198.51.100.0/24"}}, + }, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -3458,3 +3460,42 @@ func TestCompile_HostsZoneMatches(t *testing.T) { }) } } + +func TestCompile_HostsAddressMatchesFamilyGuarded(t *testing.T) { + state := mustCompile(t, hostsCfg(func(c *config.Config) { + c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64") + c.Hosts[0].Exclusions = []string{"192.0.2.9"} + })) + want := []string{ + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6 !saddr=2001:db8::/64", + "iif=wlo1 ip4 saddr=192.0.2.9", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", + } + if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) { + t.Errorf("net->fw DROP policy = %q, want %q", got, want) + } + want = []string{ + "iif=wlo1 ip4 saddr=192.0.2.0/24 !saddr=192.0.2.9", "iif=wlo1 ip6 saddr=2001:db8::/64", + "iif=enp2s0 ip4 saddr=198.51.100.0/24", + } + if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) { + t.Errorf("lan->fw policy = %q, want %q", got, want) + } + for chain, rules := range state.Rules { + for _, r := range rules { + var fam byte + for i, e := range r.Exprs { + if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyNFPROTO { + fam = r.Exprs[i+1].(*expr.Cmp).Data[0] + } + p, ok := e.(*expr.Payload) + if !ok || p.Base != expr.PayloadBaseNetworkHeader { + continue + } + if p.Len == 4 && fam != unix.NFPROTO_IPV4 || p.Len == 16 && fam != unix.NFPROTO_IPV6 { + t.Errorf("%s %s: %d-byte address compare without its family guard: %s", chain, r.Tag, p.Len, describeRule(r)) + } + } + } + } +}