Files
tomswall/internal/nftables/diff.go
T
unkin-agent e16d95fb63
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
fix: compare rule expressions by value in diff
2026-10-03 20:48:29 +10:00

114 lines
2.3 KiB
Go

package nftables
import (
"fmt"
"reflect"
"strings"
"github.com/google/nftables/expr"
)
type ManagedRule struct {
Chain string
Handle uint64
Exprs []expr.Any
Tag string
}
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 {
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()
}
func computeDiff(current, desired *FirewallState) *ChangeSet {
cs := &ChangeSet{}
currentByTag := make(map[string][]ManagedRule)
for _, rules := range current.Rules {
for _, r := range rules {
if r.Tag != "" {
currentByTag[ruleKey(r)] = append(currentByTag[ruleKey(r)], r)
}
}
}
desiredByTag := make(map[string][]ManagedRule)
for _, rules := range desired.Rules {
for _, r := range rules {
desiredByTag[ruleKey(r)] = append(desiredByTag[ruleKey(r)], r)
}
}
for tag, desiredRules := range desiredByTag {
currentRules, exists := currentByTag[tag]
if !exists {
cs.Add = append(cs.Add, desiredRules...)
continue
}
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...)
}
}
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)
}