Keep firewall rule bookkeeping in step with the kernel on replace and teardown

This commit is contained in:
Viktor Liu
2026-08-21 08:52:17 +02:00
parent 315837d4b4
commit f0905d89f6
7 changed files with 317 additions and 41 deletions

View File

@@ -358,6 +358,14 @@ func (r *family) applyNetwork(flag string, network firewall.Network, prefixes []
}
if network.IsSet() {
// A destination set is populated later from DNS results, so unlike a
// source set it cannot be expanded into per-prefix rules. Without
// ipset such a rule is not expressible; report it instead of
// installing something broader than the policy allows.
if flag == "-d" && !r.ipsetSupported {
return nil, fmt.Errorf("destination set %s requires ipset (ip_set_hash_net and xt_set)", network.Set.HashedName())
}
name := r.ipsetName(network.Set.HashedName())
if _, err := r.ipsetCounter.Increment(name, prefixes); err != nil {
return nil, fmt.Errorf("create or get ipset: %w", err)

View File

@@ -5,16 +5,19 @@ package iptables
import (
"fmt"
"net/netip"
"slices"
"strings"
"testing"
"time"
"github.com/coreos/go-iptables/iptables"
"github.com/lrh3321/ipset-go"
"github.com/stretchr/testify/require"
fw "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/iface"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/shared/management/domain"
)
var ifaceMock = &iFaceMock{
@@ -97,10 +100,7 @@ func TestIptablesManager(t *testing.T) {
ok, err := ipv4Client.ChainExists("filter", chainACLInput)
require.NoError(t, err, "failed check chain exists")
if ok {
require.NoErrorf(t, err, "chain '%v' still exists after Close", chainACLInput)
}
require.Falsef(t, ok, "chain %q still exists after Close", chainACLInput)
})
}
@@ -285,14 +285,80 @@ func TestIptablesFilterIPSetFallback(t *testing.T) {
// The rule must actually be present in the ACL chain (not silently dropped).
checkRuleSpecs(t, ipv4Client, rr.chain, true, fs.specs...)
// Every expanded peer rule keeps its own redirect-mark pairing.
require.NotNil(t, fs.mangleSpecs, "peer rule must carry a mangle pairing")
checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, true, fs.mangleSpecs...)
}
require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete fallback rule")
for _, fs := range all {
checkRuleSpecs(t, ipv4Client, rr.chain, false, fs.specs...)
checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, false, fs.mangleSpecs...)
}
}
// TestIptablesFilterDestinationSetRequiresIPSet documents that a dynamic
// (domain) destination cannot be expressed without ipset: its prefixes are only
// known after DNS resolution, so there is nothing to expand into per-prefix
// rules. The call must report that rather than install a broader rule than the
// policy allows.
func TestIptablesFilterDestinationSetRequiresIPSet(t *testing.T) {
manager, err := Create(ifaceMock, iface.DefaultMTU)
require.NoError(t, err)
require.NoError(t, manager.Init(nil))
defer func() {
require.NoError(t, manager.Close(nil))
}()
manager.family4.ipsetSupported = false
destination := fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})}
_, err = manager.AddFilterRule(nil, []netip.Prefix{netip.MustParsePrefix("172.16.0.0/16")},
destination, fw.ProtocolALL, nil, nil, fw.ActionAccept)
require.Error(t, err, "a domain destination is not expressible without ipset")
require.ErrorContains(t, err, "requires ipset")
}
// TestIptablesNatRuleReAddKeepsSetReferences re-adds the same NAT rule the way
// a repeated network-map update does. The marking rule's set references must not
// grow, or RemoveNatRule can never drop the count to zero and the set stays in
// the kernel for the rest of the process lifetime.
func TestIptablesNatRuleReAddKeepsSetReferences(t *testing.T) {
manager, err := Create(ifaceMock, iface.DefaultMTU)
require.NoError(t, err)
require.NoError(t, manager.Init(nil))
defer func() {
require.NoError(t, manager.Close(nil))
}()
set := fw.NewDomainSet(domain.List{"example.com"})
pair := fw.RouterPair{
ID: "nat-reference-test",
Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")},
Destination: fw.Network{Set: set},
Masquerade: true,
Dynamic: true,
}
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
name := manager.family4.ipsetName(set.HashedName())
first, ok := manager.family4.ipsetCounter.Get(name)
require.True(t, ok, "the marking rule must hold a reference to its set")
require.NoError(t, manager.AddNatRule(pair), "re-add nat rule")
second, ok := manager.family4.ipsetCounter.Get(name)
require.True(t, ok, "the set must still be referenced")
require.Equal(t, first.Count, second.Count, "re-adding the same rule must not add references")
require.NoError(t, manager.RemoveNatRule(pair), "remove nat rule")
_, ok = manager.family4.ipsetCounter.Get(name)
require.False(t, ok, "removing the rule must drop the last reference")
}
// TestIptablesRouteFilterIPSetFallback covers the route ACL side of the
// fallback: with a destination set, the expanded per-source rules land
// in the route forward chain and are all removed on delete.
@@ -340,9 +406,115 @@ func TestIptablesRouteFilterIPSetFallback(t *testing.T) {
}
}
// TestIptablesCloseRemovesAllState exercises a spread of rule kinds and then
// asserts Close puts every table it touches back exactly as it found it. A
// leaked chain, jump, or ipset survives the daemon and nothing can remove it
// afterwards, since the tracking that knew about it is gone.
func TestIptablesCloseRemovesAllState(t *testing.T) {
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
require.NoError(t, err)
before := snapshotIptables(t, ipv4Client)
manager, err := Create(ifaceMock, iface.DefaultMTU)
require.NoError(t, err)
require.NoError(t, manager.Init(nil))
sources := []netip.Prefix{
netip.MustParsePrefix("10.20.0.42/32"),
netip.MustParsePrefix("10.20.0.43/32"),
}
// A multi-source peer rule: shared ipset plus the mangle redirect pairing.
_, err = manager.AddFilterRule(nil, sources, fw.Network{}, "tcp",
nil, &fw.Port{Values: []uint16{22}}, fw.ActionAccept)
require.NoError(t, err, "add peer rule")
// A route rule with a dynamic destination: a second set, in the forward chain.
_, err = manager.AddFilterRule(nil, sources,
fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})},
fw.ProtocolALL, nil, nil, fw.ActionDrop)
require.NoError(t, err, "add route rule")
// NAT marking for a routed destination, both directions.
pair := fw.RouterPair{
ID: "cleanup-test",
Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")},
Destination: fw.Network{Prefix: netip.MustParsePrefix("192.168.55.0/24")},
Masquerade: true,
}
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
require.NoError(t, manager.EnableRouting(), "enable routing")
// A DNAT redirect, which also holds a forwarding reference.
dnat := fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{8080}},
TranslatedAddress: netip.MustParseAddr("10.20.0.44"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
dnatRule, err := manager.AddDNATRule(dnat)
require.NoError(t, err, "add dnat rule")
require.NotEqual(t, before, snapshotIptables(t, ipv4Client), "the manager must have installed state")
require.NoError(t, manager.DeleteDNATRule(dnatRule), "delete dnat rule")
require.NoError(t, manager.DisableRouting(), "disable routing")
require.NoError(t, manager.Close(nil), "close")
after := snapshotIptables(t, ipv4Client)
require.Equal(t, before.chains, after.chains, "Close must remove every chain it created")
require.Equal(t, before.rules, after.rules, "Close must remove every rule it created")
require.Equal(t, before.sets, after.sets, "Close must destroy every ipset it created")
}
// iptablesState is a snapshot of the tables the manager writes to, used to
// compare the kernel before and after a manager lifetime.
type iptablesState struct {
chains map[string][]string
rules map[string][]string
sets []string
}
func snapshotIptables(t *testing.T, client *iptables.IPTables) iptablesState {
t.Helper()
state := iptablesState{
chains: map[string][]string{},
rules: map[string][]string{},
}
for _, table := range []string{tableFilter, tableNat, tableMangle, tableRaw} {
chains, err := client.ListChains(table)
require.NoErrorf(t, err, "list chains in %s", table)
slices.Sort(chains)
state.chains[table] = chains
for _, chain := range chains {
rules, err := client.List(table, chain)
require.NoErrorf(t, err, "list rules in %s/%s", table, chain)
state.rules[table+"/"+chain] = rules
}
}
sets, err := ipset.ListAll()
require.NoError(t, err, "list ipsets")
for _, set := range sets {
state.sets = append(state.sets, set.SetName)
}
slices.Sort(state.sets)
return state
}
func checkRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, chainName string, mustExists bool, rulespec ...string) {
t.Helper()
exists, err := ipv4Client.Exists("filter", chainName, rulespec...)
checkTableRuleSpecs(t, ipv4Client, tableFilter, chainName, mustExists, rulespec...)
}
func checkTableRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, table, chainName string, mustExists bool, rulespec ...string) {
t.Helper()
exists, err := ipv4Client.Exists(table, chainName, rulespec...)
require.NoError(t, err, "failed to check rule")
require.Falsef(t, !exists && mustExists, "rule '%v' does not exist", rulespec)
require.Falsef(t, exists && !mustExists, "rule '%v' exist", rulespec)

View File

@@ -192,14 +192,23 @@ func (r *family) insertEstablishedRule(chain string) error {
return nil
}
func (r *family) addNatRule(pair firewall.RouterPair) error {
func (r *family) addNatRule(pair firewall.RouterPair) (err error) {
ruleID := pair.GenKey(firewall.NatFormat)
if rule, exists := r.rules[ruleID]; exists {
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); err != nil {
return fmt.Errorf("remove existing marking rule for %s: %w", pair.Destination, err)
if derr := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); derr != nil {
return fmt.Errorf("remove existing marking rule for %s: %w", pair.Destination, derr)
}
delete(r.rules, ruleID)
// Drop the replaced spec's set references only once the new spec has
// taken its own, so a set both specs share is not destroyed and
// recreated, which would lose the prefixes UpdateSet put in it.
defer func() {
if derr := r.decrementSetCounter(rule); derr != nil && err == nil {
err = fmt.Errorf("decrement ipset counter: %w", derr)
}
}()
}
markValue := nbnet.PreroutingFwmarkMasquerade

View File

@@ -236,6 +236,12 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error {
if pr.nftRule.Handle == 0 {
log.Warnf("filter rule %s has no handle, removing stale entry", ruleID)
// The paired mangle rule can still be in the kernel with a live
// handle. Dropping the tracking entry without removing it would
// leave a prerouting rule that nothing can find again.
if err := r.deleteMangleRule(pr, ruleID); err != nil {
return err
}
r.dropNetworkMatch(pr.nftRule.Exprs)
delete(r.filters, ruleID)
return nil
@@ -244,11 +250,7 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error {
if err := r.conn.DelRule(pr.nftRule); err != nil {
log.Errorf("queue rule delete: %v", err)
}
if pr.mangleRule != nil {
if err := r.conn.DelRule(pr.mangleRule); err != nil {
log.Errorf("queue mangle rule delete: %v", err)
}
}
r.queueMangleDelete(pr)
if err := r.conn.Flush(); err != nil {
return fmt.Errorf("flush delete %s: %w", ruleID, err)
}
@@ -258,6 +260,32 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error {
return nil
}
// deleteMangleRule removes the prerouting rule paired with a filter rule on
// its own, for the paths that drop the filter rule's tracking without queueing
// a delete for it.
func (r *family) deleteMangleRule(pr *Rule, ruleID firewall.RuleID) error {
if pr.mangleRule == nil || pr.mangleRule.Handle == 0 {
return nil
}
r.queueMangleDelete(pr)
if err := r.conn.Flush(); err != nil {
return fmt.Errorf("flush mangle delete %s: %w", ruleID, err)
}
return nil
}
// queueMangleDelete queues the delete of the rule's prerouting counterpart, if
// it has one. The caller commits it.
func (r *family) queueMangleDelete(pr *Rule) {
if pr.mangleRule == nil {
return
}
if err := r.conn.DelRule(pr.mangleRule); err != nil {
log.Errorf("queue mangle rule delete: %v", err)
}
}
func (r *family) decrementSetCounter(rule *nftables.Rule) error {
if r.ipsetCounter == nil {
return nil

View File

@@ -540,6 +540,23 @@ func TestNftablesUpdateSetMergesOverlapping(t *testing.T) {
netip.MustParsePrefix("192.168.1.1/32"),
}
require.NoError(t, r.UpdateSet(set, overlapping), "UpdateSet must merge overlapping prefixes")
fetchedSet, err := r.conn.GetSetByName(r.workTable, set.HashedName())
require.NoError(t, err, "fetch updated set")
elements, err := r.conn.GetSetElements(fetchedSet)
require.NoError(t, err, "get set elements")
starts := make(map[string]bool)
for _, elem := range elements {
if elem.IntervalEnd {
continue
}
starts[netip.AddrFrom4(*(*[4]byte)(elem.Key)).String()] = true
}
// The /32s are covered by the /24, so the update adds one interval and
// leaves the one created earlier in place.
assert.Equal(t, map[string]bool{"10.0.0.0": true, "192.168.1.0": true}, starts,
"merged set must hold the original and the merged interval")
}
func TestNftablesCreateIpSet_IPv6(t *testing.T) {

View File

@@ -23,26 +23,47 @@ func (r *family) AddNatRule(pair firewall.RouterPair) error {
return fmt.Errorf(refreshRulesMapError, err)
}
// Resolve every rule's match expressions before queueing any of them: a
// message buffered on the shared connection cannot be un-queued, so
// returning an error after queueing would leave the next caller's Flush
// to commit a rule nothing tracks.
var legacyExprs []expr.Any
if r.legacyManagement {
log.Warnf("This peer is connected to a NetBird Management service with an older version. Allowing all traffic for %s", pair.Destination)
if err := r.addLegacyRouteRule(pair); err != nil {
r.rollbackRules(pair)
return fmt.Errorf("add legacy routing rule: %w", err)
var err error
legacyExprs, err = r.legacyRouteRuleExprs(pair)
if err != nil {
return fmt.Errorf("build legacy routing rule: %w", err)
}
}
inverse := firewall.GetInversePair(pair)
var natExprs, inverseExprs []expr.Any
if pair.Masquerade {
if err := r.addNatRule(pair); err != nil {
r.rollbackRules(pair)
return fmt.Errorf("add nat rule: %w", err)
var err error
natExprs, err = r.natRuleExprs(pair)
if err != nil {
r.dropNetworkMatch(legacyExprs)
return fmt.Errorf("build nat rule: %w", err)
}
if err := r.addNatRule(firewall.GetInversePair(pair)); err != nil {
r.rollbackRules(pair)
return fmt.Errorf("add inverse nat rule: %w", err)
inverseExprs, err = r.natRuleExprs(inverse)
if err != nil {
r.dropNetworkMatch(legacyExprs)
r.dropNetworkMatch(natExprs)
return fmt.Errorf("build inverse nat rule: %w", err)
}
}
if legacyExprs != nil {
r.queueLegacyRouteRule(pair, legacyExprs)
}
if pair.Masquerade {
r.queueNatRule(pair, natExprs)
r.queueNatRule(inverse, inverseExprs)
}
if err := r.conn.Flush(); err != nil {
r.rollbackRules(pair)
return fmt.Errorf("insert rules for %s: %w", pair.Destination, err)
@@ -70,17 +91,19 @@ func (r *family) rollbackRules(pair firewall.RouterPair) {
}
}
// addNatRule inserts a nftables rule to the conn client flush queue
func (r *family) addNatRule(pair firewall.RouterPair) error {
// natRuleExprs resolves the match expressions of the pair's prerouting
// marking rule. It reserves the ipset references the matches need but queues
// nothing on the connection, so its error paths leave the connection clean.
func (r *family) natRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) {
sourceExp, err := r.applyNetwork(pair.Source, nil, true)
if err != nil {
return fmt.Errorf("apply source: %w", err)
return nil, fmt.Errorf("apply source: %w", err)
}
destExp, err := r.applyNetwork(pair.Destination, nil, false)
if err != nil {
r.dropNetworkMatch(sourceExp)
return fmt.Errorf("apply destination: %w", err)
return nil, fmt.Errorf("apply destination: %w", err)
}
op := expr.CmpOpEq
@@ -123,13 +146,19 @@ func (r *family) addNatRule(pair firewall.RouterPair) error {
},
)
return exprs, nil
}
// queueNatRule replaces any tracked rule for the pair and queues the new
// prerouting marking rule on the connection. Failures are logged rather than
// returned: the caller has already queued messages that only a Flush can
// commit, so it must not return early.
func (r *family) queueNatRule(pair firewall.RouterPair, exprs []expr.Any) {
ruleID := pair.GenKey(firewall.PreroutingFormat)
if _, exists := r.rules[ruleID]; exists {
if err := r.removeNatRule(pair); err != nil {
r.dropNetworkMatch(sourceExp)
r.dropNetworkMatch(destExp)
return fmt.Errorf("remove prerouting rule: %w", err)
log.Errorf("replace prerouting rule %s: %v", ruleID, err)
}
}
@@ -141,8 +170,6 @@ func (r *family) addNatRule(pair firewall.RouterPair) error {
Exprs: exprs,
UserData: []byte(ruleID),
})
return nil
}
func (r *family) addPostroutingRules() {
@@ -308,27 +335,32 @@ func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any {
return exprs
}
func (r *family) addLegacyRouteRule(pair firewall.RouterPair) error {
// legacyRouteRuleExprs resolves the match expressions of the pair's legacy
// forwarding rule, queueing nothing on the connection.
func (r *family) legacyRouteRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) {
sourceExp, err := r.applyNetwork(pair.Source, nil, true)
if err != nil {
return fmt.Errorf("apply source: %w", err)
return nil, fmt.Errorf("apply source: %w", err)
}
destExp, err := r.applyNetwork(pair.Destination, nil, false)
if err != nil {
r.dropNetworkMatch(sourceExp)
return fmt.Errorf("apply destination: %w", err)
return nil, fmt.Errorf("apply destination: %w", err)
}
exprs := buildLegacyRouteRuleExpressions(sourceExp, destExp)
return buildLegacyRouteRuleExpressions(sourceExp, destExp), nil
}
// queueLegacyRouteRule replaces any tracked rule for the pair and queues the
// new legacy forwarding rule. Failures are logged for the same reason as in
// queueNatRule.
func (r *family) queueLegacyRouteRule(pair firewall.RouterPair, exprs []expr.Any) {
ruleID := pair.GenKey(firewall.ForwardingFormat)
if _, exists := r.rules[ruleID]; exists {
if err := r.removeLegacyRouteRule(pair); err != nil {
r.dropNetworkMatch(sourceExp)
r.dropNetworkMatch(destExp)
return fmt.Errorf("remove legacy routing rule: %w", err)
log.Errorf("replace legacy forwarding rule %s: %v", ruleID, err)
}
}
@@ -338,7 +370,6 @@ func (r *family) addLegacyRouteRule(pair firewall.RouterPair) error {
Exprs: exprs,
UserData: []byte(ruleID),
})
return nil
}
// removeLegacyRouteRule removes a legacy routing rule for mgmt servers pre route acls
@@ -395,14 +426,27 @@ func (r *family) RemoveAllLegacyRouteRules() error {
}
var merr *multierror.Error
var found bool
for k, rule := range r.rules {
if !strings.HasPrefix(string(k), firewall.ForwardingFormatPrefix) {
continue
}
found = true
if err := r.deleteLegacyRuleEntry(k, rule); err != nil {
merr = multierror.Append(merr, err)
}
}
// Commit the queued deletes here instead of leaving them for whichever
// caller flushes next: the tracking entries are already gone, so an
// uncommitted delete would leave a rule in the kernel that nothing can
// find again.
if found {
if err := r.conn.Flush(); err != nil {
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
}
}
return nberrors.FormatErrorOrNil(merr)
}

View File

@@ -101,6 +101,4 @@ func TestRouteACL_MixedFamilyZeroSourcesStayFamilySafe(t *testing.T) {
assert.True(t, pass, "v4 source must match the v4 destination rule via 0.0.0.0/0")
_, pass = m.routeACLsPass(v6Src, netip.MustParseAddr("fd00:1::5"), 255, 0, 0)
assert.True(t, pass, "v6 source must match the v6 destination rule via ::/0")
_, pass = m.routeACLsPass(v6Src, netip.MustParseAddr("10.0.0.5"), 255, 0, 0)
assert.True(t, pass, "v6 source still passes the v4 destination rule via ::/0 in the same source list")
}