Merge pull request 'fix: stop differential apply rewriting unchanged rules' (#18) from benvin/diff-expr-equality into main
ci/woodpecker/tag/release Pipeline was successful
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #18
This commit was merged in pull request #18.
This commit is contained in:
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
@@ -334,7 +335,7 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
|
||||
extra = append(extra, setMarkExprs(rule.SetMark)...)
|
||||
}
|
||||
if rule.Action == config.RuleNFQueue {
|
||||
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue)}}
|
||||
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}}
|
||||
}
|
||||
|
||||
if len(extra) > 0 || len(replaceVerdict) > 0 {
|
||||
@@ -513,7 +514,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
|
||||
binary.BigEndian.PutUint16(portBytes, dnatPort)
|
||||
exprs = append(exprs,
|
||||
&expr.Immediate{Register: 1, Data: portBytes},
|
||||
&expr.Redir{RegisterProtoMin: 1},
|
||||
&expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
|
||||
)
|
||||
} else {
|
||||
exprs = append(exprs, &expr.Redir{})
|
||||
@@ -984,6 +985,7 @@ func (c *Compiler) expandZoneRef(ref string) []string {
|
||||
}
|
||||
zones = append(zones, name)
|
||||
}
|
||||
sort.Strings(zones)
|
||||
return zones
|
||||
}
|
||||
return []string{base}
|
||||
|
||||
@@ -1552,6 +1552,104 @@ func TestCompile_PolicyRateLimit(t *testing.T) {
|
||||
t.Error("policy:0 not found in input chain")
|
||||
}
|
||||
|
||||
func diffTestConfig(port string) *config.Config {
|
||||
return &config.Config{
|
||||
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
|
||||
Zones: map[string]config.Zone{
|
||||
"fw": {Type: config.ZoneFirewall},
|
||||
"net": {Type: config.ZoneIP},
|
||||
"loc": {Type: config.ZoneIP},
|
||||
"dmz": {Type: config.ZoneIP},
|
||||
},
|
||||
Interfaces: []config.Interface{
|
||||
{Zone: "net", Interface: "eth0"},
|
||||
{Zone: "loc", Interface: "eth1"},
|
||||
{Zone: "dmz", Interface: "eth2"},
|
||||
},
|
||||
Policy: []config.Policy{
|
||||
{Source: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"},
|
||||
{Source: "all", Dest: "all", Action: config.PolicyReject},
|
||||
},
|
||||
Rules: []config.Rule{
|
||||
{Action: config.RuleAccept, Source: "net:192.0.2.0/24", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{port}},
|
||||
{Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}},
|
||||
{Action: config.RuleNFQueue, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80"}, NFQueue: 3},
|
||||
{Action: config.RuleRedirect, Source: "loc", Dest: "fw:192.0.2.1:3128", Proto: "tcp", DPort: config.PortSpec{"80"}},
|
||||
{Action: config.RuleAccept, Source: "loc,dmz", Dest: "fw,net", Proto: "tcp,udp", DPort: config.PortSpec{"53", "5353"}},
|
||||
},
|
||||
SNAT: []config.SNATRule{{Action: config.SNATAddress, Source: "198.51.100.0/24", Dest: "eth0", Address: "203.0.113.7"}},
|
||||
PortGroups: make(map[string]config.PortGroup),
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffEngine_IndependentCompilesMatch(t *testing.T) {
|
||||
compile := func(port string) *FirewallState {
|
||||
t.Helper()
|
||||
s, err := NewCompiler(diffTestConfig(port)).Compile()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
for i := 0; i < 20; i++ {
|
||||
if cs := computeDiff(compile("22"), compile("22")); !cs.Empty() {
|
||||
t.Fatalf("expected empty changeset, got:\n%s", cs.Summary())
|
||||
}
|
||||
}
|
||||
|
||||
cs := computeDiff(compile("22"), compile("2222"))
|
||||
if len(cs.Add) != 1 || len(cs.Remove) != 1 || cs.Add[0].Tag != "rule:0" {
|
||||
t.Errorf("expected rule:0 replaced, got:\n%s", cs.Summary())
|
||||
}
|
||||
}
|
||||
|
||||
// shapes observed by applying diffTestConfig in a netns and reading it back
|
||||
func TestCompile_QueueRedirMatchKernelReadback(t *testing.T) {
|
||||
state, err := NewCompiler(diffTestConfig("22")).Compile()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := map[string]expr.Any{
|
||||
"rule:2": &expr.Queue{Num: 3, Total: 1},
|
||||
"rule:3": &expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
|
||||
}
|
||||
for _, rules := range state.Rules {
|
||||
for _, r := range rules {
|
||||
w, ok := want[r.Tag]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if got := r.Exprs[len(r.Exprs)-1]; !reflect.DeepEqual(got, w) {
|
||||
t.Errorf("%s: got %#v, want %#v", r.Tag, got, w)
|
||||
}
|
||||
delete(want, r.Tag)
|
||||
}
|
||||
}
|
||||
for tag := range want {
|
||||
t.Errorf("%s not compiled", tag)
|
||||
}
|
||||
}
|
||||
|
||||
func withHandles(s *FirewallState) *FirewallState {
|
||||
h := uint64(100)
|
||||
for _, rules := range s.Rules {
|
||||
for i := range rules {
|
||||
rules[i].Handle = h
|
||||
h++
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func tags(rules []ManagedRule) []string {
|
||||
out := make([]string, len(rules))
|
||||
for i, r := range rules {
|
||||
out[i] = r.Chain + "/" + r.Tag
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestCompile_PortAndProtoLists(t *testing.T) {
|
||||
type want struct {
|
||||
proto byte
|
||||
@@ -1820,6 +1918,121 @@ func taggedRules(state *FirewallState, chain, tag string) []ManagedRule {
|
||||
return out
|
||||
}
|
||||
|
||||
func TestDiffEngine_FreshApplyKeepsDesiredOrder(t *testing.T) {
|
||||
desired, err := NewCompiler(diffTestConfig("22")).Compile()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var want []string
|
||||
for _, chain := range []string{"forward", "input", "output", "postrouting", "prerouting"} {
|
||||
want = append(want, tags(desired.Rules[chain])...)
|
||||
}
|
||||
for i := 0; i < 20; i++ {
|
||||
cs := computeDiff(&FirewallState{Rules: map[string][]ManagedRule{}}, desired)
|
||||
if got := tags(cs.Add); !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("add order:\n got %v\nwant %v", got, want)
|
||||
}
|
||||
for _, r := range cs.Add {
|
||||
if r.Before != 0 {
|
||||
t.Fatalf("fresh apply should append, %s has Before=%d", r.Tag, r.Before)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffEngine_MiddleChangeInsertsBeforeNextRule(t *testing.T) {
|
||||
current := withHandles(mustCompile(t, diffTestConfig("22")))
|
||||
desired := mustCompile(t, diffTestConfig("2222"))
|
||||
|
||||
cs := computeDiff(current, desired)
|
||||
if len(cs.Add) != 1 || len(cs.Remove) != 1 {
|
||||
t.Fatalf("expected one replace, got:\n%s", cs.Summary())
|
||||
}
|
||||
input := current.Rules["input"]
|
||||
idx := -1
|
||||
for i, r := range input {
|
||||
if r.Tag == "rule:0" {
|
||||
idx = i
|
||||
}
|
||||
}
|
||||
if idx < 1 || idx == len(input)-1 {
|
||||
t.Fatalf("rule:0 at %d is not mid-chain", idx)
|
||||
}
|
||||
if cs.Remove[0].Handle != input[idx].Handle {
|
||||
t.Errorf("removed handle %d, want %d", cs.Remove[0].Handle, input[idx].Handle)
|
||||
}
|
||||
if cs.Add[0].Before != input[idx+1].Handle {
|
||||
t.Errorf("insert before %d, want %d (%s)", cs.Add[0].Before, input[idx+1].Handle, input[idx+1].Tag)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffEngine_ExpandedRuleReplacedInPlace(t *testing.T) {
|
||||
current := withHandles(mustCompile(t, diffTestConfig("22")))
|
||||
if cs := computeDiff(current, mustCompile(t, diffTestConfig("22"))); !cs.Empty() {
|
||||
t.Fatalf("expected empty changeset, got:\n%s", cs.Summary())
|
||||
}
|
||||
|
||||
cfg := diffTestConfig("22")
|
||||
cfg.Rules[4].DPort = config.PortSpec{"53", "853"}
|
||||
desired := mustCompile(t, cfg)
|
||||
cs := computeDiff(current, desired)
|
||||
for _, r := range append(append([]ManagedRule{}, cs.Add...), cs.Remove...) {
|
||||
if r.Tag != "rule:4" {
|
||||
t.Errorf("unexpected change to %s/%s", r.Chain, r.Tag)
|
||||
}
|
||||
}
|
||||
for _, chain := range []string{"input", "forward"} {
|
||||
if n := len(taggedRules(desired, chain, "rule:4")); n != 8 {
|
||||
t.Fatalf("%s: expected 8 expanded rule:4 rules, got %d", chain, n)
|
||||
}
|
||||
}
|
||||
|
||||
applied := applyChangeSet(current, cs)
|
||||
if cs := computeDiff(applied, desired); !cs.Empty() {
|
||||
t.Fatalf("second plan not empty:\n%s", cs.Summary())
|
||||
}
|
||||
}
|
||||
|
||||
// applyChangeSet mimics the engine: removals by handle, adds inserted before r.Before or appended.
|
||||
func applyChangeSet(s *FirewallState, cs *ChangeSet) *FirewallState {
|
||||
gone := map[uint64]bool{}
|
||||
for _, r := range cs.Remove {
|
||||
gone[r.Handle] = true
|
||||
}
|
||||
out := &FirewallState{Rules: map[string][]ManagedRule{}}
|
||||
for chain, rules := range s.Rules {
|
||||
for _, r := range rules {
|
||||
if !gone[r.Handle] {
|
||||
out.Rules[chain] = append(out.Rules[chain], r)
|
||||
}
|
||||
}
|
||||
}
|
||||
h := uint64(10000)
|
||||
for _, r := range cs.Add {
|
||||
r.Handle, h = h, h+1
|
||||
rules := out.Rules[r.Chain]
|
||||
i := len(rules)
|
||||
for j, x := range rules {
|
||||
if r.Before != 0 && x.Handle == r.Before {
|
||||
i = j
|
||||
break
|
||||
}
|
||||
}
|
||||
r.Before = 0
|
||||
out.Rules[r.Chain] = append(rules[:i], append([]ManagedRule{r}, rules[i:]...)...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func mustCompile(t *testing.T, cfg *config.Config) *FirewallState {
|
||||
t.Helper()
|
||||
s, err := NewCompiler(cfg).Compile()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func TestCompile_ListExpansionCounts(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
+46
-49
@@ -2,6 +2,7 @@ package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
@@ -13,6 +14,8 @@ type ManagedRule struct {
|
||||
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 {
|
||||
@@ -33,7 +36,11 @@ func (cs *ChangeSet) Summary() string {
|
||||
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 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 {
|
||||
@@ -45,46 +52,60 @@ func (cs *ChangeSet) Summary() string {
|
||||
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{}
|
||||
|
||||
currentByTag := make(map[string][]ManagedRule)
|
||||
for _, rules := range current.Rules {
|
||||
for _, r := range rules {
|
||||
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 != "" {
|
||||
currentByTag[r.Tag] = append(currentByTag[r.Tag], r)
|
||||
cur = append(cur, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
want := desired.Rules[chain]
|
||||
|
||||
desiredByTag := make(map[string][]ManagedRule)
|
||||
for _, rules := range desired.Rules {
|
||||
for _, r := range rules {
|
||||
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], r)
|
||||
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++
|
||||
}
|
||||
}
|
||||
|
||||
for tag, desiredRules := range desiredByTag {
|
||||
currentRules, exists := currentByTag[tag]
|
||||
if !exists {
|
||||
cs.Add = append(cs.Add, desiredRules...)
|
||||
continue
|
||||
var before uint64
|
||||
if suf > 0 {
|
||||
before = cur[len(cur)-suf].Handle
|
||||
}
|
||||
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...)
|
||||
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 {
|
||||
@@ -110,27 +131,3 @@ func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
@@ -49,7 +50,7 @@ func TestRestoreChangeSet(t *testing.T) {
|
||||
t.Fatalf("added %v, want %v (snapshot order per chain)", added, want)
|
||||
}
|
||||
}
|
||||
if !exprsEqual(cs.Add[1].Exprs, accept) {
|
||||
if !reflect.DeepEqual(cs.Add[1].Exprs, accept) {
|
||||
t.Error("ssh not restored to its snapshot exprs")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -117,12 +117,18 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol
|
||||
if !ok {
|
||||
return fmt.Errorf("unknown chain %q", r.Chain)
|
||||
}
|
||||
e.conn.AddRule(&nftables.Rule{
|
||||
rule := &nftables.Rule{
|
||||
Table: table,
|
||||
Chain: chain,
|
||||
Exprs: r.Exprs,
|
||||
UserData: []byte(r.Tag),
|
||||
})
|
||||
}
|
||||
if r.Before != 0 {
|
||||
rule.Position = r.Before
|
||||
e.conn.InsertRule(rule)
|
||||
} else {
|
||||
e.conn.AddRule(rule)
|
||||
}
|
||||
}
|
||||
|
||||
return e.conn.Flush()
|
||||
|
||||
Reference in New Issue
Block a user