From b6d67897def1196193a15228bec0f1f50ec58b92 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 23:48:08 +1000 Subject: [PATCH] Expand all!zone exclusions and fail closed on unknown zones --- internal/nftables/compiler.go | 59 +++++++++++++++++++++++------- internal/nftables/compiler_test.go | 42 +++++++++++++++++++++ 2 files changed, 87 insertions(+), 14 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index c8170e1..d8d6b12 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -216,7 +216,7 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for i, ct := range c.cfg.Conntrack { tag := fmt.Sprintf("conntrack:%d", i) - srcs, dsts := zoneSpecs(ct.Source), zoneSpecs(ct.Dest) + srcs, dsts := c.zoneSpecs(ct.Source), c.zoneSpecs(ct.Dest) if len(srcs) == 0 { srcs = []config.ZoneSpec{{}} } @@ -226,6 +226,10 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { for _, src := range srcs { // raw_output only sees locally generated traffic, so it applies to an fw (or omitted) source only. + if isZoneExclusion(ct.Source) && (ct.Chain == config.ConntrackPrerouting && src.Zone == fwZone || + ct.Chain == config.ConntrackOutput && src.Zone != fwZone) { + continue + } chains := []string{"raw_prerouting"} switch { case ct.Chain == config.ConntrackOutput && src.Zone != fwZone && src.Zone != "": @@ -315,7 +319,7 @@ func (c *Compiler) compileRules(state *FirewallState) error { if err != nil { return fmt.Errorf("rule[%d]: %w", i, err) } - if len(matches)*specCount(rule.Source, rule.Dest, rule.OrigDest, rule.Action) > 1 { + if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, rule.Action) > 1 { return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i) } } @@ -401,7 +405,7 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p dports, sports config.PortSpec, action config.RuleAction, logLevel string, dnatDest, origDest string, fwZone string, section config.RuleSection) error { - for _, src := range zoneSpecs(srcSpec) { + for _, src := range c.zoneSpecs(srcSpec) { for _, srcAddr := range splitAddrs(src.Addr) { for _, od := range splitAddrs(origDest) { if action == config.RuleDNAT || action == config.RuleRedirect { @@ -410,7 +414,12 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p } continue } - for _, dst := range zoneSpecs(dstSpec) { + for _, dst := range c.zoneSpecs(dstSpec) { + // Exclusion expansion never pairs fw with itself, and pairs a zone with itself only for "all+". + if src.Zone == dst.Zone && (isZoneExclusion(srcSpec) || isZoneExclusion(dstSpec)) && + (src.Zone == fwZone || !strings.Contains(srcSpec, "+!")) { + continue + } for _, dstAddr := range splitAddrs(dst.Addr) { if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, od, proto, dports, sports, action, logLevel, fwZone, section); err != nil { @@ -425,9 +434,9 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p } // specCount is how many zone/address combinations compileOneRule expands src and dst into. -func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int { +func (c *Compiler) specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int { count := func(spec string) (n int) { - for _, z := range zoneSpecs(spec) { + for _, z := range c.zoneSpecs(spec) { n += len(splitAddrs(z.Addr)) } return n @@ -439,13 +448,26 @@ func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int return n * count(dstSpec) } -// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. -func zoneSpecs(spec string) []config.ZoneSpec { +// zoneSpecs expands a comma zone list; "all"/"any" stay global and "all!x,y" becomes every zone but x and y. +func (c *Compiler) zoneSpecs(spec string) []config.ZoneSpec { zone, addr := splitZoneSpec(spec) - if base, _, _ := strings.Cut(strings.TrimSuffix(zone, "+"), "!"); base == "all" || base == "any" { - return []config.ZoneSpec{{Zone: zone, Addr: addr}} + if !isZoneExclusion(zone) { + if base := strings.TrimSuffix(zone, "+"); base == "all" || base == "any" { + return []config.ZoneSpec{{Zone: zone, Addr: addr}} + } + return config.SplitZoneList(spec) } - return config.SplitZoneList(spec) + var out []config.ZoneSpec + for _, z := range c.expandZoneRef(zone) { + out = append(out, config.ZoneSpec{Zone: z, Addr: addr}) + } + return out +} + +func isZoneExclusion(spec string) bool { + base, _, ok := strings.Cut(spec, "!") + base = strings.TrimSuffix(base, "+") + return ok && (base == "all" || base == "any") } // splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. @@ -996,9 +1018,18 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { return "forward" } -// resolveZoneInterfaces returns nil (fail closed) for a zone with no interfaces unless a non-negated address match narrows the rule. +// resolveZoneInterfaces returns nil (fail closed) for an unknown zone, or one with no interfaces unless a non-negated address match narrows the rule. func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string { - if z, ok := c.cfg.Zones[zone]; !ok || z.Type == config.ZoneFirewall { + switch zone { + case "", "all", "all+", "any", "any+": + return []string{""} + } + z, ok := c.cfg.Zones[zone] + if !ok { + slog.Warn("compiler: unknown zone, skipping its rules", "zone", zone) + return nil + } + if z.Type == config.ZoneFirewall { return []string{""} } if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 { @@ -1032,7 +1063,7 @@ func (c *Compiler) expandZoneRef(ref string) []string { } } - if base == "all" || base == "all+" { + if base == "all" || base == "all+" || base == "any" || base == "any+" { var zones []string for name := range c.cfg.Zones { if excluded != nil && excluded[name] { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 6ba6086..85e3f7f 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -2497,6 +2497,26 @@ func TestCompile_ConntrackZones(t *testing.T) { ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan"}, wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, }, + { + name: "unknown zone fails closed", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "nte", Dest: "fw"}, + want: map[string][]string{}, + }, + { + name: "all!net expands to every other zone", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all!net", Dest: "fw:192.0.2.53"}, + want: map[string][]string{"raw_output": {"daddr=192.0.2.53"}, "raw_prerouting": {"iif=eth1 daddr=192.0.2.53"}}, + }, + { + name: "all!net in prerouting skips fw", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all!net", Dest: "fw", Chain: config.ConntrackPrerouting}, + want: map[string][]string{"raw_prerouting": {"iif=eth1"}}, + }, + { + name: "omitted source and dest is global", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53"}}, + want: map[string][]string{"raw_prerouting": {""}}, + }, { name: "dest zone with address matches daddr", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"}, @@ -2530,3 +2550,25 @@ func TestCompile_ConntrackZones(t *testing.T) { }) } } + +func TestCompile_RuleZoneExclusionExpands(t *testing.T) { + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Rules = []config.Rule{ + {Source: "all!net", Dest: "fw", Action: config.RuleAccept, Proto: "tcp", DPort: config.PortSpec{"22"}}, + {Source: "nte", Dest: "fw", Action: config.RuleAccept}, + } + }) + state := mustCompile(t, cfg) + var got []string + for _, r := range taggedRules(state, "input", "rule:0") { + got = append(got, describeRule(r)) + } + if want := []string{"iif=eth1"}; !reflect.DeepEqual(got, want) { + t.Errorf("all!net -> fw input rules = %v, want %v", got, want) + } + if r := taggedRules(state, "input", "rule:1"); len(r) != 0 { + t.Errorf("unknown zone compiled %d rules, want 0", len(r)) + } +}