fix: preserve rule order in diff and insert replacements in place
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user