Files
tomswall/internal/nftables/diff.go
T
unkin-agent 1ad025a4b7
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Merge remote-tracking branch 'origin/main' into benvin/diff-expr-equality
# Conflicts:
#	internal/nftables/diff.go
2026-10-03 21:16:39 +10:00

134 lines
3.0 KiB
Go

package nftables
import (
"fmt"
"reflect"
"sort"
"strings"
"github.com/google/nftables/expr"
)
type ManagedRule struct {
Chain string
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 {
Rules map[string][]ManagedRule
}
type ChangeSet struct {
Add []ManagedRule
Remove []ManagedRule
}
func (cs *ChangeSet) Empty() bool {
return len(cs.Add) == 0 && len(cs.Remove) == 0
}
func (cs *ChangeSet) Summary() string {
var b strings.Builder
if len(cs.Add) > 0 {
fmt.Fprintf(&b, " + %d rule(s) to add\n", len(cs.Add))
for _, r := range cs.Add {
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 {
fmt.Fprintf(&b, " - %d rule(s) to remove\n", len(cs.Remove))
for _, r := range cs.Remove {
fmt.Fprintf(&b, " - [%s] %s (handle %d)\n", r.Chain, r.Tag, r.Handle)
}
}
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{}
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 != "" {
cur = append(cur, r)
}
}
want := desired.Rules[chain]
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++
}
var before uint64
if suf > 0 {
before = cur[len(cur)-suf].Handle
}
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
}
func ruleEqual(a, b ManagedRule) bool {
return a.Chain == b.Chain && a.Tag == b.Tag && reflect.DeepEqual(a.Exprs, b.Exprs)
}
// restoreChangeSet replaces every managed rule in current with the snapshot's,
// in snapshot order, so a restore cannot reorder rules.
func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
cs := &ChangeSet{}
for _, rules := range current.Rules {
for _, r := range rules {
if r.Tag != "" {
cs.Remove = append(cs.Remove, r)
}
}
}
chains := make([]string, 0, len(snap.Rules))
for c := range snap.Rules {
chains = append(chains, c)
}
sort.Strings(chains)
for _, c := range chains {
for _, r := range snap.Rules[c] {
if r.Tag != "" {
cs.Add = append(cs.Add, r)
}
}
}
return cs
}