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) } }