diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 0af760c..472794f 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -67,47 +67,79 @@ func (c *Compiler) Compile() (*FirewallState, error) { } c.compileMSSClamp(state) limitLogs(state, c.cfg.Settings.LogLimit) - familyGuards(state, c.cfg.Settings.AddressFamily) + if err := familyGuards(state, c.cfg.Settings.AddressFamily); err != nil { + return nil, err + } 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) { +// payload as nft emits it; nft list cannot decode repeated or conflicting guards. Rules are built per +// family, so a conflict is a compiler bug and fails the compile. An ip/ip6 table is its own guard: +// guards go and other-family rules are dropped. +func familyGuards(state *FirewallState, family config.AddressFamily) error { 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) + fam, ok := guardFamily(r.Exprs) + if !ok { + return fmt.Errorf("%s %s: conflicting address families", chain, r.Tag) } + if table != 0 && fam != 0 && fam != table { + continue + } + r.Exprs = normalizeGuards(r.Exprs, fam, table != 0) + out = append(out, r) } state.Rules[chain] = out } + return nil } -func normalizeGuards(in []expr.Any, table byte) ([]expr.Any, bool) { - var fam byte +// guardFamily is the family exprs' nfproto guards require (0: none); ok is false when they conflict. +func guardFamily(exprs []expr.Any) (fam byte, ok bool) { + for i := range exprs { + if p := nfprotoGuard(exprs, i); p != 0 { + if fam != 0 && p != fam { + return 0, false + } + fam = p + } + } + return fam, true +} + +// famsAgree reports whether families (0: any) can all hold for one packet. +func famsAgree(fams ...byte) bool { + var f byte + for _, p := range fams { + if p != 0 && f != 0 && p != f { + return false + } + if p != 0 { + f = p + } + } + return true +} + +func normalizeGuards(in []expr.Any, fam byte, strip bool) []expr.Any { 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) + if nfprotoGuard(in, i) != 0 { + if at < 0 { + at = len(out) } i++ continue } out = append(out, in[i]) } - if fam == 0 || table != 0 { - return out, true + if fam == 0 || strip { + return out } if l3 := slices.IndexFunc(out, func(e expr.Any) bool { p, ok := e.(*expr.Payload) @@ -115,7 +147,7 @@ func normalizeGuards(in []expr.Any, table byte) ([]expr.Any, bool) { }); l3 >= 0 && l3 < at { at = l3 } - return slices.Insert(out, at, matchNFProto(fam)...), true + return slices.Insert(out, at, matchNFProto(fam)...) } // nfprotoGuard is the family a meta nfproto == match at in[i] guards for, else 0. @@ -403,7 +435,7 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule, srcZone, srcAddr, dstZone, dstAddr string) error { if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" && - (dstAddr == "" || strings.HasPrefix(dstAddr, "!")) { + !narrows(dstAddr) { return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone) } srcIfaces, dstIfaces := c.resolveZone(srcZone, srcAddr), []zoneMatch{{}} @@ -425,6 +457,9 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, } for _, m := range matches { exprs := m.exprs + if _, ok := guardFamily(exprs); !ok { + continue + } switch ct.Action { case config.ConntrackNoTrack: exprs = append(exprs, &expr.Notrack{}) @@ -724,12 +759,33 @@ func isZoneExclusion(spec string) bool { return ok && (base == "all" || base == "any") } -// splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. +// splitAddrs yields one alternative per listed address. A negated list ("everything except") becomes +// one AND-ed match per family: the family's negations, or the bare family (a /0) when it has none. func splitAddrs(addr string) []string { - if addr == "" || strings.HasPrefix(addr, "!") { - return []string{addr} + if addr == "" { + return []string{""} } - return strings.Split(addr, ",") + if !strings.HasPrefix(addr, "!") { + return strings.Split(addr, ",") + } + v4, v6 := splitFamily(strings.Split(addr[1:], ",")) + out := []string{"0.0.0.0/0", "::/0"} + if len(v4) > 0 { + out[0] = "!" + strings.Join(v4, ",") + } + if len(v6) > 0 { + out[1] = "!" + strings.Join(v6, ",") + } + return out +} + +// narrows reports whether a splitAddrs alternative restricts addresses within its family. +func narrows(addr string) bool { + if addr == "" || strings.HasPrefix(addr, "!") { + return false + } + p, err := parsePrefix(addr) + return err != nil || p.Bits() > 0 } func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, origDest, proto string, @@ -750,7 +806,7 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, return err } if origDest != "" { - od, err := matchOrigDest(origDest) + od, err := matchDestCIDR(origDest) if err != nil { return fmt.Errorf("origdest: %w", err) } @@ -761,6 +817,9 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, for _, m := range matches { exprs := m.exprs + if _, ok := guardFamily(exprs); !ok { + continue + } if section != "" && section != config.SectionAll { exprs = append(exprs, matchSection(section)...) @@ -810,7 +869,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, var odExprs []expr.Any if origDest != "" { var err error - if odExprs, err = matchOrigDest(origDest); err != nil { + if odExprs, err = matchDestCIDR(origDest); err != nil { return fmt.Errorf("origdest: %w", err) } } @@ -824,6 +883,10 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, if ip == nil { return fmt.Errorf("invalid DNAT address %q", dnatAddr) } + natFam := byte(0) + if action != config.RuleRedirect { + natFam = addrFamily(dnatAddr) + } for _, srcIface := range srcIfaces { zm, err := zoneMatchExprs(srcIface, true) @@ -841,6 +904,9 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, exprs = append(exprs, src...) } exprs = append(exprs, odExprs...) + if f, ok := guardFamily(exprs); !ok || !famsAgree(f, natFam) { + continue + } exprs = append(exprs, m.exprs...) @@ -941,7 +1007,7 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { for _, si := range srcIfaces { for _, di := range dstIfaces { - if sz == dz && intraZoneSkip(si, di) { + if sz == dz && intraZoneSkip(si, di) || chain != "input" && !famsAgree(matchFamily(si), matchFamily(di)) { continue } exprs, err := zonePairExprs(si, di, chain) @@ -1038,25 +1104,27 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { for i, snat := range c.cfg.SNAT { tag := fmt.Sprintf("snat:%d", i) - var exprs []expr.Any - destIface, _ := splitZoneSpec(snat.Dest) - exprs = append(exprs, matchIfaceName(false, destIface)...) - - if snat.Source != "" { - srcExprs, err := matchSourceCIDR(snat.Source) - if err != nil { - return fmt.Errorf("snat[%d]: %w", i, err) + var heads [][]expr.Any + for _, src := range splitAddrs(snat.Source) { + head := matchIfaceName(false, destIface) + if src != "" { + srcExprs, err := matchSourceCIDR(src) + if err != nil { + return fmt.Errorf("snat[%d]: %w", i, err) + } + head = append(head, srcExprs...) + } + if snat.Action != config.SNATAddress || famsAgree(addrFamily(src), addrFamily(snat.Address)) { + heads = append(heads, head) } - exprs = append(exprs, srcExprs...) } matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort) if err != nil { return fmt.Errorf("snat[%d]: %w", i, err) } - head := exprs - exprs = nil + var exprs []expr.Any if snat.Mark != "" { exprs = append(exprs, matchMark(snat.Mark)...) @@ -1102,12 +1170,14 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { } } - for _, m := range matches { - state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ - Chain: "postrouting", - Exprs: append(append(append([]expr.Any{}, head...), m.exprs...), exprs...), - Tag: tag, - }) + for _, head := range heads { + for _, m := range matches { + state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ + Chain: "postrouting", + Exprs: slices.Concat(head, m.exprs, exprs), + Tag: tag, + }) + } } } @@ -1398,7 +1468,7 @@ func (c *Compiler) resolveZone(zone, addr string) []zoneMatch { if len(out) > 0 || hasHosts { return out } - if addr != "" && !strings.HasPrefix(addr, "!") { + if narrows(addr) { return []zoneMatch{{}} } if !c.warned[zone] { @@ -1925,42 +1995,28 @@ func matchDestCIDR(cidr string) ([]expr.Any, error) { return matchGuardedCIDR(cidr, false) } -// matchGuardedCIDR guards an address match with its family's nfproto, unless the list mixes families. +// matchGuardedCIDR guards a single-family splitAddrs alternative with its nfproto; a /0 is the guard alone. func matchGuardedCIDR(cidr string, isSrc bool) ([]expr.Any, error) { + if !narrows(cidr) && !strings.HasPrefix(cidr, "!") { + return matchNFProto(addrFamily(cidr)), nil + } 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 + return append(matchNFProto(addrFamily(cidr)), e...), nil } -// addrFamily is the NFPROTO shared by every address in a (negated) comma list, or 0 if they mix. +// addrFamily is the NFPROTO of a single-family splitAddrs alternative, 0 for none; unparsable is IPv6. func addrFamily(list string) byte { - var proto byte - 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 0 - } - proto = p + if list == "" { + return 0 } - 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 && addrFamily(addr) == 0 { - err = fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr) + a, _, _ := strings.Cut(strings.TrimPrefix(list, "!"), ",") + if p, err := parsePrefix(a); err == nil && p.Addr().Is4() { + return unix.NFPROTO_IPV4 } - return dst, err + return unix.NFPROTO_IPV6 } func matchNFProto(proto byte) []expr.Any { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 1bed22b..baf804f 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -2032,9 +2032,9 @@ func TestCompile_CommaZoneLists(t *testing.T) { "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", + name: "negated address list is one AND-ed rule per family", 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 ip4 !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", "iif=eth0 ip6"}}, }, { name: "zone named like all/any keyword is a plain zone", @@ -2605,25 +2605,6 @@ func TestCompile_ColonRanges(t *testing.T) { } } -func TestMatchOrigDest_FamilyGuard(t *testing.T) { - for addr, want := range map[string]byte{ - "203.0.113.5": unix.NFPROTO_IPV4, - "!203.0.113.0/24,192.0.2.1": unix.NFPROTO_IPV4, - "2001:db8::5": unix.NFPROTO_IPV6, - } { - e, err := matchOrigDest(addr) - if err != nil { - t.Fatalf("%s: %v", addr, err) - } - if m, ok := e[0].(*expr.Meta); !ok || m.Key != expr.MetaKeyNFPROTO || e[1].(*expr.Cmp).Data[0] != want { - t.Errorf("%s: missing nfproto %d guard: %v", addr, want, e[:2]) - } - } - if _, err := matchOrigDest("!203.0.113.5,2001:db8::5"); err == nil { - t.Error("mixed IPv4/IPv6 origdest: want error") - } -} - func TestCompile_OrigDestForwardRejected(t *testing.T) { cfg := &config.Config{ Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, @@ -2781,9 +2762,9 @@ func TestCompile_ConntrackZones(t *testing.T) { want: map[string][]string{"raw_output": {"oif=eth1"}}, }, { - name: "negated addresses stay one AND-ed match", + name: "negated addresses are one AND-ed match per family", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"}, - want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 !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", "iif=eth0 ip6"}}, }, { name: "sport",