137 lines
2.8 KiB
Go
137 lines
2.8 KiB
Go
package nftables
|
|
|
|
import (
|
|
"fmt"
|
|
"sort"
|
|
"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[r.Tag] = append(currentByTag[r.Tag], r)
|
|
}
|
|
}
|
|
}
|
|
|
|
desiredByTag := make(map[string][]ManagedRule)
|
|
for _, rules := range desired.Rules {
|
|
for _, r := range rules {
|
|
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], 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
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
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
|
|
}
|
|
as := fmt.Sprintf("%v", a)
|
|
bs := fmt.Sprintf("%v", b)
|
|
return as == bs
|
|
}
|