diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 7e300c6..c557f31 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -420,8 +420,8 @@ func splitAddrs(addr string) []string { func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, origDest, proto string, dports, sports config.PortSpec, action config.RuleAction, logLevel string, fwZone string, section config.RuleSection) error { - srcIfaces := c.resolveZoneInterfaces(srcZone) - dstIfaces := c.resolveZoneInterfaces(dstZone) + srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr) + dstIfaces := c.resolveZoneInterfaces(dstZone, dstAddr) chain := c.selectChain(srcZone, dstZone, fwZone) // ponytail: forward daddr is post-DNAT; lift with `ct original daddr` (expr.Ct Direction, google/nftables v0.3.0). if origDest != "" && chain == "forward" { @@ -490,7 +490,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dnatPort = uint16(p) } - srcIfaces := c.resolveZoneInterfaces(srcZone) + srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr) var odExprs []expr.Any if origDest != "" { @@ -614,8 +614,8 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { } chain := c.selectChain(sz, dz, fwZone) - srcIfaces := c.resolveZoneInterfaces(sz) - dstIfaces := c.resolveZoneInterfaces(dz) + srcIfaces := c.resolveZoneInterfaces(sz, "") + dstIfaces := c.resolveZoneInterfaces(dz, "") for _, si := range srcIfaces { for _, di := range dstIfaces { @@ -624,7 +624,7 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { if si != "" { exprs = append(exprs, matchIfaceName(true, si)...) } - if di != "" && chain == "forward" { + if di != "" && chain != "input" { exprs = append(exprs, matchIfaceName(false, di)...) } @@ -957,34 +957,25 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { return "forward" } -func (c *Compiler) resolveZoneInterfaces(zone string) []string { - if zone == "all" || zone == "" { +// resolveZoneInterfaces returns nil (fail closed) for a zone with no interfaces unless a non-negated address match narrows the rule. +func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string { + if z, ok := c.cfg.Zones[zone]; !ok || z.Type == config.ZoneFirewall { return []string{""} } - ifaces := c.cfg.ZoneInterfaces(zone) - if len(ifaces) > 0 { + if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 { return ifaces } - if z, ok := c.cfg.Zones[zone]; ok && z.Type == config.ZoneIP && !c.zoneHasHosts(zone) { - if !c.warned[zone] { - if c.warned == nil { - c.warned = map[string]bool{} - } - c.warned[zone] = true - slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) - } - return nil + if addr != "" && !strings.HasPrefix(addr, "!") { + return []string{""} } - return []string{""} -} - -func (c *Compiler) zoneHasHosts(zone string) bool { - for _, h := range c.cfg.Hosts { - if h.Zone == zone { - return true + if !c.warned[zone] { + if c.warned == nil { + c.warned = map[string]bool{} } + c.warned[zone] = true + slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) } - return false + return nil } func (c *Compiler) expandZoneRef(ref string) []string { @@ -1022,7 +1013,7 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor if srcIface != "" { exprs = append(exprs, matchIfaceName(true, srcIface)...) } - if dstIface != "" && chain == "forward" { + if dstIface != "" && chain != "input" { exprs = append(exprs, matchIfaceName(false, dstIface)...) } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 890cfc6..9e14e2a 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1,8 +1,10 @@ package nftables import ( + "bytes" "encoding/binary" "fmt" + "log/slog" "net" "reflect" "strings" @@ -142,17 +144,17 @@ func TestCompiler_ResolveZoneInterfaces(t *testing.T) { } c := NewCompiler(cfg) - ifaces := c.resolveZoneInterfaces("net") + ifaces := c.resolveZoneInterfaces("net", "") if len(ifaces) != 1 || ifaces[0] != "eth0" { t.Errorf("resolveZoneInterfaces(net) = %v, want [eth0]", ifaces) } - ifaces = c.resolveZoneInterfaces("all") + ifaces = c.resolveZoneInterfaces("all", "") if len(ifaces) != 1 || ifaces[0] != "" { t.Errorf("resolveZoneInterfaces(all) = %v, want [\"\"]", ifaces) } - ifaces = c.resolveZoneInterfaces("fw") + ifaces = c.resolveZoneInterfaces("fw", "") if len(ifaces) != 1 || ifaces[0] != "" { t.Errorf("resolveZoneInterfaces(fw) = %v, want [\"\"]", ifaces) } @@ -1759,6 +1761,83 @@ func TestCompile_PortAndProtoLists(t *testing.T) { } } +func TestCompile_OutputPolicyMatchesOif(t *testing.T) { + cfg := listCfg(func(c *config.Config) { + c.Zones["lan"] = config.Zone{Type: config.ZoneIP} + c.Interfaces = append(c.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + c.Policy = []config.Policy{{Source: "fw", Dest: "lan", Action: config.PolicyAccept}} + }) + got := taggedRules(mustCompile(t, cfg), "output", "policy:0") + if len(got) != 1 || describeRule(got[0]) != "oif=eth1" { + t.Fatalf("fw->lan policy = %v, want one rule oif=eth1", got) + } +} + +func TestCompile_InterfacelessZonesFailClosed(t *testing.T) { + ipsec := func(c *config.Config) { c.Zones["ips"] = config.Zone{Type: config.ZoneIPSec} } + rule := func(dest string) func(*config.Config) { + return func(c *config.Config) { + ipsec(c) + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "fw", Dest: dest}} + } + } + tests := []struct { + name string + mod func(*config.Config) + tag string + want int + warns []string + }{ + {"ipsec zone without interface", func(c *config.Config) { + ipsec(c) + c.Policy = []config.Policy{{Source: "fw", Dest: "ips", Action: config.PolicyAccept}} + }, "policy:0", 0, []string{"ips"}}, + {"hosts-only zone", func(c *config.Config) { + c.Zones["hst"] = config.Zone{Type: config.ZoneIP} + c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}} + c.Policy = []config.Policy{{Source: "fw", Dest: "hst", Action: config.PolicyAccept}} + }, "policy:0", 0, []string{"hst"}}, + {"fw all expansion keeps zones with interfaces", func(c *config.Config) { + ipsec(c) + c.Zones["hst"] = config.Zone{Type: config.ZoneIP} + c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}} + c.Policy = []config.Policy{{Source: "fw", Dest: "all", Action: config.PolicyDrop}} + }, "policy:0", 1, []string{"hst", "ips"}}, + {"negated address does not scope", rule("ips:!192.0.2.1"), "rule:0", 0, []string{"ips"}}, + {"address scopes", rule("ips:192.0.2.1"), "rule:0", 1, []string{"ips"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var logs bytes.Buffer + prev := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + defer slog.SetDefault(prev) + + state := mustCompile(t, listCfg(tt.mod)) + got := 0 + for chain := range state.Rules { + for _, r := range taggedRules(state, chain, tt.tag) { + got++ + if describeRule(r) == "" { + t.Errorf("%s: un-scoped rule", chain) + } + } + } + if got != tt.want { + t.Errorf("got %d %s rules, want %d", got, tt.tag, tt.want) + } + if n := strings.Count(logs.String(), "zone has no interfaces"); n != len(tt.warns) { + t.Errorf("got %d warnings, want %d:\n%s", n, len(tt.warns), logs.String()) + } + for _, z := range tt.warns { + if n := strings.Count(logs.String(), "zone="+z+"\n"); n != 1 { + t.Errorf("zone %s warned %d times, want 1", z, n) + } + } + }) + } +} + // describeRule renders a rule's iif/oif/saddr/daddr matches, e.g. "iif=eth1 oif=eth2 daddr=192.0.2.1". func describeRule(r ManagedRule) string { var parts []string @@ -1804,7 +1883,7 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "fw in source list goes to output", rule: config.Rule{Action: config.RuleAccept, Source: "fw,lan", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"22"}}, - want: map[string][]string{"output": {""}, "forward": {"iif=eth1 oif=eth2"}}, + want: map[string][]string{"output": {"oif=eth2"}, "forward": {"iif=eth1 oif=eth2"}}, }, { name: "dest list with fw splits input and forward", @@ -1847,9 +1926,24 @@ func TestCompile_CommaZoneLists(t *testing.T) { want: map[string][]string{"forward": {"iif=eth1 oif=eth3", "iif=eth1 oif=eth0"}}, }, { - name: "interface-less ipsec zone keeps zone-agnostic rule, ip zone skipped", + name: "interface-less zones are skipped", rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"}, - want: map[string][]string{"forward": {"iif=eth1"}}, + want: map[string][]string{}, + }, + { + name: "interface-less zone kept when address narrows it", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn:192.0.2.1"}, + want: map[string][]string{"forward": {"iif=eth1 daddr=192.0.2.1"}}, + }, + { + name: "fw source matches dest zone oif", + rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "lan,vpn,dmz", Proto: "tcp", DPort: config.PortSpec{"22"}}, + want: map[string][]string{"output": {"oif=eth1"}}, + }, + { + name: "fw to all has no oif", + rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "all:192.0.2.1"}, + want: map[string][]string{"output": {"daddr=192.0.2.1"}}, }, { name: "dnat origdest", @@ -1871,6 +1965,11 @@ func TestCompile_CommaZoneLists(t *testing.T) { rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"}, want: map[string][]string{"input": {"iif=eth0 daddr=203.0.113.5"}}, }, + { + name: "origdest does not scope interface-less zone", + rule: config.Rule{Action: config.RuleAccept, Source: "vpn", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"}, + want: map[string][]string{}, + }, { name: "accept ipv6 origdest", rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"},