diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 0bb1f1a..416ca06 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -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} diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index ac36125..623b4c3 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -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 diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index faa9860..e0a8167 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -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 -} diff --git a/internal/nftables/diff_test.go b/internal/nftables/diff_test.go index 33c9b2d..f0eff4c 100644 --- a/internal/nftables/diff_test.go +++ b/internal/nftables/diff_test.go @@ -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") } } diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 5311490..77a4afe 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -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()