From 695869c80b4cc5af179688c80575fcde4b2c0fe1 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Fri, 9 Oct 2026 22:48:12 +1100 Subject: [PATCH] Match zones defined by hosts entries --- internal/config/config_test.go | 7 ++ internal/config/hosts.go | 13 ++- internal/nftables/compiler.go | 175 +++++++++++++++++++++++------ internal/nftables/compiler_test.go | 168 ++++++++++++++++++++++++--- internal/shorewall/convert.go | 24 ++-- internal/shorewall/convert_test.go | 28 +++++ 6 files changed, 355 insertions(+), 60 deletions(-) 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 47244d4..77d2e81 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) @@ -870,18 +871,14 @@ 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 { - 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 != "" { @@ -1213,11 +1210,19 @@ 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. +type zoneMatch struct { + iface, addr string + excl []string +} + +// 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 { @@ -1225,24 +1230,126 @@ 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) + out = append(out, zoneMatch{iface: iface, excl: sub}) + 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) + for _, a := range h.Addresses { + if addrsOverlap(a, addr) { + out = append(out, zoneMatch{iface: h.Interface, addr: a, excl: slices.Concat(h.Exclusions, sub)}) + } + } + } + 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 +} + +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, guarded by +// the address family so an IPv4 host never matches IPv6 bytes in an inet table. +// ponytail: exclusions are unguarded, so an IPv6 packet whose bytes hit an IPv4 exclusion skips the zone. +func zoneMatchExprs(m zoneMatch, src bool) ([]expr.Any, error) { + var out []expr.Any + if m.iface != "" { + out = matchIfaceName(src, m.iface) + } + if m.addr != "" { + p, err := parsePrefix(m.addr) + if err != nil { + return nil, fmt.Errorf("invalid host address %q", m.addr) + } + e, err := matchAddrCIDR(m.addr, src) + if err != nil { + return nil, err + } + fam := byte(unix.NFPROTO_IPV6) + if p.Addr().Is4() { + fam = unix.NFPROTO_IPV4 + } + out = append(append(out, matchNFProto(fam)...), e...) + } + 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 @@ -1272,14 +1379,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 5d70611..108ca25 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 @@ -3320,3 +3332,129 @@ 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 saddr=192.0.2.0/24", "iif=enp2s0 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 !saddr=192.0.2.0/24", + "iif=enp2s0 !saddr=198.51.100.0/24", + } + 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 saddr=192.0.2.0/24", "iif=enp2s0 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 daddr=192.0.2.0/24", "iif=tun0 oif=enp2s0 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 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 !saddr=192.0.2.0/24", + "iif=enp2s0 !saddr=198.51.100.0/24", + "iif=enp2s0 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 saddr=192.0.2.0/24", "iif=enp2s0 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 saddr=192.0.2.0/24", "iif=enp2s0 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 saddr=192.0.2.0/24", "iif=enp2s0 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 saddr=192.0.2.0/24", "iif=enp2s0 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 saddr=192.0.2.0/24", "iif=enp2s0 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) + } + }) + } +} 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) + } +}