diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 900eecb..af9b5d6 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -858,16 +858,21 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, func (c *Compiler) compilePolicies(state *FirewallState) error { fwZone := c.cfg.FirewallZone() + overridden := map[string]bool{} for i, pol := range c.cfg.Policy { tag := fmt.Sprintf("policy:%d", i) + explicitIntra := pol.Source == pol.Dest && !isGlobalZone(pol.Source) srcZones := c.expandZoneRef(pol.Source) dstZones := c.expandZoneRef(pol.Dest) for _, sz := range srcZones { for _, dz := range dstZones { - if sz == dz && !strings.HasSuffix(pol.Source, "+") { - continue + if sz == dz { + if sz == fwZone || (!explicitIntra && !strings.HasSuffix(pol.Source, "+") && !strings.HasSuffix(pol.Dest, "+")) { + continue + } + overridden[sz] = true } chain := c.selectChain(sz, dz, fwZone) @@ -876,6 +881,9 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { for _, si := range srcIfaces { for _, di := range dstIfaces { + if sz == dz && intraZoneSkip(si, di) { + continue + } exprs, err := zonePairExprs(si, di, chain) if err != nil { return fmt.Errorf("policy[%d]: %w", i, err) @@ -906,9 +914,66 @@ func (c *Compiler) compilePolicies(state *FirewallState) error { } } + 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) error { + fwZone := c.cfg.FirewallZone() + zones := make([]string, 0, len(c.cfg.Zones)) + for z := range c.cfg.Zones { + zones = append(zones, z) + } + sort.Strings(zones) + for _, z := range zones { + if z == fwZone || overridden[z] { + continue + } + 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(exprs, &expr.Verdict{Kind: expr.VerdictAccept}), + Tag: "intra:" + z, + }) + } + } + } return nil } +// intraZoneSkip drops intra-zone pairs on one interface (routeback's job) and pairs whose address +// families can never both match. +func intraZoneSkip(si, di zoneMatch) bool { + if si.iface != "" && si.iface == di.iface { + 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 { for i, snat := range c.cfg.SNAT { tag := fmt.Sprintf("snat:%d", i) diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index dad2976..519f1ea 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3461,6 +3461,74 @@ func TestCompile_HostsZoneMatches(t *testing.T) { } } +func TestCompile_IntraZoneMultiInterface(t *testing.T) { + tests := []struct { + name string + policy []config.Policy + tag string + want []string + }{ + { + name: "implicit accept between distinct interfaces", + policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, + tag: "intra:lxd", + want: []string{"iif=lxdbr0 oif=docker0", "iif=lxdbr0 oif=br-", "iif=docker0 oif=lxdbr0", "iif=docker0 oif=br-", "iif=br- oif=lxdbr0", "iif=br- oif=docker0"}, + }, + { + name: "explicit zone policy overrides", + policy: []config.Policy{{Source: "lxd", Dest: "lxd", Action: config.PolicyDrop, Log: "info"}, {Source: "all", Dest: "all", Action: config.PolicyDrop}}, + tag: "policy:0", + want: []string{"iif=lxdbr0 oif=docker0", "iif=lxdbr0 oif=br-", "iif=docker0 oif=lxdbr0", "iif=docker0 oif=br-", "iif=br- oif=lxdbr0", "iif=br- oif=docker0"}, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{ + "fw": {Type: config.ZoneFirewall}, + "net": {Type: config.ZoneIP}, + "lxd": {Type: config.ZoneIP}, + }, + Interfaces: []config.Interface{ + {Zone: "net", Interface: "eth0"}, + {Zone: "lxd", Interface: "lxdbr0"}, + {Zone: "lxd", Interface: "docker0"}, + {Zone: "lxd", Interface: "br-+"}, + }, + Policy: tc.policy, + PortGroups: map[string]config.PortGroup{}, + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + var got, all []string + for _, r := range state.Rules["forward"] { + if r.Tag == tc.tag { + got = append(got, describeRule(r)) + } + if strings.HasPrefix(r.Tag, "intra:") { + all = append(all, r.Tag) + } + } + if !reflect.DeepEqual(got, tc.want) { + t.Errorf("%s rules = %q, want %q", tc.tag, got, tc.want) + } + if tc.tag != "intra:lxd" && len(all) != 0 { + t.Errorf("explicit policy must replace implicit accept, got %q", all) + } + last := taggedRules(state, "forward", tc.tag) + if len(last) == 0 { + return + } + if v, ok := last[0].Exprs[len(last[0].Exprs)-1].(*expr.Verdict); !ok || (tc.tag == "intra:lxd") != (v.Kind == expr.VerdictAccept) { + t.Errorf("%s verdict = %#v", tc.tag, last[0].Exprs[len(last[0].Exprs)-1]) + } + }) + } +} + 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") @@ -3499,3 +3567,73 @@ func TestCompile_HostsAddressMatchesFamilyGuarded(t *testing.T) { } } } + +func TestCompile_FirewallSelfPolicySkipped(t *testing.T) { + for _, action := range []config.PolicyAction{config.PolicyAccept, config.PolicyDrop} { + t.Run(string(action), func(t *testing.T) { + state := mustCompile(t, listCfg(func(c *config.Config) { + c.Policy = []config.Policy{ + {Source: "fw", Dest: "fw", Action: action}, + {Source: "net", Dest: "fw", Action: config.PolicyDrop, Log: "info"}, + } + })) + for _, chain := range []string{"input", "output", "forward"} { + if got := taggedRules(state, chain, "policy:0"); len(got) != 0 { + t.Errorf("fw->fw emitted %d rules in %s", len(got), chain) + } + } + got := taggedRules(state, "input", "policy:1") + if len(got) != 1 || describeRule(got[0]) != "iif=eth0" { + t.Errorf("net->fw input rules = %d, want one scoped to eth0", len(got)) + } + }) + } +} + +func TestCompile_DestPlusOverridesIntraZone(t *testing.T) { + state := mustCompile(t, listCfg(func(c *config.Config) { + c.Zones["lxd"] = config.Zone{Type: config.ZoneIP} + c.Interfaces = append(c.Interfaces, config.Interface{Zone: "lxd", Interface: "lxdbr0"}, config.Interface{Zone: "lxd", Interface: "docker0"}) + c.Policy = []config.Policy{{Source: "lxd", Dest: "all+", Action: config.PolicyDrop}} + })) + if got := taggedRules(state, "forward", "intra:lxd"); len(got) != 0 { + t.Errorf("lxd all+ must override implicit intra-zone accept, got %d rules", len(got)) + } + if got := taggedRules(state, "forward", "policy:0"); len(got) == 0 { + 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) + } +}