fix: compare rule expressions by value in diff
This commit is contained in:
@@ -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}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user