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 23137a2..cd7964c 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -3,6 +3,7 @@ package nftables import ( "encoding/binary" "fmt" + "log/slog" "net" "sort" "strconv" @@ -15,7 +16,8 @@ import ( ) type Compiler struct { - cfg *config.Config + cfg *config.Config + warned map[string]bool } func NewCompiler(cfg *config.Config) *Compiler { @@ -26,6 +28,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 { @@ -221,34 +224,30 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { chains = []string{"prerouting", "output"} } + matches, err := l4Matches(ct.Proto, ct.DPort, nil) + if err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + for _, chain := range chains { - var exprs []expr.Any + for _, m := range matches { + exprs := append([]expr.Any{}, m.exprs...) - if ct.Proto != "" { - exprs = append(exprs, matchProto(ct.Proto)...) - } - for _, p := range ct.DPort { - pe, err := parsePortOrRange(p) - if err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) + switch ct.Action { + case config.ConntrackNoTrack: + exprs = append(exprs, &expr.Notrack{}) + case config.ConntrackHelper: + continue + case config.ConntrackDrop: + exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) } - exprs = append(exprs, pe...) - } - switch ct.Action { - case config.ConntrackNoTrack: - exprs = append(exprs, &expr.Notrack{}) - case config.ConntrackHelper: - continue - case config.ConntrackDrop: - exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag + ":" + chain, + }) } - - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag + ":" + chain, - }) } } return nil @@ -272,6 +271,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 { @@ -290,11 +299,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 { @@ -350,51 +362,101 @@ 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) for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { - exprs, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) + matches, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) if err != nil { return err } - if section != "" && section != config.SectionAll { - exprs = append(exprs, matchSection(section)...) - } + for _, m := range matches { + exprs := m.exprs - if logLevel != "" { - exprs = append(exprs, buildLog(logLevel, tag)...) - } + if section != "" && section != config.SectionAll { + exprs = append(exprs, matchSection(section)...) + } - verdict := actionVerdict(action, proto, c.cfg.Settings.AddressFamily) - if verdict != nil { - exprs = append(exprs, verdict...) - } + if logLevel != "" { + exprs = append(exprs, buildLog(logLevel, tag)...) + } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + verdict := actionVerdict(action, m.proto, c.cfg.Settings.AddressFamily) + if verdict != nil { + exprs = append(exprs, verdict...) + } + + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } } 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) @@ -414,102 +476,99 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, srcIfaces := c.resolveZoneInterfaces(srcZone) + matches, err := l4Matches(proto, dports, nil) + if err != nil { + return err + } + + ip := net.ParseIP(dnatAddr) + if ip == nil { + return fmt.Errorf("invalid DNAT address %q", dnatAddr) + } + for _, srcIface := range srcIfaces { - var exprs []expr.Any + for _, m := range matches { + var exprs []expr.Any - if srcIface != "" { - exprs = append(exprs, matchIfaceName(true, srcIface)...) - } - - if srcAddr != "" { - src, err := matchSourceCIDR(srcAddr) - if err != nil { - return err + if srcIface != "" { + exprs = append(exprs, matchIfaceName(true, srcIface)...) } - exprs = append(exprs, src...) - } - if proto != "" { - exprs = append(exprs, matchProto(proto)...) - } - - for _, portStr := range dports { - pe, err := parsePortOrRange(portStr) - if err != nil { - return err - } - exprs = append(exprs, pe...) - } - - if logLevel != "" { - exprs = append(exprs, buildLog(logLevel, tag)...) - } - - ip := net.ParseIP(dnatAddr) - if ip == nil { - return fmt.Errorf("invalid DNAT address %q", dnatAddr) - } - - if action == config.RuleRedirect { - if dnatPort > 0 { - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, dnatPort) - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: portBytes}, - &expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED}, - ) - } else { - exprs = append(exprs, &expr.Redir{}) - } - } else { - if ip4 := ip.To4(); ip4 != nil { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip4}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV4, - RegAddrMin: 1, - RegAddrMax: 1, + if srcAddr != "" { + src, err := matchSourceCIDR(srcAddr) + if err != nil { + return err } + exprs = append(exprs, src...) + } + + exprs = append(exprs, m.exprs...) + + if logLevel != "" { + exprs = append(exprs, buildLog(logLevel, tag)...) + } + + if action == config.RuleRedirect { if dnatPort > 0 { portBytes := make([]byte, 2) binary.BigEndian.PutUint16(portBytes, dnatPort) exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: portBytes}, + &expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + } else { + exprs = append(exprs, &expr.Redir{}) } - exprs = append(exprs, natExpr) } else { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip.To16()}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV6, - RegAddrMin: 1, - RegAddrMax: 1, - } - if dnatPort > 0 { - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, dnatPort) + if ip4 := ip.To4(); ip4 != nil { exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: ip4}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV4, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) + } else { + exprs = append(exprs, + &expr.Immediate{Register: 1, Data: ip.To16()}, + ) + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV6, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) } - exprs = append(exprs, natExpr) } - } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } return nil } @@ -589,25 +648,12 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { exprs = append(exprs, srcExprs...) } - if snat.Proto != "" { - exprs = append(exprs, matchProto(snat.Proto)...) - } - - for _, portStr := range snat.DPort { - pe, err := parsePortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] dport: %w", i, err) - } - exprs = append(exprs, pe...) - } - - for _, portStr := range snat.SPort { - pe, err := parseSPortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] sport: %w", i, err) - } - exprs = append(exprs, pe...) + matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort) + if err != nil { + return fmt.Errorf("snat[%d]: %w", i, err) } + head := exprs + exprs = nil if snat.Mark != "" { exprs = append(exprs, matchMark(snat.Mark)...) @@ -653,11 +699,13 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { } } - state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ - Chain: "postrouting", - Exprs: exprs, - Tag: tag, - }) + 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, + }) + } } return nil @@ -889,10 +937,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 { @@ -924,7 +991,7 @@ func (c *Compiler) expandZoneRef(ref string) []string { return []string{base} } -func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]expr.Any, error) { +func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { var exprs []expr.Any if srcIface != "" { @@ -950,34 +1017,92 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor exprs = append(exprs, dst...) } + matches, err := l4Matches(proto, dports, sports) + if err != nil { + return nil, err + } + for i := range matches { + matches[i].exprs = append(append([]expr.Any{}, exprs...), matches[i].exprs...) + } + return matches, nil +} + +type l4Match struct { + proto byte + exprs []expr.Any +} + +// l4Matches yields one alternative per (proto, dport, sport): nft ANDs a rule's exprs, so lists need one rule each. +func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) { + protos := []string{""} if proto != "" { - exprs = append(exprs, matchProto(proto)...) + protos = strings.Split(proto, ",") } - isICMP := strings.EqualFold(proto, "icmp") || strings.EqualFold(proto, "icmpv6") || strings.EqualFold(proto, "ipv6-icmp") - - for _, portStr := range dports { - if isICMP { - pe := matchICMPType(portStr) - exprs = append(exprs, pe...) - } else { - pe, err := parsePortOrRange(portStr) + 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) + } + var pm []expr.Any + var n byte + isICMP := false + if p != "" { + var err error + n, err = protoNumber(p) if err != nil { return nil, err } - exprs = append(exprs, pe...) + isICMP = n == unix.IPPROTO_ICMP || n == unix.IPPROTO_ICMPV6 + if (len(dports) > 0 && !isICMP || len(sports) > 0) && !hasPorts(n) { + return nil, fmt.Errorf("protocol %q does not support ports", p) + } + pm = matchProtoNum(n) } - } - - for _, portStr := range sports { - pe, err := parseSPortOrRange(portStr) + parseD := parsePortOrRange + if isICMP { + if len(protos) > 1 && len(dports) > 0 { + return nil, fmt.Errorf("dport %v is ambiguous with %s in proto list %q", dports, p, proto) + } + parseD = matchICMPType + } + dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err } - exprs = append(exprs, pe...) + salts, err := portAlternatives(sports, parseSPortOrRange) + if err != nil { + return nil, err + } + for _, d := range dalts { + for _, sp := range salts { + e := append(append(append([]expr.Any{}, pm...), d...), sp...) + out = append(out, l4Match{proto: n, exprs: e}) + } + } } + return out, nil +} - return exprs, nil +func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, error)) ([][]expr.Any, error) { + var alts [][]expr.Any + for _, item := range ports { + for _, s := range strings.Split(item, ",") { + if s = strings.TrimSpace(s); s == "" { + return nil, fmt.Errorf("empty element in port list %q", item) + } + e, err := parse(s) + if err != nil { + return nil, err + } + alts = append(alts, e) + } + } + if len(alts) == 0 { + return [][]expr.Any{nil}, nil + } + return alts, nil } // matchIfaceName matches an interface name, supporting wildcard "+" suffix. @@ -1005,33 +1130,27 @@ func matchIfaceName(input bool, name string) []expr.Any { } } -func matchProto(proto string) []expr.Any { - var protoNum byte - switch strings.ToLower(proto) { - case "tcp": - protoNum = unix.IPPROTO_TCP - case "udp": - protoNum = unix.IPPROTO_UDP - case "icmp": - protoNum = unix.IPPROTO_ICMP - case "icmpv6", "ipv6-icmp": - protoNum = unix.IPPROTO_ICMPV6 - case "gre": - protoNum = 47 - case "esp": - protoNum = 50 - case "ah": - protoNum = 51 - case "sctp": - protoNum = unix.IPPROTO_SCTP - default: - n, _ := strconv.Atoi(proto) - protoNum = byte(n) +var protoNumbers = map[string]byte{ + "icmp": 1, "igmp": 2, "ipip": 4, "ipencap": 4, "tcp": 6, "udp": 17, + "gre": 47, "esp": 50, "ah": 51, "icmpv6": 58, "ipv6-icmp": 58, + "ospf": 89, "ospfigp": 89, "pim": 103, "vrrp": 112, "l2tp": 115, + "sctp": 132, "udplite": 136, +} + +// protoNumber resolves a common IANA protocol name or a 0-255 number. +func protoNumber(proto string) (byte, error) { + if n, ok := protoNumbers[strings.ToLower(proto)]; ok { + return n, nil } - return []expr.Any{ - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}}, + n, err := strconv.ParseUint(proto, 10, 8) + if err != nil { + return 0, fmt.Errorf("unknown protocol %q", proto) } + return byte(n), nil +} + +func hasPorts(proto byte) bool { + return proto == unix.IPPROTO_TCP || proto == unix.IPPROTO_UDP || proto == unix.IPPROTO_SCTP || proto == unix.IPPROTO_UDPLITE } func matchDPort(port uint16) []expr.Any { @@ -1143,33 +1262,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) { @@ -1220,8 +1339,8 @@ func matchConnLimit(spec string) []expr.Any { } func parsePortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid port range low %q: %w", parts[0], err) @@ -1240,8 +1359,8 @@ func parsePortOrRange(s string) ([]expr.Any, error) { } func parseSPortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid sport range low %q: %w", parts[0], err) @@ -1272,6 +1391,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 @@ -1621,7 +1751,7 @@ func logLevelToNF(level string) expr.LogLevel { } } -func actionVerdict(action config.RuleAction, proto string, family config.AddressFamily) []expr.Any { +func actionVerdict(action config.RuleAction, proto byte, family config.AddressFamily) []expr.Any { switch action { case config.RuleAccept: return []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} @@ -1646,8 +1776,8 @@ func actionVerdict(action config.RuleAction, proto string, family config.Address } } -func rejectExprs(proto string, family config.AddressFamily) []expr.Any { - if strings.ToLower(proto) == "tcp" { +func rejectExprs(proto byte, family config.AddressFamily) []expr.Any { + if proto == unix.IPPROTO_TCP { return []expr.Any{&expr.Reject{ Type: unix.NFT_REJECT_TCP_RST, Code: 0, @@ -1666,7 +1796,7 @@ func policyVerdict(action config.PolicyAction, family config.AddressFamily) []ex case config.PolicyDrop: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} case config.PolicyReject: - return rejectExprs("", family) + return rejectExprs(0, family) default: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index e97738c..8c650c8 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1,7 +1,10 @@ package nftables import ( + "encoding/binary" + "fmt" "reflect" + "strings" "testing" "github.com/google/nftables/expr" @@ -975,7 +978,7 @@ func TestNegatedAddress(t *testing.T) { } func TestRejectTCPRST(t *testing.T) { - exprs := rejectExprs("tcp", config.FamilyINET) + exprs := rejectExprs(unix.IPPROTO_TCP, config.FamilyINET) if len(exprs) != 1 { t.Fatalf("expected 1 expr, got %d", len(exprs)) } @@ -984,7 +987,7 @@ func TestRejectTCPRST(t *testing.T) { t.Errorf("TCP reject should use NFT_REJECT_TCP_RST (1), got %d", rej.Type) } - exprs = rejectExprs("udp", config.FamilyINET) + exprs = rejectExprs(unix.IPPROTO_UDP, config.FamilyINET) rej = exprs[0].(*expr.Reject) if rej.Type != 2 { t.Errorf("non-TCP reject should use NFT_REJECT_ICMPX_UNREACH (2), got %d", rej.Type) @@ -1267,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) } @@ -1557,6 +1563,7 @@ func diffTestConfig(port string) *config.Config { {Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}}, {Action: config.RuleNFQueue, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80"}, NFQueue: 3}, {Action: config.RuleRedirect, Source: "loc", Dest: "fw:192.0.2.1:3128", Proto: "tcp", DPort: config.PortSpec{"80"}}, + {Action: config.RuleAccept, Source: "loc,dmz", Dest: "fw,net", Proto: "tcp,udp", DPort: config.PortSpec{"53", "5353"}}, }, SNAT: []config.SNATRule{{Action: config.SNATAddress, Source: "198.51.100.0/24", Dest: "eth0", Address: "203.0.113.7"}}, PortGroups: make(map[string]config.PortGroup), @@ -1631,6 +1638,274 @@ func tags(rules []ManagedRule) []string { return out } +func TestCompile_PortAndProtoLists(t *testing.T) { + type want struct { + proto byte + dport string // "80" for an exact compare, "8000-8100" for a range + } + tests := []struct { + name string + rule config.Rule + chain string + want []want + }{ + { + name: "multi-port", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80", "443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "range in list", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22", "8000:8100"}}, + chain: "input", + want: []want{{6, "22"}, {6, "8000-8100"}}, + }, + { + name: "comma string", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80,443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "tcp,udp", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "input", + want: []want{{6, "53"}, {17, "53"}}, + }, + { + name: "dnat tcp,udp", + rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "prerouting", + want: []want{{6, "53"}, {17, "53"}}, + }, + { + name: "protocol names", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "ospf,OSPFIGP,igmp,gre,esp,ah,vrrp,pim,ipencap,ipv6-icmp"}, + chain: "input", + want: []want{{89, ""}, {89, ""}, {2, ""}, {47, ""}, {50, ""}, {51, ""}, {112, ""}, {103, ""}, {4, ""}, {58, ""}}, + }, + { + name: "protocol numbers", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "0,89,255"}, + chain: "input", + want: []want{{0, ""}, {89, ""}, {255, ""}}, + }, + { + name: "numeric tcp with port", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "6,udplite", DPort: config.PortSpec{"22"}}, + chain: "input", + want: []want{{6, "22"}, {136, "22"}}, + }, + } + + 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}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, + Rules: []config.Rule{tt.rule}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + + var got []want + for _, r := range state.Rules[tt.chain] { + if r.Tag != "rule:0" { + continue + } + var w want + var portCmps []string + for i, e := range r.Exprs { + if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyL4PROTO { + w.proto = r.Exprs[i+1].(*expr.Cmp).Data[0] + } + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == 2 { + for _, c := range r.Exprs[i+1:] { + cmp, ok := c.(*expr.Cmp) + if !ok { + break + } + portCmps = append(portCmps, fmt.Sprint(binary.BigEndian.Uint16(cmp.Data))) + } + } + } + w.dport = strings.Join(portCmps, "-") + got = append(got, w) + } + + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %+v, want %+v", got, tt.want) + } + }) + } +} + +// 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}, + 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), + } + 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 TestDiffEngine_FreshApplyKeepsDesiredOrder(t *testing.T) { desired, err := NewCompiler(diffTestConfig("22")).Compile() if err != nil { @@ -1679,6 +1954,64 @@ func TestDiffEngine_MiddleChangeInsertsBeforeNextRule(t *testing.T) { } } +func TestDiffEngine_ExpandedRuleReplacedInPlace(t *testing.T) { + current := withHandles(mustCompile(t, diffTestConfig("22"))) + if cs := computeDiff(current, mustCompile(t, diffTestConfig("22"))); !cs.Empty() { + t.Fatalf("expected empty changeset, got:\n%s", cs.Summary()) + } + + cfg := diffTestConfig("22") + cfg.Rules[4].DPort = config.PortSpec{"53", "853"} + desired := mustCompile(t, cfg) + cs := computeDiff(current, desired) + for _, r := range append(append([]ManagedRule{}, cs.Add...), cs.Remove...) { + if r.Tag != "rule:4" { + t.Errorf("unexpected change to %s/%s", r.Chain, r.Tag) + } + } + for _, chain := range []string{"input", "forward"} { + if n := len(taggedRules(desired, chain, "rule:4")); n != 8 { + t.Fatalf("%s: expected 8 expanded rule:4 rules, got %d", chain, n) + } + } + + applied := applyChangeSet(current, cs) + if cs := computeDiff(applied, desired); !cs.Empty() { + t.Fatalf("second plan not empty:\n%s", cs.Summary()) + } +} + +// applyChangeSet mimics the engine: removals by handle, adds inserted before r.Before or appended. +func applyChangeSet(s *FirewallState, cs *ChangeSet) *FirewallState { + gone := map[uint64]bool{} + for _, r := range cs.Remove { + gone[r.Handle] = true + } + out := &FirewallState{Rules: map[string][]ManagedRule{}} + for chain, rules := range s.Rules { + for _, r := range rules { + if !gone[r.Handle] { + out.Rules[chain] = append(out.Rules[chain], r) + } + } + } + h := uint64(10000) + for _, r := range cs.Add { + r.Handle, h = h, h+1 + rules := out.Rules[r.Chain] + i := len(rules) + for j, x := range rules { + if r.Before != 0 && x.Handle == r.Before { + i = j + break + } + } + r.Before = 0 + out.Rules[r.Chain] = append(rules[:i], append([]ManagedRule{r}, rules[i:]...)...) + } + return out +} + func mustCompile(t *testing.T, cfg *config.Config) *FirewallState { t.Helper() s, err := NewCompiler(cfg).Compile() @@ -1687,3 +2020,226 @@ func mustCompile(t *testing.T, cfg *config.Config) *FirewallState { } return s } + +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_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 + want []uint32 + }{ + {"tcp,udp", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, + {"6", []uint32{unix.NFT_REJECT_TCP_RST}}, + {"6,17", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, + } + for _, tt := range tests { + t.Run(tt.proto, func(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: tt.proto, DPort: config.PortSpec{"53"}}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, "input", "rule:0") + if len(rules) != len(tt.want) { + t.Fatalf("got %d rules, want %d", len(rules), len(tt.want)) + } + 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 rej.Type != tt.want[i] { + t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, tt.want[i]) + } + } + }) + } +} + +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 + 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"}}}, + {"unknown icmp type with code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus/0"}}}, + {"invalid icmp code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"destination-unreachable/x"}}}, + {"icmp code out of range", config.Rule{Proto: "icmp", DPort: config.PortSpec{"3/256"}}}, + {"unknown proto in list", config.Rule{Proto: "tcp,udpp", DPort: config.PortSpec{"53"}}}, + {"proto number out of range", config.Rule{Proto: "256"}}, + {"unknown proto name", config.Rule{Proto: "bogus"}}, + {"dport with ospf", config.Rule{Proto: "ospf", DPort: config.PortSpec{"80"}}}, + {"sport with gre", config.Rule{Proto: "gre", SPort: config.PortSpec{"80"}}}, + {"dport with icmp in proto list", config.Rule{Proto: "icmp,tcp", DPort: config.PortSpec{"80"}}}, + {"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_ICMPList(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "icmp", DPort: config.PortSpec{"echo-request,echo-reply", "destination-unreachable/4"}}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + var got [][]byte + for _, r := range taggedRules(state, "input", "rule:0") { + var tc []byte + for i, e := range r.Exprs { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Len == 1 { + tc = append(tc, r.Exprs[i+1].(*expr.Cmp).Data[0]) + } + } + got = append(got, tc) + } + want := [][]byte{{8}, {0}, {3, 4}} + if !reflect.DeepEqual(got, want) { + t.Errorf("icmp type/code per rule = %v, want %v", got, want) + } +} + +func TestCompile_ColonRanges(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + offset uint32 + }{ + {"rule sport", 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", 0}, + {"snat sport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "udp", SPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 0}, + {"snat dport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp", DPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 2}, + {"conntrack dport", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"1024:2048"}}} + }, "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) + } + rules := taggedRules(state, tt.chain, tt.tag) + if len(rules) != 1 { + t.Fatalf("got %d %s rules, want 1", len(rules), tt.tag) + } + var got []string + ex := rules[0].Exprs + for i, e := range ex { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == tt.offset && p.Len == 2 && i+2 < len(ex) { + lo, hi := ex[i+1].(*expr.Cmp), ex[i+2].(*expr.Cmp) + got = append(got, fmt.Sprintf("%d>=%d,%d<=%d", lo.Op, binary.BigEndian.Uint16(lo.Data), hi.Op, binary.BigEndian.Uint16(hi.Data))) + } + } + want := []string{fmt.Sprintf("%d>=1024,%d<=2048", expr.CmpOpGte, expr.CmpOpLte)} + if !reflect.DeepEqual(got, want) { + t.Errorf("range match = %v, want %v", got, want) + } + }) + } +}