Merge remote-tracking branch 'origin/main' into benvin/diff-expr-equality
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
This commit is contained in:
2026-10-03 21:15:27 +10:00
5 changed files with 1003 additions and 232 deletions
+14 -6
View File
@@ -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)
}
}
}
}
+46
View File
@@ -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)
}
}
}
+40 -9
View File
@@ -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 {
+344 -214
View File
@@ -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}}
}
+559 -3
View File
@@ -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)
}
})
}
}