diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index e5a9737..f8ef196 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -268,6 +268,16 @@ func (c *Compiler) compileRules(state *FirewallState) error { sport = rule.SPort } + if rule.RateLimit != "" || rule.ConnLimit != "" { + matches, err := l4Matches(proto, dports, sport) + if err != nil { + return fmt.Errorf("rule[%d]: %w", i, err) + } + if len(matches)*specCount(rule.Source, rule.Dest, 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) + } + } + if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, proto, dports, sport, rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil { @@ -368,6 +378,20 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p return nil } +// specCount is how many zone/address combinations compileOneRule expands src and dst into. +func specCount(srcSpec, dstSpec string, action config.RuleAction) int { + count := func(spec string) (n int) { + for _, z := range zoneSpecs(spec) { + n += len(splitAddrs(z.Addr)) + } + return n + } + if action == config.RuleDNAT || action == config.RuleRedirect { + return count(srcSpec) + } + return count(srcSpec) * count(dstSpec) +} + // zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. func zoneSpecs(spec string) []config.ZoneSpec { if strings.HasPrefix(spec, "all") || strings.HasPrefix(spec, "any") { @@ -1006,10 +1030,13 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) var out []l4Match for _, p := range protos { p = strings.TrimSpace(p) + if proto != "" && p == "" { + return nil, fmt.Errorf("empty element in proto list %q", proto) + } isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") parseD := parsePortOrRange if isICMP { - parseD = func(s string) ([]expr.Any, error) { return matchICMPType(s), nil } + parseD = matchICMPType } dalts, err := portAlternatives(dports, parseD) if err != nil { @@ -1038,7 +1065,7 @@ func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, err for _, item := range ports { for _, s := range strings.Split(item, ",") { if s = strings.TrimSpace(s); s == "" { - continue + return nil, fmt.Errorf("empty element in port list %q", item) } e, err := parse(s) if err != nil { @@ -1216,33 +1243,33 @@ var icmpTypeNames = map[string]byte{ "address-mask-reply": 18, } -func matchICMPType(spec string) []expr.Any { +func matchICMPType(spec string) ([]expr.Any, error) { if strings.Contains(spec, "/") { parts := strings.SplitN(spec, "/", 2) typeVal, ok := resolveICMPType(parts[0]) if !ok { - return nil + return nil, fmt.Errorf("invalid icmp type %q", parts[0]) } code, err := strconv.ParseUint(parts[1], 10, 8) if err != nil { - return nil + return nil, fmt.Errorf("invalid icmp code %q: %w", parts[1], err) } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}}, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 1, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{byte(code)}}, - } + }, nil } typeVal, ok := resolveICMPType(spec) if !ok { - return nil + return nil, fmt.Errorf("invalid icmp type %q", spec) } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}}, - } + }, nil } func resolveICMPType(s string) (byte, bool) { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 6ae3a27..e9cfe1b 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/google/nftables/expr" + "golang.org/x/sys/unix" "git.unkin.net/unkin/tomswall/internal/config" ) @@ -1269,7 +1270,10 @@ func TestMatchICMPType(t *testing.T) { } for _, tt := range tests { - exprs := matchICMPType(tt.input) + exprs, err := matchICMPType(tt.input) + if err != nil { + t.Fatalf("matchICMPType(%q) error: %v", tt.input, err) + } if len(exprs) != tt.wantLen { t.Errorf("matchICMPType(%q) returned %d exprs, want %d", tt.input, len(exprs), tt.wantLen) } @@ -1729,37 +1733,102 @@ func TestCompile_CommaZoneLists(t *testing.T) { } } -func TestCompile_CommaZoneListExtras(t *testing.T) { +func listCfg(mod func(*config.Config)) *config.Config { cfg := &config.Config{ Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, - Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, "lan": {Type: config.ZoneIP}}, - Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}}, - Rules: []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"}}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, PortGroups: make(map[string]config.PortGroup), } - state, err := NewCompiler(cfg).Compile() + mod(cfg) + return cfg +} + +func taggedRules(state *FirewallState, chain, tag string) []ManagedRule { + var out []ManagedRule + for _, r := range state.Rules[chain] { + if r.Tag == tag { + out = append(out, r) + } + } + return out +} + +func TestCompile_ListExpansionCounts(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + want int + }{ + {"proto x dport cross product", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80", "443"}}} + }, "input", "rule:0", 4}, + {"sport list", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", SPort: config.PortSpec{"1024,2048"}}} + }, "input", "rule:0", 2}, + {"snat proto x dport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} + }, "postrouting", "snat:0", 4}, + {"conntrack dport list", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53", "123"}}} + }, "prerouting", "conntrack:0:prerouting", 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state, err := NewCompiler(listCfg(tt.mod)).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + if got := len(taggedRules(state, tt.chain, tt.tag)); got != tt.want { + t.Errorf("%s rules in %s = %d, want %d", tt.tag, tt.chain, got, tt.want) + } + }) + } +} + +func TestCompile_CommaZoneListLimitErrors(t *testing.T) { + for _, r := range []config.Rule{ + {Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"}, + {Action: config.RuleAccept, Source: "net,lan", Dest: "fw", ConnLimit: "10"}, + {Action: config.RuleAccept, Source: "net", Dest: "fw:192.0.2.1,198.51.100.1", RateLimit: "10/sec"}, + } { + t.Run(r.Source+">"+r.Dest, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, "lan": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}}, + Rules: []config.Rule{r}, + PortGroups: make(map[string]config.PortGroup), + } + if _, err := NewCompiler(cfg).Compile(); err == nil { + t.Fatal("Compile() succeeded, want error") + } + }) + } +} + +func TestCompile_RejectPerProto(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}} + })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } - for _, chain := range []string{"input", "forward"} { - found := false - for _, r := range state.Rules[chain] { - if r.Tag != "rule:0" { - continue - } - found = true - hasLimit := false - for _, e := range r.Exprs { - if _, ok := e.(*expr.Limit); ok { - hasLimit = true - } - } - if !hasLimit { - t.Errorf("%s rule missing Limit expression", chain) - } + rules := taggedRules(state, "input", "rule:0") + if len(rules) != 2 { + t.Fatalf("got %d rules, want 2", len(rules)) + } + want := []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH} + for i, r := range rules { + rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) + if !ok { + t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) } - if !found { - t.Errorf("rule:0 not found in %s chain", chain) + if rej.Type != want[i] { + t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, want[i]) } } } @@ -1778,3 +1847,44 @@ func TestNegatedAddressList(t *testing.T) { } } } + +func TestCompile_ListErrors(t *testing.T) { + tests := []struct { + name string + rule config.Rule + }{ + {"invalid port in list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,abc"}}}, + {"trailing empty proto", config.Rule{Proto: "tcp,", DPort: config.PortSpec{"80"}}}, + {"leading empty proto", config.Rule{Proto: ",udp", DPort: config.PortSpec{"80"}}}, + {"empty port element", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,"}}}, + {"unknown icmp type", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus"}}}, + {"ratelimit with port list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,443"}, RateLimit: "10/sec"}}, + {"connlimit with proto list", config.Rule{Proto: "tcp,udp", DPort: config.PortSpec{"53"}, ConnLimit: "10"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := tt.rule + r.Action, r.Source, r.Dest = config.RuleAccept, "net", "fw" + _, err := NewCompiler(listCfg(func(c *config.Config) { c.Rules = []config.Rule{r} })).Compile() + if err == nil { + t.Fatal("Compile() succeeded, want error") + } + }) + } +} + +func TestCompile_ListExpansionDiffStable(t *testing.T) { + cfg := listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} + }) + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + if n := len(taggedRules(state, "input", "rule:0")); n != 4 { + t.Fatalf("got %d rule:0 rules, want 4", n) + } + if cs := computeDiff(state, state); !cs.Empty() { + t.Errorf("diff not empty:\n%s", cs.Summary()) + } +}