From e16d95fb63e741cd54d0fa07190758f2c2029f84 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:48:29 +1000 Subject: [PATCH 1/2] fix: compare rule expressions by value in diff --- internal/nftables/compiler.go | 6 ++-- internal/nftables/compiler_test.go | 49 ++++++++++++++++++++++++++++++ internal/nftables/diff.go | 14 ++++++--- 3 files changed, 62 insertions(+), 7 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 45e6549..23137a2 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -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} diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 84fa751..0cba130 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -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()) + } +} diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index 82cfa30..85cd3c8 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -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) } From df7ebb8efe9180eaa4366a33515facaf7a81417b Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:52:01 +1000 Subject: [PATCH 2/2] fix: preserve rule order in diff and insert replacements in place --- internal/nftables/compiler_test.go | 107 +++++++++++++++++++++++++++++ internal/nftables/diff.go | 94 ++++++++++++------------- internal/nftables/engine.go | 10 ++- 3 files changed, 159 insertions(+), 52 deletions(-) diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 0cba130..e97738c 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -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 +} diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index 85cd3c8..e4bd988 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -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) } diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 0618a6f..32616cc 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -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()