Expand address matches per family, fail compile on conflicting guards

This commit is contained in:
2026-10-09 23:48:53 +11:00
parent 42a4dab6a3
commit 9a65973d7f
2 changed files with 130 additions and 93 deletions
+126 -70
View File
@@ -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 {
+4 -23
View File
@@ -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",