fix: preserve rule order in diff and insert replacements in place
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:52:01 +10:00
parent e16d95fb63
commit df7ebb8efe
3 changed files with 159 additions and 52 deletions
+107
View File
@@ -1,9 +1,11 @@
package nftables
import (
"reflect"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
@@ -1553,6 +1555,8 @@ func diffTestConfig(port string) *config.Config {
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"}},
},
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),
@@ -1580,3 +1584,106 @@ func TestDiffEngine_IndependentCompilesMatch(t *testing.T) {
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 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 mustCompile(t *testing.T, cfg *config.Config) *FirewallState {
t.Helper()
s, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatal(err)
}
return s
}
+44 -50
View File
@@ -3,6 +3,7 @@ package nftables
import (
"fmt"
"reflect"
"sort"
"strings"
"github.com/google/nftables/expr"
@@ -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,69 +52,56 @@ 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[ruleKey(r)] = append(currentByTag[ruleKey(r)], r)
cur = append(cur, r)
}
}
}
want := desired.Rules[chain]
desiredByTag := make(map[string][]ManagedRule)
for _, rules := range desired.Rules {
for _, r := range rules {
desiredByTag[ruleKey(r)] = append(desiredByTag[ruleKey(r)], 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
}
// 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
}
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
}
return reflect.DeepEqual(a, b)
func ruleEqual(a, b ManagedRule) bool {
return a.Chain == b.Chain && a.Tag == b.Tag && reflect.DeepEqual(a.Exprs, b.Exprs)
}
+8 -2
View File
@@ -109,12 +109,18 @@ func (e *Engine) Apply(changes *ChangeSet) error {
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()