diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 94e7ba9..aa27b32 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -531,6 +531,13 @@ func TestValidateHosts(t *testing.T) { }, wantErr: "interface required", }, + { + name: "invalid exclusion", + zones: map[string]Zone{"fw": {Type: ZoneFirewall}, "net": {Type: ZoneIP}, "loc": {Type: ZoneIP}}, + interfaces: []Interface{{Zone: "net", Interface: "eth0"}}, + hosts: []Host{{Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}, Exclusions: []string{"192.0.2.0/24!192.0.2.7"}}}, + wantErr: "invalid address", + }, } for _, tt := range tests { diff --git a/internal/config/hosts.go b/internal/config/hosts.go index f9847dc..37affa4 100644 --- a/internal/config/hosts.go +++ b/internal/config/hosts.go @@ -1,6 +1,10 @@ package config -import "fmt" +import ( + "fmt" + "net/netip" + "slices" +) type Host struct { Zone string `yaml:"zone"` @@ -52,6 +56,13 @@ func (c *Config) validateHosts() error { if !h.Dynamic && len(h.Addresses) == 0 { return fmt.Errorf("host[%d]: at least one address required (or set dynamic: true)", i) } + for _, a := range slices.Concat(h.Addresses, h.Exclusions) { + if _, err := netip.ParsePrefix(a); err != nil { + if _, err := netip.ParseAddr(a); err != nil { + return fmt.Errorf("host[%d]: invalid address %q", i, a) + } + } + } } return nil } diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 9a28ee6..53d5b55 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -5,6 +5,7 @@ import ( "fmt" "log/slog" "net" + "net/netip" "slices" "sort" "strconv" @@ -345,12 +346,12 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, (dstAddr == "" || strings.HasPrefix(dstAddr, "!")) { return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone) } - srcIfaces, dstIfaces := c.resolveZoneInterfaces(srcZone, srcAddr), []string{""} - if chain == "raw_prerouting" && c.resolveZoneInterfaces(dstZone, dstAddr) == nil { + srcIfaces, dstIfaces := c.resolveZone(srcZone, srcAddr), []zoneMatch{{}} + if chain == "raw_prerouting" && c.resolveZone(dstZone, dstAddr) == nil { return nil } if chain == "raw_output" { - srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr) + srcIfaces, dstIfaces = []zoneMatch{{}}, c.resolveZone(dstZone, dstAddr) } out := chain if ct.Action == config.ConntrackHelper { @@ -674,8 +675,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, srcAddr) - dstIfaces := c.resolveZoneInterfaces(dstZone, dstAddr) + srcIfaces := c.resolveZone(srcZone, srcAddr) + dstIfaces := c.resolveZone(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" { @@ -744,7 +745,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dnatPort = uint16(p) } - srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr) + srcIfaces := c.resolveZone(srcZone, srcAddr) var odExprs []expr.Any if origDest != "" { @@ -765,12 +766,12 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, } for _, srcIface := range srcIfaces { + zm, err := zoneMatchExprs(srcIface, true) + if err != nil { + return err + } for _, m := range matches { - var exprs []expr.Any - - if srcIface != "" { - exprs = append(exprs, matchIfaceName(true, srcIface)...) - } + exprs := slices.Clone(zm) if srcAddr != "" { src, err := matchSourceCIDR(srcAddr) @@ -875,21 +876,17 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { } chain := c.selectChain(sz, dz, fwZone) - srcIfaces := c.resolveZoneInterfaces(sz, "") - dstIfaces := c.resolveZoneInterfaces(dz, "") + srcIfaces := c.resolveZone(sz, "") + dstIfaces := c.resolveZone(dz, "") for _, si := range srcIfaces { for _, di := range dstIfaces { - if sz == dz && si != "" && si == di { + if sz == dz && intraZoneSkip(si, di) { continue } - var exprs []expr.Any - - if si != "" { - exprs = append(exprs, matchIfaceName(true, si)...) - } - if di != "" && chain != "input" { - exprs = append(exprs, matchIfaceName(false, di)...) + exprs, err := zonePairExprs(si, di, chain) + if err != nil { + return fmt.Errorf("policy[%d]: %w", i, err) } if pol.RateLimit != "" { @@ -917,12 +914,11 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { } } - c.compileImplicitIntraZone(state, overridden) - return nil + return c.compileImplicitIntraZone(state, overridden) } // compileImplicitIntraZone accepts traffic between different interfaces of one zone, shorewall's implicit intra-zone ACCEPT policy. -func (c *Compiler) compileImplicitIntraZone(state *FirewallState, overridden map[string]bool) { +func (c *Compiler) compileImplicitIntraZone(state *FirewallState, overridden map[string]bool) error { fwZone := c.cfg.FirewallZone() zones := make([]string, 0, len(c.cfg.Zones)) for z := range c.cfg.Zones { @@ -933,21 +929,49 @@ func (c *Compiler) compileImplicitIntraZone(state *FirewallState, overridden map if z == fwZone || overridden[z] { continue } - ifaces := c.cfg.ZoneInterfaces(z) - for _, si := range ifaces { - for _, di := range ifaces { - if si == di { + if len(c.cfg.ZoneInterfaces(z)) == 0 && !slices.ContainsFunc(c.cfg.Hosts, func(h config.Host) bool { return h.Zone == z }) { + continue + } + matches := c.resolveZone(z, "") + for _, si := range matches { + for _, di := range matches { + if intraZoneSkip(si, di) { continue } + exprs, err := zonePairExprs(si, di, "forward") + if err != nil { + return fmt.Errorf("zone %s: %w", z, err) + } state.Rules["forward"] = append(state.Rules["forward"], ManagedRule{ Chain: "forward", - Exprs: append(append(matchIfaceName(true, si), matchIfaceName(false, di)...), - &expr.Verdict{Kind: expr.VerdictAccept}), - Tag: "intra:" + z, + Exprs: append(exprs, &expr.Verdict{Kind: expr.VerdictAccept}), + Tag: "intra:" + z, }) } } } + return nil +} + +// intraZoneSkip drops intra-zone pairs on one interface unless both are routeback hosts entries +// (interface routeback is compileIntraZone's job), and pairs whose address families can never both match. +func intraZoneSkip(si, di zoneMatch) bool { + if si.iface != "" && si.iface == di.iface && !(si.routeback && di.routeback) { + return true + } + a, b := matchFamily(si), matchFamily(di) + return a != 0 && b != 0 && a != b +} + +// matchFamily is the NFPROTO a zoneMatch is guarded by: its host address's family, else fam. +func matchFamily(m zoneMatch) byte { + if p, err := parsePrefix(m.addr); err == nil { + if p.Addr().Is4() { + return unix.NFPROTO_IPV4 + } + return unix.NFPROTO_IPV6 + } + return m.fam } func (c *Compiler) compileSNAT(state *FirewallState) error { @@ -1251,11 +1275,23 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { return "forward" } -// resolveZoneInterfaces returns nil (fail closed) for an unknown zone, or one with no interfaces unless a non-negated address match narrows the rule. -func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string { +// zoneMatch classifies a packet into a zone: an interface (empty: any) and, for a hosts entry, one host +// address; excl carves out hosts exclusions and the hosts of sub-zones, which shorewall matches first. +// fam (an NFPROTO; addr's family when set, else 0: any) guards addr and excl so IPv4 offsets are +// never compared against IPv6 bytes. routeback marks a hosts entry with the routeback option. +type zoneMatch struct { + iface, addr string + excl []string + fam byte + routeback bool +} + +// resolveZone returns nil (fail closed) for an unknown zone, or one with neither interfaces nor hosts +// unless a non-negated address match narrows the rule. +func (c *Compiler) resolveZone(zone, addr string) []zoneMatch { switch zone { case "", "all", "all+", "any", "any+": - return []string{""} + return []zoneMatch{{}} } z, ok := c.cfg.Zones[zone] if !ok { @@ -1263,24 +1299,151 @@ func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string { return nil } if z.Type == config.ZoneFirewall { - return []string{""} + return []zoneMatch{{}} } - if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 { - return ifaces + var out []zoneMatch + hasHosts := false + for _, iface := range c.cfg.ZoneInterfaces(zone) { + sub, back := c.subZoneHosts(zone, iface) + if len(sub) == 0 { + out = append(out, zoneMatch{iface: iface}) + } else { + v4, v6 := splitFamily(sub) + out = append(out, zoneMatch{iface: iface, excl: v4, fam: unix.NFPROTO_IPV4}, + zoneMatch{iface: iface, excl: v6, fam: unix.NFPROTO_IPV6}) + } + for _, b := range back { + if addrsOverlap(b, addr) { + out = append(out, zoneMatch{iface: iface, addr: b}) + } + } + } + for _, h := range c.cfg.Hosts { + if h.Zone != zone { + continue + } + hasHosts = true + sub, _ := c.subZoneHosts(zone, h.Interface) + v4, v6 := splitFamily(slices.Concat(h.Exclusions, sub)) + for _, a := range h.Addresses { + if addrsOverlap(a, addr) { + m := zoneMatch{iface: h.Interface, addr: a, excl: v6, routeback: h.Options.RouteBack} + if p, err := parsePrefix(a); err == nil && p.Addr().Is4() { + m.excl = v4 + } + out = append(out, m) + } + } + } + if len(out) > 0 || hasHosts { + return out } if addr != "" && !strings.HasPrefix(addr, "!") { - return []string{""} + return []zoneMatch{{}} } 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) + slog.Warn("compiler: zone has no interfaces or hosts, skipping its rules", "zone", zone) } return nil } +// subZoneHosts lists the host addresses on iface that belong to sub-zones of zone, and those +// sub-zone hosts' exclusions, which fall back to zone. +func (c *Compiler) subZoneHosts(zone, iface string) (sub, back []string) { + for _, h := range c.cfg.Hosts { + if h.Interface == iface && c.cfg.IsSubZone(h.Zone, zone) { + sub = append(sub, h.Addresses...) + back = append(back, h.Exclusions...) + } + } + return sub, back +} + +// addrsOverlap reports whether a host address can match a rule address; unparsable or negated rule addresses keep the host. +func addrsOverlap(host, rule string) bool { + if rule == "" || strings.HasPrefix(rule, "!") { + return true + } + h, err := parsePrefix(host) + if err != nil { + return true + } + for _, r := range strings.Split(rule, ",") { + if p, err := parsePrefix(r); err != nil || p.Overlaps(h) { + return true + } + } + return false +} + +// splitFamily partitions addresses by family; unparsable ones go to v6 so zoneMatchExprs still rejects them. +func splitFamily(addrs []string) (v4, v6 []string) { + for _, a := range addrs { + if p, err := parsePrefix(a); err == nil && p.Addr().Is4() { + v4 = append(v4, a) + } else { + v6 = append(v6, a) + } + } + return v4, v6 +} + +func parsePrefix(s string) (netip.Prefix, error) { + if a, err := netip.ParseAddr(s); err == nil { + return netip.PrefixFrom(a, a.BitLen()), nil + } + return netip.ParsePrefix(s) +} + +// zoneMatchExprs matches a zone on the in (src) or out interface plus its host address and exclusions, +// all guarded by m.fam so an IPv4 address never matches IPv6 bytes in an inet table. +func zoneMatchExprs(m zoneMatch, src bool) ([]expr.Any, error) { + var out []expr.Any + if m.iface != "" { + out = matchIfaceName(src, m.iface) + } + var addr []expr.Any + if m.addr != "" { + p, err := parsePrefix(m.addr) + if err != nil { + return nil, fmt.Errorf("invalid host address %q", m.addr) + } + if addr, err = matchAddrCIDR(m.addr, src); err != nil { + return nil, err + } + m.fam = unix.NFPROTO_IPV6 + if p.Addr().Is4() { + m.fam = unix.NFPROTO_IPV4 + } + } + if m.fam != 0 { + out = append(out, matchNFProto(m.fam)...) + } + out = append(out, addr...) + if len(m.excl) > 0 { + e, err := matchAddrCIDR("!"+strings.Join(m.excl, ","), src) + if err != nil { + return nil, err + } + out = append(out, e...) + } + return out, nil +} + +// zonePairExprs matches the source zone inbound and, outside input, the dest zone outbound. +func zonePairExprs(src, dst zoneMatch, chain string) ([]expr.Any, error) { + out, err := zoneMatchExprs(src, true) + if err != nil || chain == "input" { + return out, err + } + d, err := zoneMatchExprs(dst, false) + return append(out, d...), err +} + func (c *Compiler) expandZoneRef(ref string) []string { base := ref var excluded map[string]bool @@ -1310,14 +1473,10 @@ func (c *Compiler) expandZoneRef(ref string) []string { return []string{base} } -func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { - var exprs []expr.Any - - if srcIface != "" { - exprs = append(exprs, matchIfaceName(true, srcIface)...) - } - if dstIface != "" && chain != "input" { - exprs = append(exprs, matchIfaceName(false, dstIface)...) +func (c *Compiler) buildMatchExprs(srcIface, dstIface zoneMatch, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { + exprs, err := zonePairExprs(srcIface, dstIface, chain) + if err != nil { + return nil, err } if srcAddr != "" { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 4371c1f..dccbc3a 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -7,6 +7,7 @@ import ( "log/slog" "net" "reflect" + "slices" "strings" "testing" @@ -144,19 +145,17 @@ func TestCompiler_ResolveZoneInterfaces(t *testing.T) { } c := NewCompiler(cfg) - ifaces := c.resolveZoneInterfaces("net", "") - if len(ifaces) != 1 || ifaces[0] != "eth0" { - t.Errorf("resolveZoneInterfaces(net) = %v, want [eth0]", ifaces) - } - - ifaces = c.resolveZoneInterfaces("all", "") - if len(ifaces) != 1 || ifaces[0] != "" { - t.Errorf("resolveZoneInterfaces(all) = %v, want [\"\"]", ifaces) - } - - ifaces = c.resolveZoneInterfaces("fw", "") - if len(ifaces) != 1 || ifaces[0] != "" { - t.Errorf("resolveZoneInterfaces(fw) = %v, want [\"\"]", ifaces) + for _, tt := range []struct { + zone string + want []zoneMatch + }{ + {"net", []zoneMatch{{iface: "eth0"}}}, + {"all", []zoneMatch{{}}}, + {"fw", []zoneMatch{{}}}, + } { + if got := c.resolveZone(tt.zone, ""); !reflect.DeepEqual(got, tt.want) { + t.Errorf("resolveZone(%s) = %v, want %v", tt.zone, got, tt.want) + } } } @@ -1849,13 +1848,13 @@ func TestCompile_InterfacelessZonesFailClosed(t *testing.T) { 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"}}, + }, "policy:0", 1, nil}, {"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"}}, + }, "policy:0", 2, []string{"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"}}, } @@ -1895,6 +1894,19 @@ func TestCompile_InterfacelessZonesFailClosed(t *testing.T) { func describeRule(r ManagedRule) string { var parts []string for i, e := range r.Exprs { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader && i+2 < len(r.Exprs) { + bw, okb := r.Exprs[i+1].(*expr.Bitwise) + cmp, okc := r.Exprs[i+2].(*expr.Cmp) + if okb && okc { + name := map[uint32]string{12: "saddr", 16: "daddr", 8: "saddr", 24: "daddr"}[p.Offset] + if cmp.Op == expr.CmpOpNeq { + name = "!" + name + } + ones, _ := net.IPMask(bw.Mask).Size() + parts = append(parts, fmt.Sprintf("%s=%s/%d", name, net.IP(cmp.Data), ones)) + continue + } + } cmp, ok := func() (*expr.Cmp, bool) { if i+1 >= len(r.Exprs) { return nil, false @@ -1912,6 +1924,8 @@ func describeRule(r ManagedRule) string { parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00")) case expr.MetaKeyOIFNAME: parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) + case expr.MetaKeyNFPROTO: + parts = append(parts, map[byte]string{unix.NFPROTO_IPV4: "ip4", unix.NFPROTO_IPV6: "ip6"}[cmp.Data[0]]) } case *expr.Payload: if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) { @@ -2050,25 +2064,25 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "dnat origdest", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"}, - want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5"}, + want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { name: "dnat origdest list", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5,203.0.113.6"}, - want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5", "iif=eth0 daddr=203.0.113.6"}, + want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5", "iif=eth0 ip4 daddr=203.0.113.6"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { name: "dnat negated origdest list", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,203.0.113.6"}, - want: map[string][]string{"prerouting": {"iif=eth0 !daddr=203.0.113.5 !daddr=203.0.113.6"}, + want: map[string][]string{"prerouting": {"iif=eth0 ip4 !daddr=203.0.113.5 !daddr=203.0.113.6"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { name: "accept origdest", 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"}}, + want: map[string][]string{"input": {"iif=eth0 ip4 daddr=203.0.113.5"}}, }, { name: "origdest does not scope interface-less zone", @@ -2078,7 +2092,7 @@ func TestCompile_CommaZoneLists(t *testing.T) { { name: "accept ipv6 origdest", rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"}, - want: map[string][]string{"input": {"iif=eth0 daddr=2001:db8::5"}}, + want: map[string][]string{"input": {"iif=eth0 ip6 daddr=2001:db8::5"}}, }, { name: "blrule zone list", @@ -3321,6 +3335,132 @@ func TestCompile_LogLimitSplitsAroundExtrasAndNAT(t *testing.T) { } } +// hostsCfg models a shorewall setup where lan:net is defined by hosts on net's interfaces. +func hostsCfg(mod func(*config.Config)) *config.Config { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{ + "fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, + "lan": {Type: config.ZoneIP, Parents: []string{"net"}}, "vpn": {Type: config.ZoneIP}, + }, + Interfaces: []config.Interface{ + {Zone: "net", Interface: "wlo1"}, {Zone: "net", Interface: "enp2s0"}, {Zone: "vpn", Interface: "tun0"}, + }, + Hosts: []config.Host{ + {Zone: "lan", Interface: "wlo1", Addresses: []string{"192.0.2.0/24"}}, + {Zone: "lan", Interface: "enp2s0", Addresses: []string{"198.51.100.0/24"}}, + }, + Policy: []config.Policy{ + {Source: "net", Dest: "all", Action: config.PolicyDrop}, + {Source: "lan", Dest: "fw", Action: config.PolicyReject}, + {Source: "all", Dest: "all", Action: config.PolicyReject}, + }, + Rules: []config.Rule{ + {Action: config.RuleAccept, Source: "lan,vpn", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"6768"}}, + }, + PortGroups: make(map[string]config.PortGroup), + } + if mod != nil { + mod(cfg) + } + return cfg +} + +func describeTagged(state *FirewallState, chain, tag string) []string { + var out []string + for _, r := range taggedRules(state, chain, tag) { + out = append(out, describeRule(r)) + } + return out +} + +func TestCompile_HostsZoneRules(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, hostsCfg(nil)) + want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24", "iif=tun0"} + if got := describeTagged(state, "input", "rule:0"); !reflect.DeepEqual(got, want) { + t.Errorf("input rule:0 = %q, want %q", got, want) + } + if strings.Contains(logs.String(), "skipping") { + t.Errorf("unexpected warning:\n%s", logs.String()) + } +} + +func TestCompile_HostsSubZoneBeforeParent(t *testing.T) { + state := mustCompile(t, hostsCfg(nil)) + want := []string{ + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", + } + if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) { + t.Errorf("net->fw policy = %q, want %q", got, want) + } + want = []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"} + if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) { + t.Errorf("lan->fw policy = %q, want %q", got, want) + } + for _, r := range taggedRules(state, "input", "policy:1") { + if !slices.ContainsFunc(r.Exprs, func(e expr.Any) bool { + m, ok := e.(*expr.Meta) + return ok && m.Key == expr.MetaKeyNFPROTO + }) { + t.Errorf("host match lacks an nfproto guard: %v", describeRule(r)) + } + } +} + +func TestCompile_HostsZoneMatches(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + want []string + }{ + {"forward to hosts zone", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "vpn", Dest: "lan", Proto: "tcp"}} + }, "forward", "rule:0", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24", "iif=tun0 oif=enp2s0 ip4 daddr=198.51.100.0/24"}}, + {"rule address prunes non-overlapping hosts", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "lan:192.0.2.5", Dest: "fw", Proto: "tcp"}} + }, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 saddr=192.0.2.5"}}, + {"host exclusions and sub-zone exclusions fall back to the parent", func(c *config.Config) { + c.Hosts[1].Exclusions = []string{"198.51.100.7"} + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp"}} + }, "input", "rule:0", []string{ + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", + "iif=enp2s0 ip4 saddr=198.51.100.7", + }}, + {"host exclusion", func(c *config.Config) { + c.Hosts[1].Exclusions = []string{"198.51.100.7"} + }, "input", "policy:1", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24 !saddr=198.51.100.7"}}, + {"DNAT from hosts zone", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "lan", Dest: "vpn:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"80"}}} + }, "prerouting", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, + {"conntrack from hosts zone", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Source: "lan", Proto: "udp"}} + }, "raw_prerouting", "conntrack:0:raw_prerouting", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, + {"blrule from hosts zone", func(c *config.Config) { + c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "lan", Dest: "all"}} + }, "input", "blrule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, + {"all expansion includes hosts zone", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "all!net,vpn", Dest: "fw", Proto: "tcp"}} + }, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state := mustCompile(t, hostsCfg(tt.mod)) + 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) + } + }) + } +} + func TestCompile_IntraZoneMultiInterface(t *testing.T) { tests := []struct { name string @@ -3389,6 +3529,45 @@ func TestCompile_IntraZoneMultiInterface(t *testing.T) { } } +func TestCompile_HostsAddressMatchesFamilyGuarded(t *testing.T) { + state := mustCompile(t, hostsCfg(func(c *config.Config) { + c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64") + c.Hosts[0].Exclusions = []string{"192.0.2.9"} + })) + want := []string{ + "iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6 !saddr=2001:db8::/64", + "iif=wlo1 ip4 saddr=192.0.2.9", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6", + } + if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) { + t.Errorf("net->fw DROP policy = %q, want %q", got, want) + } + want = []string{ + "iif=wlo1 ip4 saddr=192.0.2.0/24 !saddr=192.0.2.9", "iif=wlo1 ip6 saddr=2001:db8::/64", + "iif=enp2s0 ip4 saddr=198.51.100.0/24", + } + if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) { + t.Errorf("lan->fw policy = %q, want %q", got, want) + } + for chain, rules := range state.Rules { + for _, r := range rules { + var fam byte + for i, e := range r.Exprs { + if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyNFPROTO { + fam = r.Exprs[i+1].(*expr.Cmp).Data[0] + } + p, ok := e.(*expr.Payload) + if !ok || p.Base != expr.PayloadBaseNetworkHeader { + continue + } + if p.Len == 4 && fam != unix.NFPROTO_IPV4 || p.Len == 16 && fam != unix.NFPROTO_IPV6 { + t.Errorf("%s %s: %d-byte address compare without its family guard: %s", chain, r.Tag, p.Len, describeRule(r)) + } + } + } + } +} + func TestCompile_FirewallSelfPolicySkipped(t *testing.T) { for _, action := range []config.PolicyAction{config.PolicyAccept, config.PolicyDrop} { t.Run(string(action), func(t *testing.T) { @@ -3424,3 +3603,63 @@ func TestCompile_DestPlusOverridesIntraZone(t *testing.T) { t.Error("lxd all+ emitted no forward rules") } } + +func TestCompile_HostsIntraZone(t *testing.T) { + state := mustCompile(t, hostsCfg(nil)) + want := []string{ + "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24", + "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24", + } + if got := describeTagged(state, "forward", "intra:lan"); !reflect.DeepEqual(got, want) { + t.Errorf("lan intra = %q, want %q", got, want) + } + want = []string{ + "iif=wlo1 ip4 !saddr=192.0.2.0/24 oif=enp2s0 ip4 !daddr=198.51.100.0/24", + "iif=wlo1 ip6 oif=enp2s0 ip6", + "iif=enp2s0 ip4 !saddr=198.51.100.0/24 oif=wlo1 ip4 !daddr=192.0.2.0/24", + "iif=enp2s0 ip6 oif=wlo1 ip6", + } + if got := describeTagged(state, "forward", "intra:net"); !reflect.DeepEqual(got, want) { + t.Errorf("net intra = %q, want %q", got, want) + } + + state = mustCompile(t, hostsCfg(func(c *config.Config) { + c.Policy = append([]config.Policy{{Source: "lan", Dest: "lan", Action: config.PolicyDrop}}, c.Policy...) + })) + if got := describeTagged(state, "forward", "intra:lan"); len(got) != 0 { + t.Errorf("explicit lan lan policy must replace implicit accept, got %q", got) + } + want = []string{ + "iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24", + "iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24", + } + if got := describeTagged(state, "forward", "policy:0"); !reflect.DeepEqual(got, want) { + t.Errorf("lan lan policy = %q, want %q", got, want) + } +} + +func TestCompile_HostsRouteBack(t *testing.T) { + sameIface := func(routeback bool) func(*config.Config) { + return func(c *config.Config) { + c.Hosts = []config.Host{ + {Zone: "lan", Interface: "wlo1", Addresses: []string{"192.0.2.0/24"}, Options: config.HostOptions{RouteBack: routeback}}, + {Zone: "lan", Interface: "wlo1", Addresses: []string{"198.51.100.0/24"}}, + } + } + } + if got := describeTagged(mustCompile(t, hostsCfg(sameIface(false))), "forward", "intra:lan"); len(got) != 0 { + t.Errorf("same-interface hosts without routeback = %q, want none", got) + } + want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24"} + if got := describeTagged(mustCompile(t, hostsCfg(sameIface(true))), "forward", "intra:lan"); !reflect.DeepEqual(got, want) { + t.Errorf("routeback hosts intra = %q, want %q", got, want) + } + + state := mustCompile(t, hostsCfg(func(c *config.Config) { + sameIface(true)(c) + c.Policy = append([]config.Policy{{Source: "lan", Dest: "lan", Action: config.PolicyDrop}}, c.Policy...) + })) + if got := describeTagged(state, "forward", "policy:0"); !reflect.DeepEqual(got, want) { + t.Errorf("routeback hosts lan lan policy = %q, want %q", got, want) + } +} diff --git a/internal/shorewall/convert.go b/internal/shorewall/convert.go index e9bc9ea..e64b8ed 100644 --- a/internal/shorewall/convert.go +++ b/internal/shorewall/convert.go @@ -248,6 +248,9 @@ func convertInterfaces(dir string, cfg *config.Config, params map[string]string) for _, row := range rows { zone := subst(field(row, 0), params) iface := subst(field(row, 1), params) + if isDash(zone) { + zone = "" + } intf := config.Interface{ Zone: zone, @@ -392,12 +395,14 @@ func convertHosts(dir string, cfg *config.Config, params map[string]string) erro zone := subst(field(row, 0), params) hostDef := subst(field(row, 1), params) + hostDef, excl, _ := strings.Cut(hostDef, "!") iface, addrs := splitHostDef(hostDef) host := config.Host{ - Zone: zone, - Interface: iface, - Addresses: addrs, + Zone: zone, + Interface: iface, + Addresses: addrs, + Exclusions: splitAddrList(excl), } optsStr := subst(field(row, 2), params) @@ -415,16 +420,19 @@ func splitHostDef(s string) (string, []string) { if idx < 0 { return s, nil } - iface := s[:idx] - addrPart := s[idx+1:] + return s[:idx], splitAddrList(s[idx+1:]) +} + +// splitAddrList splits a comma address list, unwrapping shorewall6 [addr]/len brackets. +func splitAddrList(s string) []string { var addrs []string - for _, a := range strings.Split(addrPart, ",") { - a = strings.TrimSpace(a) + for _, a := range strings.Split(s, ",") { + a = strings.NewReplacer("[", "", "]", "").Replace(strings.TrimSpace(a)) if a != "" { addrs = append(addrs, a) } } - return iface, addrs + return addrs } func parseHostOptions(s string) config.HostOptions { diff --git a/internal/shorewall/convert_test.go b/internal/shorewall/convert_test.go index 5431b5a..193111e 100644 --- a/internal/shorewall/convert_test.go +++ b/internal/shorewall/convert_test.go @@ -3,6 +3,7 @@ package shorewall import ( "os" "path/filepath" + "reflect" "testing" "git.unkin.net/unkin/tomswall/internal/config" @@ -776,3 +777,30 @@ func TestConvert_LogLimit(t *testing.T) { } } } + +func TestConvert_HostsExclusions(t *testing.T) { + dir := minimalShorewallDir(t) + writeFile(t, dir, "interfaces", ` +net eth0 +- eth1 +`) + writeFile(t, dir, "hosts", ` +loc eth0:192.0.2.0/24,198.51.100.0/24!192.0.2.7,192.0.2.8 routeback +loc eth1:[2001:db8::]/64 +`) + cfg, err := Convert(dir) + if err != nil { + t.Fatalf("Convert: %v", err) + } + want := []config.Host{ + {Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24", "198.51.100.0/24"}, + Exclusions: []string{"192.0.2.7", "192.0.2.8"}, Options: config.HostOptions{RouteBack: true}}, + {Zone: "loc", Interface: "eth1", Addresses: []string{"2001:db8::/64"}}, + } + if !reflect.DeepEqual(cfg.Hosts, want) { + t.Errorf("hosts = %+v, want %+v", cfg.Hosts, want) + } + if err := cfg.Validate(); err != nil { + t.Errorf("Validate: %v", err) + } +}