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()