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