fix: compare rule expressions by value in diff
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-03 20:48:29 +10:00
parent 410109515e
commit e16d95fb63
3 changed files with 62 additions and 7 deletions
+4 -2
View File
@@ -4,6 +4,7 @@ import (
"encoding/binary"
"fmt"
"net"
"sort"
"strconv"
"strings"
@@ -322,7 +323,7 @@ func (c *Compiler) applyRuleExtras(state *FirewallState, tag string, rule config
extra = append(extra, setMarkExprs(rule.SetMark)...)
}
if rule.Action == config.RuleNFQueue {
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue)}}
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}}
}
if len(extra) > 0 || len(replaceVerdict) > 0 {
@@ -455,7 +456,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec,
binary.BigEndian.PutUint16(portBytes, dnatPort)
exprs = append(exprs,
&expr.Immediate{Register: 1, Data: portBytes},
&expr.Redir{RegisterProtoMin: 1},
&expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
)
} else {
exprs = append(exprs, &expr.Redir{})
@@ -917,6 +918,7 @@ func (c *Compiler) expandZoneRef(ref string) []string {
}
zones = append(zones, name)
}
sort.Strings(zones)
return zones
}
return []string{base}
+49
View File
@@ -1531,3 +1531,52 @@ func TestCompile_PolicyRateLimit(t *testing.T) {
}
t.Error("policy:0 not found in input chain")
}
func diffTestConfig(port string) *config.Config {
return &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"loc": {Type: config.ZoneIP},
"dmz": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
{Zone: "dmz", Interface: "eth2"},
},
Policy: []config.Policy{
{Source: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"},
{Source: "all", Dest: "all", Action: config.PolicyReject},
},
Rules: []config.Rule{
{Action: config.RuleAccept, Source: "net:192.0.2.0/24", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{port}},
{Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}},
},
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),
}
}
func TestDiffEngine_IndependentCompilesMatch(t *testing.T) {
compile := func(port string) *FirewallState {
t.Helper()
s, err := NewCompiler(diffTestConfig(port)).Compile()
if err != nil {
t.Fatal(err)
}
return s
}
for i := 0; i < 20; i++ {
if cs := computeDiff(compile("22"), compile("22")); !cs.Empty() {
t.Fatalf("expected empty changeset, got:\n%s", cs.Summary())
}
}
cs := computeDiff(compile("22"), compile("2222"))
if len(cs.Add) != 1 || len(cs.Remove) != 1 || cs.Add[0].Tag != "rule:0" {
t.Errorf("expected rule:0 replaced, got:\n%s", cs.Summary())
}
}
+9 -5
View File
@@ -2,6 +2,7 @@ package nftables
import (
"fmt"
"reflect"
"strings"
"github.com/google/nftables/expr"
@@ -51,7 +52,7 @@ func computeDiff(current, desired *FirewallState) *ChangeSet {
for _, rules := range current.Rules {
for _, r := range rules {
if r.Tag != "" {
currentByTag[r.Tag] = append(currentByTag[r.Tag], r)
currentByTag[ruleKey(r)] = append(currentByTag[ruleKey(r)], r)
}
}
}
@@ -59,7 +60,7 @@ func computeDiff(current, desired *FirewallState) *ChangeSet {
desiredByTag := make(map[string][]ManagedRule)
for _, rules := range desired.Rules {
for _, r := range rules {
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], r)
desiredByTag[ruleKey(r)] = append(desiredByTag[ruleKey(r)], r)
}
}
@@ -84,6 +85,11 @@ func computeDiff(current, desired *FirewallState) *ChangeSet {
return cs
}
// a tag can span chains and map iteration order is random, so group per chain
func ruleKey(r ManagedRule) string {
return r.Chain + "\x00" + r.Tag
}
func rulesMatch(a, b []ManagedRule) bool {
if len(a) != len(b) {
return false
@@ -103,7 +109,5 @@ func exprsEqual(a, b []expr.Any) bool {
if len(a) != len(b) {
return false
}
as := fmt.Sprintf("%v", a)
bs := fmt.Sprintf("%v", b)
return as == bs
return reflect.DeepEqual(a, b)
}