Files
tomswall/internal/nftables/guards_test.go
T
unkin-agent bc65312647
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Cover per-family expansion across rules, NAT, tunnels and blrules
2026-10-09 23:48:53 +11:00

240 lines
10 KiB
Go

package nftables
import (
"os"
"os/exec"
"reflect"
"strings"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
func guardCfg(af config.AddressFamily) *config.Config {
return hostsCfg(func(c *config.Config) {
c.Settings.AddressFamily = af
c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64")
c.Hosts[0].Exclusions = []string{"192.0.2.9", "2001:db8::9"}
c.Interfaces[0].Options.NoSmurfs = true
c.Rules = append(c.Rules,
config.Rule{Action: config.RuleDrop, Source: "net:203.0.113.7", Dest: "lan:192.0.2.5"},
config.Rule{Action: config.RuleDrop, Source: "net:2001:db8:1::7", Dest: "lan:2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1,2001:db8::1"},
config.Rule{Action: config.RuleDrop, Source: "vpn:203.0.113.7", Dest: "net:!2001:db8::5"},
config.Rule{Action: config.RuleDNAT, Source: "vpn", Dest: "lan:192.0.2.10", Proto: "tcp",
DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn:192.0.2.77,2001:db8:7::7", Dest: "fw"})
c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "vpn:!192.0.2.1,2001:db8::1", Dest: "fw"}}
c.SNAT = []config.SNATRule{
{Action: config.SNATMasquerade, Dest: "wlo1", Source: "!192.0.2.0/24,2001:db8::/48"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "!192.0.2.9,2001:db8::9"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "2001:db8::/48"},
}
c.Tunnels = []config.Tunnel{{Type: "gre", Zone: "vpn", Gateways: []string{"203.0.113.50", "2001:db8:5::1"}}}
c.StaticNAT = []config.StaticNAT{
{External: "203.0.113.60", Interface: "wlo1", Internal: "192.0.2.60"},
{External: "2001:db8:6::1", Interface: "wlo1", Internal: "2001:db8::60"},
}
})
}
func TestCompile_FamilyGuardsDecodable(t *testing.T) {
addrLen := map[config.AddressFamily]uint32{config.FamilyIP: 4, config.FamilyIP6: 16}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
t.Run(string(af), func(t *testing.T) {
for chain, rules := range mustCompile(t, guardCfg(af)).Rules {
for _, r := range rules {
guards, l3 := 0, false
for i, e := range r.Exprs {
if nfprotoGuard(r.Exprs, i) != 0 {
guards++
if l3 {
t.Errorf("%s %s: guard after a network payload: %s", chain, r.Tag, describeRule(r))
}
}
if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader {
l3 = true
if n := addrLen[af]; n != 0 && (p.Len == 4 || p.Len == 16) && p.Len != n {
t.Errorf("%s %s: other-family address in %s table: %s", chain, r.Tag, af, describeRule(r))
}
}
}
if max := map[bool]int{true: 1, false: 0}[af == config.FamilyINET]; guards > max {
t.Errorf("%s %s: %d family guards in %s table: %s", chain, r.Tag, guards, af, describeRule(r))
}
}
}
})
}
}
func TestCompile_FamilyGuardsPerFamily(t *testing.T) {
state := mustCompile(t, guardCfg(config.FamilyINET))
for _, tt := range []struct {
chain, tag string
want []string
}{
{"forward", "rule:3", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9",
"iif=tun0 oif=enp2s0 ip6",
}},
{"forward", "rule:4", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64 !daddr=2001:db8::1",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9 !daddr=2001:db8::1",
"iif=tun0 oif=enp2s0 ip6 !daddr=2001:db8::1",
}},
{"forward", "rule:5", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 saddr=203.0.113.7",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 saddr=203.0.113.7",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 saddr=203.0.113.7",
}},
{"prerouting", "rule:6", []string{"iif=tun0 ip4 !daddr=203.0.113.5"}},
{"forward", "rule:6:accept", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24 !daddr=192.0.2.9 daddr=192.0.2.10"}},
{"input", "rule:7", []string{"iif=tun0 ip4 saddr=192.0.2.77", "iif=tun0 ip6 saddr=2001:db8:7::7"}},
{"input", "blrule:0", []string{"iif=tun0 ip4 !saddr=192.0.2.1", "iif=tun0 ip6 !saddr=2001:db8::1"}},
{"postrouting", "snat:0", []string{"oif=wlo1 ip4 !saddr=192.0.2.0/24", "oif=wlo1 ip6 !saddr=2001:db8::/48"}},
{"postrouting", "snat:1", []string{"oif=wlo1 ip4 !saddr=192.0.2.9"}},
{"postrouting", "snat:2", nil},
{"input", "tunnel:0", []string{"ip4 saddr=203.0.113.50", "ip6 saddr=2001:db8:5::1"}},
{"prerouting", "staticnat:dnat:0", []string{"iif=wlo1 ip4 daddr=203.0.113.60"}},
{"postrouting", "staticnat:snat:1", []string{"oif=wlo1 ip6 saddr=2001:db8::60"}},
} {
if got := describeTagged(state, tt.chain, tt.tag); !reflect.DeepEqual(got, tt.want) {
t.Errorf("%s %s = %q, want %q", tt.chain, tt.tag, got, tt.want)
}
}
}
// TestCompile_NegatedV4DropKeepsV6 checks a single-family table keeps its family's half of a
// negated DROP: "everything except 192.0.2.1" still drops all IPv6.
func TestCompile_NegatedV4DropKeepsV6(t *testing.T) {
for af, want := range map[config.AddressFamily][]string{
config.FamilyIP: {"iif=tun0 oif=wlo1 !daddr=192.0.2.0/24 !daddr=192.0.2.1", "iif=tun0 oif=wlo1 daddr=192.0.2.9 !daddr=192.0.2.1", "iif=tun0 oif=enp2s0 !daddr=198.51.100.0/24 !daddr=192.0.2.1"},
config.FamilyIP6: {"iif=tun0 oif=wlo1 !daddr=2001:db8::/64", "iif=tun0 oif=wlo1 daddr=2001:db8::9", "iif=tun0 oif=enp2s0"},
} {
if got := describeTagged(mustCompile(t, guardCfg(af)), "forward", "rule:3"); !reflect.DeepEqual(got, want) {
t.Errorf("%s rule:3 = %q, want %q", af, got, want)
}
}
}
func TestSplitAddrs(t *testing.T) {
for in, want := range map[string][]string{
"": {""},
"192.0.2.1,2001:db8::1": {"192.0.2.1", "2001:db8::1"},
"!192.0.2.1": {"!192.0.2.1", "::/0"},
"!2001:db8::1": {"0.0.0.0/0", "!2001:db8::1"},
"!192.0.2.1,2001:db8::1,198.51.100.0/24": {"!192.0.2.1,198.51.100.0/24", "!2001:db8::1"},
} {
if got := splitAddrs(in); !reflect.DeepEqual(got, want) {
t.Errorf("splitAddrs(%q) = %q, want %q", in, got, want)
}
}
}
func TestMatchGuardedCIDR(t *testing.T) {
for in, want := range map[string]string{
"192.0.2.1": "ip4 saddr=192.0.2.1",
"2001:db8::/48": "ip6 saddr=2001:db8::/48",
"!192.0.2.1,198.51.100.0/24": "ip4 !saddr=192.0.2.1 !saddr=198.51.100.0/24",
"!2001:db8::1": "ip6 !saddr=2001:db8::1",
"0.0.0.0/0": "ip4",
"::/0": "ip6",
} {
e, err := matchSourceCIDR(in)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if got := describeRule(ManagedRule{Exprs: e}); got != want {
t.Errorf("matchSourceCIDR(%q) = %q, want %q", in, got, want)
}
}
if e, _ := matchDestCIDR("!2001:db8::5"); describeRule(ManagedRule{Exprs: e}) != "ip6 !daddr=2001:db8::5" {
t.Errorf("matchDestCIDR(!2001:db8::5) = %q", describeRule(ManagedRule{Exprs: e}))
}
if _, err := matchSourceCIDR("!nonsense"); err == nil {
t.Error("invalid negated address: want error")
}
}
func TestFamilyGuards(t *testing.T) {
v4, _ := matchSourceCIDR("192.0.2.1")
v6, _ := matchDestCIDR("2001:db8::1")
both, _ := matchDestCIDR("198.51.100.1")
state := func(e ...[]expr.Any) *FirewallState {
var r []expr.Any
for _, x := range e {
r = append(r, x...)
}
return &FirewallState{Rules: map[string][]ManagedRule{"input": {{Exprs: r, Tag: "t"}}}}
}
s := state(v4, both)
if err := familyGuards(s, config.FamilyINET); err != nil || describeRule(s.Rules["input"][0]) != "ip4 saddr=192.0.2.1 daddr=198.51.100.1" {
t.Errorf("same-family guards not merged: %v %q", err, describeRule(s.Rules["input"][0]))
}
if err := familyGuards(state(v4, v6), config.FamilyINET); err == nil || !strings.Contains(err.Error(), "conflicting") {
t.Errorf("conflicting guards: want error, got %v", err)
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP); err != nil || len(s.Rules["input"]) != 0 {
t.Errorf("ip table must drop IPv6 rules: %v %v", err, s.Rules["input"])
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP6); err != nil || describeRule(s.Rules["input"][0]) != "daddr=2001:db8::1" {
t.Errorf("ip6 table must strip the guard: %v %q", err, describeRule(s.Rules["input"][0]))
}
if !famsAgree(0, unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV4) || famsAgree(unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV6) {
t.Error("famsAgree")
}
}
// TestNetnsNftListDecodes applies each family in a fresh user+net namespace and requires
// nft(8) to list the ruleset and a second plan to be empty. Needs unshare and nft.
func TestNetnsNftListDecodes(t *testing.T) {
if af := os.Getenv("TOMSWALL_NETNS_CHILD"); af != "" {
netnsChild(t, config.AddressFamily(af))
return
}
if os.Getenv("TOMSWALL_NETNS_TEST") == "" {
t.Skip("set TOMSWALL_NETNS_TEST=1 to run (needs unshare and nft)")
}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
cmd := exec.Command("unshare", "-rn", os.Args[0], "-test.run=^TestNetnsNftListDecodes$", "-test.v")
cmd.Env = append(os.Environ(), "TOMSWALL_NETNS_CHILD="+string(af))
if out, err := cmd.CombinedOutput(); err != nil {
t.Errorf("%s: %v\n%s", af, err, out)
}
}
}
func netnsChild(t *testing.T, af config.AddressFamily) {
e, err := NewEngine(guardCfg(af))
if err != nil {
t.Fatal(err)
}
cs, err := e.Plan()
if err != nil {
t.Fatal(err)
}
if err := e.Apply(cs); err != nil {
t.Fatal(err)
}
if out, err := exec.Command("nft", "list", "ruleset").CombinedOutput(); err != nil {
t.Fatalf("nft list ruleset: %v\n%s", err, out)
}
if cs, err = e.Plan(); err != nil || len(cs.Add)+len(cs.Remove) != 0 {
t.Fatalf("second plan not empty: %d add, %d remove, err %v", len(cs.Add), len(cs.Remove), err)
}
}