114 lines
2.3 KiB
Go
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)
|
|
}
|