diff --git a/internal/config/blrules.go b/internal/config/blrules.go index b1fc8e3..4118d22 100644 --- a/internal/config/blrules.go +++ b/internal/config/blrules.go @@ -57,17 +57,25 @@ func (c *Config) validateBlrules() error { if r.Source != "all" && r.Source != "any" && r.Source != "none" && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - srcZone := zoneFromSpec(r.Source) - if _, ok := c.Zones[srcZone]; !ok { - return fmt.Errorf("blrules[%d]: source zone %q not defined", i, srcZone) + for _, zs := range SplitZoneList(r.Source) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("blrules[%d]: source zone %q not defined", i, zs.Zone) + } + if !validAddrList(zs.Addr) { + return fmt.Errorf("blrules[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - dstZone := zoneFromSpec(r.Dest) - if _, ok := c.Zones[dstZone]; !ok { - return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, dstZone) + for _, zs := range SplitZoneList(r.Dest) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, zs.Zone) + } + if !validAddrList(zs.Addr) { + return fmt.Errorf("blrules[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4041c1b..7485bae 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -3,6 +3,7 @@ package config import ( "os" "path/filepath" + "reflect" "strings" "testing" ) @@ -854,6 +855,32 @@ func TestValidateRules(t *testing.T) { }, wantErr: "source zone \"missing\" not defined", }, + { + name: "comma zone lists are valid", + rules: []Rule{ + {Action: RuleAccept, Source: "fw,loc", Dest: "loc,net:192.0.2.1,198.51.100.1"}, + }, + }, + { + name: "undefined zone in dest list", + rules: []Rule{ + {Action: RuleAccept, Source: "loc", Dest: "net,missing"}, + }, + wantErr: "dest zone \"missing\" not defined", + }, + { + name: "negation prefixing the whole address list is valid", + rules: []Rule{ + {Action: RuleAccept, Source: "net:!192.0.2.1,198.51.100.1", Dest: "fw"}, + }, + }, + { + name: "negation inside an address list", + rules: []Rule{ + {Action: RuleAccept, Source: "net", Dest: "loc:192.0.2.1,!198.51.100.1"}, + }, + wantErr: "'!' may only prefix the whole address list", + }, { name: "all keyword is valid source", rules: []Rule{ @@ -1006,3 +1033,22 @@ func TestValidateSNAT(t *testing.T) { }) } } + +func TestSplitZoneList(t *testing.T) { + tests := []struct { + in string + want []ZoneSpec + }{ + {"net", []ZoneSpec{{Zone: "net"}}}, + {"fw,lan,svr", []ZoneSpec{{Zone: "fw"}, {Zone: "lan"}, {Zone: "svr"}}}, + {"svr:192.0.2.17", []ZoneSpec{{Zone: "svr", Addr: "192.0.2.17"}}}, + {"net:192.0.2.1,198.51.100.1", []ZoneSpec{{Zone: "net", Addr: "192.0.2.1,198.51.100.1"}}}, + {"lan,svr:192.0.2.17", []ZoneSpec{{Zone: "lan"}, {Zone: "svr", Addr: "192.0.2.17"}}}, + {"net:2001:db8::1", []ZoneSpec{{Zone: "net", Addr: "2001:db8::1"}}}, + } + for _, tt := range tests { + if got := SplitZoneList(tt.in); !reflect.DeepEqual(got, tt.want) { + t.Errorf("SplitZoneList(%q) = %+v, want %+v", tt.in, got, tt.want) + } + } +} diff --git a/internal/config/rules.go b/internal/config/rules.go index 0a52040..fb94e31 100644 --- a/internal/config/rules.go +++ b/internal/config/rules.go @@ -1,6 +1,9 @@ package config -import "fmt" +import ( + "fmt" + "strings" +) type RuleAction string @@ -172,10 +175,12 @@ func (c *Config) validateRules() error { if r.Source != "all" && r.Source != "any" && r.Source != "none" && !hasPrefix(r.Source, "all+") && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - for _, srcPart := range splitZones(r.Source) { - srcZone := zoneFromSpec(srcPart) - if _, ok := c.Zones[srcZone]; !ok { - return fmt.Errorf("rule[%d]: source zone %q not defined", i, srcZone) + for _, zs := range SplitZoneList(r.Source) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("rule[%d]: source zone %q not defined", i, zs.Zone) + } + if !validAddrList(zs.Addr) { + return fmt.Errorf("rule[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) } } } @@ -183,10 +188,12 @@ func (c *Config) validateRules() error { if r.Action != RuleDNAT && r.Action != RuleRedirect && r.Action != RuleNoNAT { if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && !hasPrefix(r.Dest, "all+") && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - for _, dstPart := range splitZones(r.Dest) { - dstZone := zoneFromSpec(dstPart) - if _, ok := c.Zones[dstZone]; !ok { - return fmt.Errorf("rule[%d]: dest zone %q not defined", i, dstZone) + for _, zs := range SplitZoneList(r.Dest) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("rule[%d]: dest zone %q not defined", i, zs.Zone) + } + if !validAddrList(zs.Addr) { + return fmt.Errorf("rule[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) } } } @@ -223,6 +230,30 @@ func (c *Config) validateRules() error { return nil } +// ZoneSpec is one zone of a SOURCE/DEST list; Addr is its comma-separated address list, if any. +type ZoneSpec struct{ Zone, Addr string } + +// SplitZoneList parses "lan,svr:a,b": commas before the first colon separate zones, +// commas after it separate addresses of the last zone (shorewall semantics). +func SplitZoneList(spec string) []ZoneSpec { + zones, addr, _ := strings.Cut(spec, ":") + var out []ZoneSpec + for _, z := range strings.Split(zones, ",") { + if z = strings.TrimSpace(z); z != "" { + out = append(out, ZoneSpec{Zone: z}) + } + } + if len(out) > 0 { + out[len(out)-1].Addr = addr + } + return out +} + +// validAddrList reports whether '!' appears only at the start, negating the whole list. +func validAddrList(addr string) bool { + return !strings.Contains(strings.TrimPrefix(addr, "!"), "!") +} + // zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24". func zoneFromSpec(spec string) string { for i, c := range spec { diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 2994a80..b054955 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -3,6 +3,7 @@ package nftables import ( "encoding/binary" "fmt" + "log/slog" "net" "strconv" "strings" @@ -14,7 +15,8 @@ import ( ) type Compiler struct { - cfg *config.Config + cfg *config.Config + warned map[string]bool } func NewCompiler(cfg *config.Config) *Compiler { @@ -25,6 +27,7 @@ func (c *Compiler) Compile() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } + c.warned = nil c.compileLoopback(state) if err := c.compileConntrackFastPath(state); err != nil { @@ -272,8 +275,8 @@ func (c *Compiler) compileRules(state *FirewallState) error { if err != nil { return fmt.Errorf("rule[%d]: %w", i, err) } - if len(matches) > 1 { - return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto or port lists (each expanded rule would get its own limiter)", i) + 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) } } @@ -295,11 +298,14 @@ func (c *Compiler) compileRules(state *FirewallState) error { } func (c *Compiler) applyRuleExtras(state *FirewallState, tag string, rule config.Rule) { - fwZone := c.cfg.FirewallZone() - srcZone, _ := splitZoneSpec(rule.Source) - dstZone, _ := splitZoneSpec(rule.Dest) - chain := c.selectChain(srcZone, dstZone, fwZone) + for chain := range state.Rules { + if chain != "prerouting" { + c.applyChainExtras(state, chain, tag, rule) + } + } +} +func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rule config.Rule) { rules := state.Rules[chain] for idx := len(rules) - 1; idx >= 0; idx-- { if rules[idx].Tag != tag { @@ -355,13 +361,61 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p dports, sports config.PortSpec, action config.RuleAction, logLevel string, dnatDest string, fwZone string, section config.RuleSection) error { - srcZone, srcAddr := splitZoneSpec(srcSpec) - dstZone, dstAddr := splitZoneSpec(dstSpec) - - if action == config.RuleDNAT || action == config.RuleRedirect { - return c.compileDNATRule(state, tag, srcSpec, dstSpec, proto, dports, sports, action, logLevel, fwZone) + for _, src := range zoneSpecs(srcSpec) { + for _, srcAddr := range splitAddrs(src.Addr) { + if action == config.RuleDNAT || action == config.RuleRedirect { + if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, dstSpec, proto, dports, action, logLevel); err != nil { + return err + } + continue + } + for _, dst := range zoneSpecs(dstSpec) { + for _, dstAddr := range splitAddrs(dst.Addr) { + if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, proto, + dports, sports, action, logLevel, fwZone, section); err != nil { + return err + } + } + } + } } + 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 { + zone, addr := splitZoneSpec(spec) + if base, _, _ := strings.Cut(strings.TrimSuffix(zone, "+"), "!"); base == "all" || base == "any" { + return []config.ZoneSpec{{Zone: zone, Addr: addr}} + } + return config.SplitZoneList(spec) +} + +// splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. +func splitAddrs(addr string) []string { + if addr == "" || strings.HasPrefix(addr, "!") { + return []string{addr} + } + return strings.Split(addr, ",") +} + +func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, proto string, + dports, sports config.PortSpec, action config.RuleAction, logLevel string, + fwZone string, section config.RuleSection) error { srcIfaces := c.resolveZoneInterfaces(srcZone) dstIfaces := c.resolveZoneInterfaces(dstZone) chain := c.selectChain(srcZone, dstZone, fwZone) @@ -400,10 +454,8 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p return nil } -func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, proto string, - dports, sports config.PortSpec, action config.RuleAction, logLevel, fwZone string) error { - - srcZone, srcAddr := splitZoneSpec(srcSpec) +func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dstSpec, proto string, + dports config.PortSpec, action config.RuleAction, logLevel string) error { chain := "prerouting" parts := strings.SplitN(dstSpec, ":", 3) @@ -884,10 +936,29 @@ func (c *Compiler) resolveZoneInterfaces(zone string) []string { return []string{""} } ifaces := c.cfg.ZoneInterfaces(zone) - if len(ifaces) == 0 { - return []string{""} + if len(ifaces) > 0 { + return ifaces } - return ifaces + if z, ok := c.cfg.Zones[zone]; ok && z.Type == config.ZoneIP && !c.zoneHasHosts(zone) { + if !c.warned[zone] { + if c.warned == nil { + c.warned = map[string]bool{} + } + c.warned[zone] = true + slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) + } + return nil + } + return []string{""} +} + +func (c *Compiler) zoneHasHosts(zone string) bool { + for _, h := range c.cfg.Hosts { + if h.Zone == zone { + return true + } + } + return false } func (c *Compiler) expandZoneRef(ref string) []string { @@ -1318,6 +1389,17 @@ func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) { if strings.HasPrefix(cidr, "!") { negated = true cidr = cidr[1:] + if strings.Contains(cidr, ",") { + var all []expr.Any + for _, a := range strings.Split(cidr, ",") { + e, err := matchAddrCIDR("!"+a, isSrc) + if err != nil { + return nil, err + } + all = append(all, e...) + } + return all, nil + } } cmpOp := expr.CmpOpEq diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 1a221b5..37c3f7a 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1648,6 +1648,144 @@ func TestCompile_PortAndProtoLists(t *testing.T) { } } +// describeRule renders a rule's iif/oif/saddr/daddr matches, e.g. "iif=eth1 oif=eth2 daddr=192.0.2.1". +func describeRule(r ManagedRule) string { + var parts []string + for i, e := range r.Exprs { + cmp, ok := func() (*expr.Cmp, bool) { + if i+1 >= len(r.Exprs) { + return nil, false + } + c, ok := r.Exprs[i+1].(*expr.Cmp) + return c, ok + }() + if !ok { + continue + } + switch m := e.(type) { + case *expr.Meta: + switch m.Key { + case expr.MetaKeyIIFNAME: + parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00")) + case expr.MetaKeyOIFNAME: + parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) + } + case *expr.Payload: + if m.Base == expr.PayloadBaseNetworkHeader && m.Len == 4 { + name := map[uint32]string{12: "saddr", 16: "daddr"}[m.Offset] + if cmp.Op == expr.CmpOpNeq { + name = "!" + name + } + parts = append(parts, fmt.Sprintf("%s=%d.%d.%d.%d", name, cmp.Data[0], cmp.Data[1], cmp.Data[2], cmp.Data[3])) + } + } + } + return strings.Join(parts, " ") +} + +func TestCompile_CommaZoneLists(t *testing.T) { + tests := []struct { + name string + rule config.Rule + blrule *config.BlruleRule + want map[string][]string + }{ + { + name: "fw in source list goes to output", + rule: config.Rule{Action: config.RuleAccept, Source: "fw,lan", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"22"}}, + want: map[string][]string{"output": {""}, "forward": {"iif=eth1 oif=eth2"}}, + }, + { + name: "dest list with fw splits input and forward", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "fw,svr,net"}, + want: map[string][]string{"input": {"iif=eth1"}, "forward": {"iif=eth1 oif=eth2", "iif=eth1 oif=eth0"}}, + }, + { + name: "zone without interfaces emits nothing", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "svr,dmz"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth2"}}, + }, + { + name: "address list after colon belongs to one zone", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth0 daddr=192.0.2.1", "iif=eth1 oif=eth0 daddr=198.51.100.1"}}, + }, + { + name: "zone:address inside a list", + rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"}, + want: map[string][]string{"input": {"iif=eth1", "iif=eth2 saddr=203.0.113.7"}}, + }, + { + name: "dnat source list", + rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, + want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}}, + }, + { + name: "dnat source address list", + rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, + want: map[string][]string{"prerouting": {"iif=eth0 saddr=192.0.2.5", "iif=eth0 saddr=198.51.100.5"}}, + }, + { + name: "negated address list stays one AND-ed rule", + 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 !saddr=192.0.2.5 !saddr=198.51.100.5"}}, + }, + { + name: "zone named like all/any keyword is a plain zone", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "anycast,net"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth3", "iif=eth1 oif=eth0"}}, + }, + { + name: "interface-less ipsec zone keeps zone-agnostic rule, ip zone skipped", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"}, + want: map[string][]string{"forward": {"iif=eth1"}}, + }, + { + name: "blrule zone list", + blrule: &config.BlruleRule{Action: config.BlruleDrop, Source: "net,anycast", Dest: "fw"}, + want: map[string][]string{"input": {"iif=eth0", "iif=eth3"}}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, 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}, "svr": {Type: config.ZoneIP}, "dmz": {Type: config.ZoneIP}, + "anycast": {Type: config.ZoneIP}, "vpn": {Type: config.ZoneIPSec}, + }, + Interfaces: []config.Interface{ + {Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}, {Zone: "svr", Interface: "eth2"}, + {Zone: "anycast", Interface: "eth3"}, + }, + Rules: []config.Rule{tt.rule}, + PortGroups: make(map[string]config.PortGroup), + } + tag := "rule:0" + if tt.blrule != nil { + cfg.Rules, cfg.Blrules, tag = nil, []config.BlruleRule{*tt.blrule}, "blrule:0" + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + got := map[string][]string{} + for chain, rules := range state.Rules { + for _, r := range rules { + if r.Tag == tag { + got[chain] = append(got[chain], describeRule(r)) + } + } + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %v, want %v", got, tt.want) + } + }) + } +} + func listCfg(mod func(*config.Config)) *config.Config { cfg := &config.Config{ Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, @@ -1704,6 +1842,42 @@ func TestCompile_ListExpansionCounts(t *testing.T) { } } +func TestCompile_DNATGetsNoRuleExtras(t *testing.T) { + compile := func(mark string) []expr.Any { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}, Mark: mark}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + return state.Rules["prerouting"][len(state.Rules["prerouting"])-1].Exprs + } + if plain, marked := compile(""), compile("0x1"); !reflect.DeepEqual(plain, marked) { + t.Errorf("DNAT prerouting rule changed by mark extra: %d exprs vs %d", len(plain), len(marked)) + } +} + +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) { tests := []struct { proto string @@ -1738,6 +1912,21 @@ func TestCompile_RejectPerProto(t *testing.T) { } } +func TestNegatedAddressList(t *testing.T) { + exprs, err := matchDestCIDR("!192.0.2.1,198.51.100.1") + if err != nil { + t.Fatalf("matchDestCIDR error: %v", err) + } + if len(exprs) != 4 { + t.Fatalf("expected 4 exprs, got %d", len(exprs)) + } + for _, i := range []int{1, 3} { + if exprs[i].(*expr.Cmp).Op != expr.CmpOpNeq { + t.Errorf("expr %d should be CmpOpNeq", i) + } + } +} + func TestCompile_ListErrors(t *testing.T) { tests := []struct { name string