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 // Helpers are the ct helper objects in kernel (insertion) order. Helpers []Helper } // Helper is a named ct helper object. type Helper struct { Name string `json:"name"` Helper expr.CtHelper `json:"helper"` } type ChangeSet struct { Add []ManagedRule Remove []ManagedRule AddHelpers []Helper RemoveHelpers []string } func (cs *ChangeSet) Empty() bool { return len(cs.Add) == 0 && len(cs.Remove) == 0 && len(cs.AddHelpers) == 0 && len(cs.RemoveHelpers) == 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) } } for _, h := range cs.AddHelpers { fmt.Fprintf(&b, " + ct helper %q\n", h.Name) } for _, n := range cs.RemoveHelpers { fmt.Fprintf(&b, " - ct helper %q\n", n) } 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 := diffHelpers(current, desired) 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{AddHelpers: snap.Helpers} for _, h := range current.Helpers { cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name) } 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 } // diffHelpers replaces any ct helper object that is missing or differs. // L3Proto is ignored: the kernel narrows inet to ip/ip6 for single-family helpers such as pptp. func diffHelpers(current, desired *FirewallState) *ChangeSet { cs := &ChangeSet{} same := func(a, b expr.CtHelper) bool { return a.Name == b.Name && a.L4Proto == b.L4Proto } find := func(hs []Helper, name string) (expr.CtHelper, bool) { for _, h := range hs { if h.Name == name { return h.Helper, true } } return expr.CtHelper{}, false } for _, h := range current.Helpers { if want, ok := find(desired.Helpers, h.Name); !ok || !same(want, h.Helper) { cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name) } } for _, h := range desired.Helpers { if have, ok := find(current.Helpers, h.Name); !ok || !same(have, h.Helper) { cs.AddHelpers = append(cs.AddHelpers, h) } } return cs }