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