Merge pull request 'fix: stop differential apply rewriting unchanged rules' (#18) from benvin/diff-expr-equality into main
ci/woodpecker/tag/release Pipeline was successful

Reviewed-on: #18
This commit was merged in pull request #18.
This commit is contained in:
2026-10-03 21:24:48 +10:00
5 changed files with 273 additions and 54 deletions
+4 -2
View File
@@ -5,6 +5,7 @@ import (
"fmt"
"log/slog"
"net"
"sort"
"strconv"
"strings"
@@ -334,7 +335,7 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
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 {
@@ -513,7 +514,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
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{})
@@ -984,6 +985,7 @@ func (c *Compiler) expandZoneRef(ref string) []string {
}
zones = append(zones, name)
}
sort.Strings(zones)
return zones
}
return []string{base}
+213
View File
@@ -1552,6 +1552,104 @@ 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"}},
{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),
}
}
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())
}
}
// shapes observed by applying diffTestConfig in a netns and reading it back
func TestCompile_QueueRedirMatchKernelReadback(t *testing.T) {
state, err := NewCompiler(diffTestConfig("22")).Compile()
if err != nil {
t.Fatal(err)
}
want := map[string]expr.Any{
"rule:2": &expr.Queue{Num: 3, Total: 1},
"rule:3": &expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
}
for _, rules := range state.Rules {
for _, r := range rules {
w, ok := want[r.Tag]
if !ok {
continue
}
if got := r.Exprs[len(r.Exprs)-1]; !reflect.DeepEqual(got, w) {
t.Errorf("%s: got %#v, want %#v", r.Tag, got, w)
}
delete(want, r.Tag)
}
}
for tag := range want {
t.Errorf("%s not compiled", tag)
}
}
func withHandles(s *FirewallState) *FirewallState {
h := uint64(100)
for _, rules := range s.Rules {
for i := range rules {
rules[i].Handle = h
h++
}
}
return s
}
func tags(rules []ManagedRule) []string {
out := make([]string, len(rules))
for i, r := range rules {
out[i] = r.Chain + "/" + r.Tag
}
return out
}
func TestCompile_PortAndProtoLists(t *testing.T) {
type want struct {
proto byte
@@ -1820,6 +1918,121 @@ func taggedRules(state *FirewallState, chain, tag string) []ManagedRule {
return out
}
func TestDiffEngine_FreshApplyKeepsDesiredOrder(t *testing.T) {
desired, err := NewCompiler(diffTestConfig("22")).Compile()
if err != nil {
t.Fatal(err)
}
var want []string
for _, chain := range []string{"forward", "input", "output", "postrouting", "prerouting"} {
want = append(want, tags(desired.Rules[chain])...)
}
for i := 0; i < 20; i++ {
cs := computeDiff(&FirewallState{Rules: map[string][]ManagedRule{}}, desired)
if got := tags(cs.Add); !reflect.DeepEqual(got, want) {
t.Fatalf("add order:\n got %v\nwant %v", got, want)
}
for _, r := range cs.Add {
if r.Before != 0 {
t.Fatalf("fresh apply should append, %s has Before=%d", r.Tag, r.Before)
}
}
}
}
func TestDiffEngine_MiddleChangeInsertsBeforeNextRule(t *testing.T) {
current := withHandles(mustCompile(t, diffTestConfig("22")))
desired := mustCompile(t, diffTestConfig("2222"))
cs := computeDiff(current, desired)
if len(cs.Add) != 1 || len(cs.Remove) != 1 {
t.Fatalf("expected one replace, got:\n%s", cs.Summary())
}
input := current.Rules["input"]
idx := -1
for i, r := range input {
if r.Tag == "rule:0" {
idx = i
}
}
if idx < 1 || idx == len(input)-1 {
t.Fatalf("rule:0 at %d is not mid-chain", idx)
}
if cs.Remove[0].Handle != input[idx].Handle {
t.Errorf("removed handle %d, want %d", cs.Remove[0].Handle, input[idx].Handle)
}
if cs.Add[0].Before != input[idx+1].Handle {
t.Errorf("insert before %d, want %d (%s)", cs.Add[0].Before, input[idx+1].Handle, input[idx+1].Tag)
}
}
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()
if err != nil {
t.Fatal(err)
}
return s
}
func TestCompile_ListExpansionCounts(t *testing.T) {
tests := []struct {
name string
+46 -49
View File
@@ -2,6 +2,7 @@ package nftables
import (
"fmt"
"reflect"
"sort"
"strings"
@@ -13,6 +14,8 @@ type ManagedRule struct {
Handle uint64
Exprs []expr.Any
Tag string
// Before is the handle of the existing rule an added rule is inserted ahead of; 0 appends.
Before uint64
}
type FirewallState struct {
@@ -33,7 +36,11 @@ func (cs *ChangeSet) Summary() string {
if len(cs.Add) > 0 {
fmt.Fprintf(&b, " + %d rule(s) to add\n", len(cs.Add))
for _, r := range cs.Add {
fmt.Fprintf(&b, " + [%s] %s\n", r.Chain, r.Tag)
if r.Before != 0 {
fmt.Fprintf(&b, " + [%s] %s (before handle %d)\n", r.Chain, r.Tag, r.Before)
} else {
fmt.Fprintf(&b, " + [%s] %s\n", r.Chain, r.Tag)
}
}
}
if len(cs.Remove) > 0 {
@@ -45,46 +52,60 @@ func (cs *ChangeSet) Summary() string {
return b.String()
}
// Rules are first-match, so order matters: keep the common prefix and suffix of
// each chain, replace the middle, and insert the new rules before the first kept
// suffix rule (or append when there is none).
func computeDiff(current, desired *FirewallState) *ChangeSet {
cs := &ChangeSet{}
currentByTag := make(map[string][]ManagedRule)
for _, rules := range current.Rules {
for _, r := range rules {
chains := make([]string, 0, len(current.Rules)+len(desired.Rules))
for c := range current.Rules {
chains = append(chains, c)
}
for c := range desired.Rules {
if _, ok := current.Rules[c]; !ok {
chains = append(chains, c)
}
}
sort.Strings(chains)
for _, chain := range chains {
var cur []ManagedRule
for _, r := range current.Rules[chain] {
if r.Tag != "" {
currentByTag[r.Tag] = append(currentByTag[r.Tag], r)
cur = append(cur, r)
}
}
}
want := desired.Rules[chain]
desiredByTag := make(map[string][]ManagedRule)
for _, rules := range desired.Rules {
for _, r := range rules {
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], r)
pre := 0
for pre < len(cur) && pre < len(want) && ruleEqual(cur[pre], want[pre]) {
pre++
}
suf := 0
for suf < len(cur)-pre && suf < len(want)-pre &&
ruleEqual(cur[len(cur)-1-suf], want[len(want)-1-suf]) {
suf++
}
}
for tag, desiredRules := range desiredByTag {
currentRules, exists := currentByTag[tag]
if !exists {
cs.Add = append(cs.Add, desiredRules...)
continue
var before uint64
if suf > 0 {
before = cur[len(cur)-suf].Handle
}
if !rulesMatch(currentRules, desiredRules) {
cs.Remove = append(cs.Remove, currentRules...)
cs.Add = append(cs.Add, desiredRules...)
}
}
for tag, currentRules := range currentByTag {
if _, exists := desiredByTag[tag]; !exists {
cs.Remove = append(cs.Remove, currentRules...)
cs.Remove = append(cs.Remove, cur[pre:len(cur)-suf]...)
for _, r := range want[pre : len(want)-suf] {
r.Before = before
cs.Add = append(cs.Add, r)
}
}
return cs
}
func ruleEqual(a, b ManagedRule) bool {
return a.Chain == b.Chain && a.Tag == b.Tag && reflect.DeepEqual(a.Exprs, b.Exprs)
}
// restoreChangeSet replaces every managed rule in current with the snapshot's,
// in snapshot order, so a restore cannot reorder rules.
func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
@@ -110,27 +131,3 @@ func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
}
return cs
}
func rulesMatch(a, b []ManagedRule) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i].Chain != b[i].Chain {
return false
}
if !exprsEqual(a[i].Exprs, b[i].Exprs) {
return false
}
}
return true
}
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
}
+2 -1
View File
@@ -1,6 +1,7 @@
package nftables
import (
"reflect"
"testing"
"github.com/google/nftables/expr"
@@ -49,7 +50,7 @@ func TestRestoreChangeSet(t *testing.T) {
t.Fatalf("added %v, want %v (snapshot order per chain)", added, want)
}
}
if !exprsEqual(cs.Add[1].Exprs, accept) {
if !reflect.DeepEqual(cs.Add[1].Exprs, accept) {
t.Error("ssh not restored to its snapshot exprs")
}
}
+8 -2
View File
@@ -117,12 +117,18 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol
if !ok {
return fmt.Errorf("unknown chain %q", r.Chain)
}
e.conn.AddRule(&nftables.Rule{
rule := &nftables.Rule{
Table: table,
Chain: chain,
Exprs: r.Exprs,
UserData: []byte(r.Tag),
})
}
if r.Before != 0 {
rule.Position = r.Before
e.conn.InsertRule(rule)
} else {
e.conn.AddRule(rule)
}
}
return e.conn.Flush()