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