diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 472794f..894cf59 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -79,7 +79,7 @@ func (c *Compiler) Compile() (*FirewallState, error) { // 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] + table := tableFamily(family) for chain, rules := range state.Rules { var out []ManagedRule for _, r := range rules { @@ -98,6 +98,11 @@ func familyGuards(state *FirewallState, family config.AddressFamily) error { return nil } +// tableFamily is the NFPROTO an ip/ip6 table is restricted to, 0 for inet. +func tableFamily(family config.AddressFamily) byte { + return map[config.AddressFamily]byte{config.FamilyIP: unix.NFPROTO_IPV4, config.FamilyIP6: unix.NFPROTO_IPV6}[family] +} + // 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 { @@ -503,7 +508,7 @@ func (c *Compiler) compileRules(state *FirewallState) error { return fmt.Errorf("rule[%d]: %w", i, err) } if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, fwZone, 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) + return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists, or with a negated address in an inet table (it expands to one rule per family); each expanded rule would get its own limiter", i) } } @@ -681,19 +686,39 @@ func (c *Compiler) compileDNATAccept(state *FirewallState, tag, srcZone, srcAddr return nil } -// specCount is how many zone/address combinations compileOneRule expands src and dst into. +// specCount is how many rules compileOneRule emits for src and dst once familyGuards has dropped +// cross-family address combinations and those outside an ip/ip6 table's family. func (c *Compiler) specCount(srcSpec, dstSpec, origDest, fwZone string, action config.RuleAction) int { + table := tableFamily(c.cfg.Settings.AddressFamily) + count := func(addrs ...string) int { + n := 0 + var walk func(i int, fams []byte) + walk = func(i int, fams []byte) { + if !famsAgree(fams...) { + return + } + if i == len(addrs) { + n++ + return + } + for _, a := range splitAddrs(addrs[i]) { + walk(i+1, append(fams, addrFamily(a))) + } + } + walk(0, []byte{table}) + return n + } n := 0 if action == config.RuleDNAT || action == config.RuleRedirect { for _, src := range c.dnatSourceSpecs(srcSpec, fwZone) { - n += len(splitAddrs(src.Addr)) + n += count(src.Addr, origDest) } - return n * len(splitAddrs(origDest)) + return n } for _, p := range c.zonePairs(srcSpec, dstSpec, fwZone) { - n += len(splitAddrs(p[0].Addr)) * len(splitAddrs(p[1].Addr)) + n += count(p[0].Addr, p[1].Addr, origDest) } - return n * len(splitAddrs(origDest)) + return n } // zonePairs is the src/dst zone expansion of a non-DNAT rule, with fw added beside all/any. diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index baf804f..3dde993 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3644,3 +3644,48 @@ func TestCompile_HostsRouteBack(t *testing.T) { t.Errorf("routeback hosts lan lan policy = %q, want %q", got, want) } } + +func TestCompile_LimitNegatedAddrPerTableFamily(t *testing.T) { + for _, tt := range []struct { + name string + family config.AddressFamily + source string + wantErr bool + }{ + {"ip v4 negation", config.FamilyIP, "net:!192.0.2.1", false}, + {"ip v6 negation", config.FamilyIP, "net:!2001:db8::1", false}, + {"ip mixed negation", config.FamilyIP, "net:!192.0.2.1,2001:db8::1", false}, + {"ip6 v6 negation", config.FamilyIP6, "net:!2001:db8::1", false}, + {"ip6 v4 negation", config.FamilyIP6, "net:!192.0.2.1", false}, + {"inet v4 negation", config.FamilyINET, "net:!192.0.2.1", true}, + } { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: tt.family}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Rules: []config.Rule{{Action: config.RuleAccept, Source: tt.source, Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, RateLimit: "10/sec"}}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if tt.wantErr { + if err == nil || !strings.Contains(err.Error(), "negated address in an inet table") { + t.Fatalf("Compile() error = %v, want negated-address inet error", err) + } + return + } + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + n := 0 + for _, r := range state.Rules["input"] { + if r.Tag == "rule:0" { + n++ + } + } + if n != 1 { + t.Errorf("got %d rule:0 rules, want 1", n) + } + }) + } +}