diff --git a/internal/config/blrules.go b/internal/config/blrules.go index 4118d22..1dbd02c 100644 --- a/internal/config/blrules.go +++ b/internal/config/blrules.go @@ -55,28 +55,11 @@ func (c *Config) validateBlrules() error { return fmt.Errorf("blrules[%d]: dest required", i) } - if r.Source != "all" && r.Source != "any" && r.Source != "none" && - !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - for _, zs := range SplitZoneList(r.Source) { - if _, ok := c.Zones[zs.Zone]; !ok { - return fmt.Errorf("blrules[%d]: source zone %q not defined", i, zs.Zone) - } - if !validAddrList(zs.Addr) { - return fmt.Errorf("blrules[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) - } - } + if err := c.validateZoneRef(r.Source); err != nil { + return fmt.Errorf("blrules[%d]: source %w", i, err) } - - if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && - !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - for _, zs := range SplitZoneList(r.Dest) { - if _, ok := c.Zones[zs.Zone]; !ok { - return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, zs.Zone) - } - if !validAddrList(zs.Addr) { - return fmt.Errorf("blrules[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) - } - } + if err := c.validateZoneRef(r.Dest); err != nil { + return fmt.Errorf("blrules[%d]: dest %w", i, err) } } return nil diff --git a/internal/config/conntrack.go b/internal/config/conntrack.go index 527a00f..9b5be24 100644 --- a/internal/config/conntrack.go +++ b/internal/config/conntrack.go @@ -1,6 +1,9 @@ package config -import "fmt" +import ( + "fmt" + "strings" +) type ConntrackAction string @@ -68,8 +71,14 @@ func (c *Config) validateConntrack() error { return fmt.Errorf("conntrack[%d]: helper name required for helper action", i) } - if ct.Source == "" && ct.Dest == "" && ct.Action != ConntrackHelper { - return fmt.Errorf("conntrack[%d]: source or dest required", i) + if HasZoneExclusion(ct.Source) || HasZoneExclusion(ct.Dest) { + return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i) + } + if err := c.validateZoneRef(ct.Source); err != nil { + return fmt.Errorf("conntrack[%d]: source %w", i, err) + } + if err := c.validateZoneRef(ct.Dest); err != nil { + return fmt.Errorf("conntrack[%d]: dest %w", i, err) } if ct.User != "" { @@ -84,3 +93,9 @@ func (c *Config) validateConntrack() error { } return nil } + +// HasZoneExclusion reports an all/any zone ref with a "+" or "!" modifier (all+, all!x, any+!x, ...). +func HasZoneExclusion(spec string) bool { + zones, _, _ := strings.Cut(spec, ":") + return (strings.HasPrefix(zones, "all") || strings.HasPrefix(zones, "any")) && strings.ContainsAny(zones[3:], "+!") +} diff --git a/internal/config/extras_test.go b/internal/config/extras_test.go index f83941d..d312ed9 100644 --- a/internal/config/extras_test.go +++ b/internal/config/extras_test.go @@ -53,11 +53,52 @@ func TestValidateConntrack(t *testing.T) { }, }, { - name: "source or dest required for non-helper", - rules: []ConntrackRule{ - {Action: ConntrackDrop}, - }, - wantErr: "source or dest required", + name: "omitted source and dest is valid", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Proto: "udp", DPort: PortSpec{"53"}}}, + }, + { + name: "unknown source zone", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "nte"}}, + wantErr: `source zone "nte" not defined`, + }, + { + name: "unknown dest zone", + rules: []ConntrackRule{{Action: ConntrackDrop, Source: "net", Dest: "nte:192.0.2.1"}}, + wantErr: `dest zone "nte" not defined`, + }, + { + name: "all and plain zone forms are valid", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "net,fw", Dest: "all:192.0.2.1"}}, + }, + { + name: "Source all!net rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "all!net"}}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "Dest all!net:192.0.2.1 rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "all!net:192.0.2.1"}}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "Source all+ rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "all+"}}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "Dest all+!net rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "all+!net"}}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "Source any!net rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "any!net"}}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "Dest any+ rejected", + rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "any+"}}, + wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "helper without source/dest is valid", diff --git a/internal/config/rules.go b/internal/config/rules.go index fb94e31..4099b79 100644 --- a/internal/config/rules.go +++ b/internal/config/rules.go @@ -173,29 +173,13 @@ func (c *Config) validateRules() error { return fmt.Errorf("rule[%d]: dest required", i) } - if r.Source != "all" && r.Source != "any" && r.Source != "none" && - !hasPrefix(r.Source, "all+") && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - for _, zs := range SplitZoneList(r.Source) { - if _, ok := c.Zones[zs.Zone]; !ok { - return fmt.Errorf("rule[%d]: source zone %q not defined", i, zs.Zone) - } - if !validAddrList(zs.Addr) { - return fmt.Errorf("rule[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) - } - } + if err := c.validateZoneRef(r.Source); err != nil { + return fmt.Errorf("rule[%d]: source %w", i, err) } if r.Action != RuleDNAT && r.Action != RuleRedirect && r.Action != RuleNoNAT { - if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && - !hasPrefix(r.Dest, "all+") && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - for _, zs := range SplitZoneList(r.Dest) { - if _, ok := c.Zones[zs.Zone]; !ok { - return fmt.Errorf("rule[%d]: dest zone %q not defined", i, zs.Zone) - } - if !validAddrList(zs.Addr) { - return fmt.Errorf("rule[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) - } - } + if err := c.validateZoneRef(r.Dest); err != nil { + return fmt.Errorf("rule[%d]: dest %w", i, err) } } @@ -264,6 +248,32 @@ func zoneFromSpec(spec string) string { return spec } -func hasPrefix(s, prefix string) bool { - return len(s) >= len(prefix) && s[:len(prefix)] == prefix +// validateZoneRef checks a SOURCE/DEST spec: all/any[+][!excluded,...][:addr], none, or a declared zone list. +func (c *Config) validateZoneRef(spec string) error { + zones, addr, _ := strings.Cut(spec, ":") + base, excl, isExcl := strings.Cut(zones, "!") + switch base { + case "", "none", "all", "all+", "any", "any+": + if base == "" && isExcl { + return fmt.Errorf("%q: exclusion needs all or any", spec) + } + for _, z := range strings.Split(excl, ",") { + if _, ok := c.Zones[strings.TrimSpace(z)]; isExcl && !ok { + return fmt.Errorf("excluded zone %q not defined", z) + } + } + if !validAddrList(addr) { + return fmt.Errorf("%q: '!' may only prefix the whole address list", addr) + } + return nil + } + for _, zs := range SplitZoneList(spec) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("zone %q not defined", zs.Zone) + } + if !validAddrList(zs.Addr) { + return fmt.Errorf("%q: '!' may only prefix the whole address list", zs.Addr) + } + } + return nil } diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 62543cd..a86c5b5 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -219,97 +219,138 @@ var helperProtos = map[string]string{ "Q.931": "tcp", "RAS": "udp", "sane": "tcp", "sip": "udp", "snmp": "udp", "tftp": "udp", } -// compileHelper declares a ct helper object per (helper, proto) and assigns it -// after conntrack (-200) has created the entry; a raw-priority assignment is a no-op. -func (c *Compiler) compileHelper(state *FirewallState, tag string, ct config.ConntrackRule) error { - if ct.Source != "" || ct.Dest != "" { - slog.Warn("conntrack helper source/dest not supported, assigning globally", "rule", tag, "helper", ct.Helper) - } - chains := []string{"helper_prerouting", "helper_output"} - switch ct.Chain { - case config.ConntrackPrerouting: - chains = chains[:1] - case config.ConntrackOutput: - chains = chains[1:] +func helperObjName(helper, proto string) string { + if proto != helperProtos[helper] { + return helper + "-" + proto } + return helper +} +// expandHelper declares a ct helper object per (helper, proto) and returns one entry per proto. +func expandHelper(state *FirewallState, ct config.ConntrackRule) ([]config.ConntrackRule, error) { proto := ct.Proto if proto == "" { if proto = helperProtos[ct.Helper]; proto == "" { - return fmt.Errorf("proto required for helper %q", ct.Helper) + return nil, fmt.Errorf("proto required for helper %q", ct.Helper) } } + var out []config.ConntrackRule for _, p := range strings.Split(proto, ",") { p = strings.TrimSpace(p) n, err := protoNumber(p) if err != nil { - return err - } - name := ct.Helper - if p != helperProtos[ct.Helper] { - name += "-" + p + return nil, err } + name := helperObjName(ct.Helper, p) if !slices.ContainsFunc(state.Helpers, func(h Helper) bool { return h.Name == name }) { state.Helpers = append(state.Helpers, Helper{Name: name, Helper: expr.CtHelper{Name: ct.Helper, L3Proto: unix.NFPROTO_INET, L4Proto: n}}) } + pct := ct + pct.Proto = p + out = append(out, pct) + } + return out, nil +} - matches, err := l4Matches(p, ct.DPort, ct.SPort) - if err != nil { - return err +func (c *Compiler) compileConntrack(state *FirewallState) error { + fwZone := c.cfg.FirewallZone() + for i, ct := range c.cfg.Conntrack { + tag := fmt.Sprintf("conntrack:%d", i) + if config.HasZoneExclusion(ct.Source) || config.HasZoneExclusion(ct.Dest) { + return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i) } - for _, chain := range chains { - for _, m := range matches { - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: append(append([]expr.Any{}, m.exprs...), - &expr.Objref{Type: unix.NFT_OBJECT_CT_HELPER, Name: name}), - Tag: tag + ":" + chain, - }) + cts := []config.ConntrackRule{ct} + if ct.Action == config.ConntrackHelper { + if ct.Chain == "" { + ct.Chain = config.ConntrackBoth + } + var err error + if cts, err = expandHelper(state, ct); err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + } + srcs, dsts := c.zoneSpecs(ct.Source), c.zoneSpecs(ct.Dest) + if len(srcs) == 0 { + srcs = []config.ZoneSpec{{}} + } + if len(dsts) == 0 { + dsts = []config.ZoneSpec{{}} + } + + for _, src := range srcs { + if src.Zone == "all" || src.Zone == "any" { + src.Zone = "" + } + chains := []string{"raw_prerouting"} + switch { + case ct.Chain == config.ConntrackOutput && src.Zone != fwZone && src.Zone != "": + return fmt.Errorf("conntrack[%d]: chain output needs SOURCE %s, got %q", i, fwZone, src.Zone) + case ct.Chain == config.ConntrackPrerouting && src.Zone == fwZone: + return fmt.Errorf("conntrack[%d]: SOURCE %s cannot use chain prerouting", i, fwZone) + case ct.Chain != config.ConntrackPrerouting && src.Zone == fwZone, ct.Chain == config.ConntrackOutput: + chains = []string{"raw_output"} + case ct.Chain == config.ConntrackBoth && src.Zone == "": + chains = []string{"raw_prerouting", "raw_output"} + } + for _, srcAddr := range splitAddrs(src.Addr) { + for _, dst := range dsts { + for _, dstAddr := range splitAddrs(dst.Addr) { + for _, chain := range chains { + for _, pct := range cts { + if err := c.compileConntrackPair(state, tag, chain, pct, src.Zone, srcAddr, dst.Zone, dstAddr); err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + } + } + } + } } } } return nil } -func (c *Compiler) compileConntrack(state *FirewallState) error { - for i, ct := range c.cfg.Conntrack { - tag := fmt.Sprintf("conntrack:%d", i) - if ct.Action == config.ConntrackHelper { - if err := c.compileHelper(state, tag, ct); err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) +// compileConntrackPair matches iif of the source zone in raw_prerouting and oif of the dest zone in raw_output. +// Helpers are assigned after conntrack (-200) has created the entry, so they go to the mangle-priority +// helper_* chain instead; a raw-priority assignment is a no-op. +func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule, + srcZone, srcAddr, dstZone, dstAddr string) error { + if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" && + (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 { + return nil + } + if chain == "raw_output" { + srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr) + } + out := chain + if ct.Action == config.ConntrackHelper { + out = "helper_" + strings.TrimPrefix(chain, "raw_") + } + for _, srcIface := range srcIfaces { + for _, dstIface := range dstIfaces { + matches, err := c.buildMatchExprs(srcIface, dstIface, chain, ct.Proto, ct.DPort, ct.SPort, srcAddr, dstAddr) + if err != nil { + return err } - continue - } - - chains := []string{"prerouting"} - switch ct.Chain { - case config.ConntrackOutput: - chains = []string{"output"} - case config.ConntrackBoth: - chains = []string{"prerouting", "output"} - } - - matches, err := l4Matches(ct.Proto, ct.DPort, nil) - if err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) - } - - for _, chain := range chains { for _, m := range matches { - exprs := append([]expr.Any{}, m.exprs...) - + exprs := m.exprs switch ct.Action { case config.ConntrackNoTrack: exprs = append(exprs, &expr.Notrack{}) case config.ConntrackDrop: exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) + case config.ConntrackHelper: + exprs = append(exprs, &expr.Objref{Type: unix.NFT_OBJECT_CT_HELPER, Name: helperObjName(ct.Helper, ct.Proto)}) } - - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, + state.Rules[out] = append(state.Rules[out], ManagedRule{ + Chain: out, Exprs: exprs, - Tag: tag + ":" + chain, + Tag: tag + ":" + out, }) } } @@ -340,7 +381,7 @@ func (c *Compiler) compileRules(state *FirewallState) error { if err != nil { return fmt.Errorf("rule[%d]: %w", i, err) } - if len(matches)*specCount(rule.Source, rule.Dest, rule.OrigDest, rule.Action) > 1 { + if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, rule.Action) > 1 { return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i) } } @@ -426,7 +467,7 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p dports, sports config.PortSpec, action config.RuleAction, logLevel string, dnatDest, origDest string, fwZone string, section config.RuleSection) error { - for _, src := range zoneSpecs(srcSpec) { + for _, src := range c.zoneSpecs(srcSpec) { for _, srcAddr := range splitAddrs(src.Addr) { for _, od := range splitAddrs(origDest) { if action == config.RuleDNAT || action == config.RuleRedirect { @@ -435,7 +476,12 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p } continue } - for _, dst := range zoneSpecs(dstSpec) { + for _, dst := range c.zoneSpecs(dstSpec) { + // Exclusion expansion never pairs fw with itself, and pairs a zone with itself only for "all+". + if src.Zone == dst.Zone && (isZoneExclusion(srcSpec) || isZoneExclusion(dstSpec)) && + (src.Zone == fwZone || !strings.Contains(srcSpec, "+!") && !strings.Contains(dstSpec, "+!")) { + continue + } for _, dstAddr := range splitAddrs(dst.Addr) { if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, od, proto, dports, sports, action, logLevel, fwZone, section); err != nil { @@ -450,9 +496,9 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p } // specCount is how many zone/address combinations compileOneRule expands src and dst into. -func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int { +func (c *Compiler) specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int { count := func(spec string) (n int) { - for _, z := range zoneSpecs(spec) { + for _, z := range c.zoneSpecs(spec) { n += len(splitAddrs(z.Addr)) } return n @@ -464,13 +510,26 @@ func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int return n * count(dstSpec) } -// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. -func zoneSpecs(spec string) []config.ZoneSpec { +// zoneSpecs expands a comma zone list; "all"/"any" stay global and "all!x,y" becomes every zone but x and y. +func (c *Compiler) zoneSpecs(spec string) []config.ZoneSpec { zone, addr := splitZoneSpec(spec) - if base, _, _ := strings.Cut(strings.TrimSuffix(zone, "+"), "!"); base == "all" || base == "any" { - return []config.ZoneSpec{{Zone: zone, Addr: addr}} + if !isZoneExclusion(zone) { + if base := strings.TrimSuffix(zone, "+"); base == "all" || base == "any" { + return []config.ZoneSpec{{Zone: zone, Addr: addr}} + } + return config.SplitZoneList(spec) } - return config.SplitZoneList(spec) + var out []config.ZoneSpec + for _, z := range c.expandZoneRef(zone) { + out = append(out, config.ZoneSpec{Zone: z, Addr: addr}) + } + return out +} + +func isZoneExclusion(spec string) bool { + base, _, ok := strings.Cut(spec, "!") + base = strings.TrimSuffix(base, "+") + return ok && (base == "all" || base == "any") } // splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. @@ -1023,9 +1082,18 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { return "forward" } -// resolveZoneInterfaces returns nil (fail closed) for a zone with no interfaces unless a non-negated address match narrows the rule. +// 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 { - if z, ok := c.cfg.Zones[zone]; !ok || z.Type == config.ZoneFirewall { + switch zone { + case "", "all", "all+", "any", "any+": + return []string{""} + } + z, ok := c.cfg.Zones[zone] + if !ok { + slog.Warn("compiler: unknown zone, skipping its rules", "zone", zone) + return nil + } + if z.Type == config.ZoneFirewall { return []string{""} } if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 { @@ -1059,7 +1127,7 @@ func (c *Compiler) expandZoneRef(ref string) []string { } } - if base == "all" || base == "all+" { + if base == "all" || base == "all+" || base == "any" || base == "any+" { var zones []string for name := range c.cfg.Zones { if excluded != nil && excluded[name] { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index f970417..7be75de 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -583,7 +583,7 @@ func TestCompile_ConntrackNoTrack(t *testing.T) { { Action: config.ConntrackNoTrack, Source: "net", - Dest: "fw", + Dest: "fw:192.0.2.1", Proto: "udp", DPort: config.PortSpec{"53"}, }, @@ -597,14 +597,17 @@ func TestCompile_ConntrackNoTrack(t *testing.T) { } found := false - for _, r := range state.Rules["prerouting"] { - if r.Tag == "conntrack:0:prerouting" { + for _, r := range state.Rules["raw_prerouting"] { + if r.Tag == "conntrack:0:raw_prerouting" { found = true break } } if !found { - t.Error("no notrack rule found in prerouting chain") + t.Error("no notrack rule found in raw_prerouting chain") + } + if len(state.Rules["prerouting"]) != 0 { + t.Error("conntrack rule leaked into the nat prerouting chain") } } @@ -1872,6 +1875,9 @@ func describeRule(r ManagedRule) string { } parts = append(parts, name+"="+net.IP(cmp.Data).String()) } + if m.Base == expr.PayloadBaseTransportHeader && m.Offset == 0 && m.Len == 2 && cmp.Op == expr.CmpOpEq { + parts = append(parts, fmt.Sprintf("sport=%d", binary.BigEndian.Uint16(cmp.Data))) + } } } return strings.Join(parts, " ") @@ -2181,7 +2187,7 @@ func TestCompile_ListExpansionCounts(t *testing.T) { }, "postrouting", "snat:0", 4}, {"conntrack dport list", func(c *config.Config) { c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53", "123"}}} - }, "prerouting", "conntrack:0:prerouting", 2}, + }, "raw_prerouting", "conntrack:0:raw_prerouting", 2}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -2357,7 +2363,7 @@ func TestCompile_ColonRanges(t *testing.T) { }, "postrouting", "snat:0", 2}, {"conntrack dport", func(c *config.Config) { c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"1024:2048"}}} - }, "prerouting", "conntrack:0:prerouting", 2}, + }, "raw_prerouting", "conntrack:0:raw_prerouting", 2}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -2477,3 +2483,319 @@ func TestCompile_ConntrackHelper(t *testing.T) { t.Error("expected error for unknown helper without proto") } } + +func TestCompile_ConntrackZones(t *testing.T) { + tests := []struct { + name string + ct config.ConntrackRule + want map[string][]string + wantErr string + }{ + { + name: "source zone matches iif", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", DPort: config.PortSpec{"53"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0"}}, + }, + { + name: "source and dest addresses", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:192.0.2.1,198.51.100.1", Dest: "fw:203.0.113.1"}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 saddr=198.51.100.1 daddr=203.0.113.1"}}, + }, + { + name: "fw source goes to raw_output with dest oif", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net,lan"}, + want: map[string][]string{"raw_output": {"oif=eth0", "oif=eth1"}}, + }, + { + name: "all matches no interface", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Dest: "fw:192.0.2.53"}, + want: map[string][]string{"raw_prerouting": {"daddr=192.0.2.53"}}, + }, + { + name: "interface-less zone fails closed", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "dmz"}, + want: map[string][]string{}, + }, + { + name: "chain output with fw source matches dest oif", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net", Chain: config.ConntrackOutput}, + want: map[string][]string{"raw_output": {"oif=eth0"}}, + }, + { + name: "chain output with non-fw source is rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw", Chain: config.ConntrackOutput}, + wantErr: "chain output needs SOURCE fw", + }, + { + name: "chain prerouting with fw source is rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net", Chain: config.ConntrackPrerouting}, + wantErr: "SOURCE fw cannot use chain prerouting", + }, + { + name: "omitted source and dest is global in both chains", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, + want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, + }, + { + name: "all source with chain both is global in both chains", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, + want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, + }, + { + name: "all source with chain output is global in raw_output", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackOutput}, + want: map[string][]string{"raw_output": {""}}, + }, + { + name: "any source with chain both is global in both chains", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, + want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, + }, + { + name: "any source with chain output is global in raw_output", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackOutput}, + want: map[string][]string{"raw_output": {""}}, + }, + { + name: "chain both with non-fw source emits prerouting only", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, + want: map[string][]string{"raw_prerouting": {"iif=eth0"}}, + }, + { + name: "chain both with fw source emits output only", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "lan", Chain: config.ConntrackBoth}, + want: map[string][]string{"raw_output": {"oif=eth1"}}, + }, + { + name: "negated addresses stay one AND-ed match", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 !saddr=192.0.2.1 !saddr=198.51.100.1"}}, + }, + { + name: "sport", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", SPort: config.PortSpec{"123"}}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 sport=123"}}, + }, + { + name: "dest zone without address is rejected in prerouting", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan"}, + wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, + }, + { + name: "fw dest zone without address is rejected in prerouting", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "fw"}, + wantErr: `conntrack DEST zone "fw" needs an address in prerouting`, + }, + { + name: "fw dest zone with address matches daddr", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw:192.0.2.1"}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=192.0.2.1"}}, + }, + { + name: "unknown zone fails closed", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "nte"}, + want: map[string][]string{}, + }, + { + name: "unknown dest zone fails closed in prerouting", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "typo"}, + want: map[string][]string{}, + }, + { + name: "none dest zone yields no rule", + ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "none"}, + want: map[string][]string{}, + }, + { + name: "all!net rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "all!net rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "all!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "all+ rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all+"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "all+!net rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "all+!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "any!net rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "any+ rejected", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "any+"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + { + name: "omitted source and dest is global", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53"}}, + want: map[string][]string{"raw_prerouting": {""}}, + }, + { + name: "dest zone with address matches daddr", + ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"}, + want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Zones["dmz"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Conntrack = []config.ConntrackRule{tt.ct} + }) + if tt.wantErr != "" { + if _, err := NewCompiler(cfg).Compile(); err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("Compile() error = %v, want %q", err, tt.wantErr) + } + return + } + state := mustCompile(t, cfg) + got := map[string][]string{} + for _, chain := range []string{"raw_prerouting", "raw_output"} { + for _, r := range taggedRules(state, chain, "conntrack:0:"+chain) { + got[chain] = append(got[chain], describeRule(r)) + } + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %v, want %v", got, tt.want) + } + }) + } +} + +func TestCompile_RuleZoneExclusionExpands(t *testing.T) { + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Rules = []config.Rule{ + {Source: "all!net", Dest: "fw", Action: config.RuleAccept, Proto: "tcp", DPort: config.PortSpec{"22"}}, + {Source: "nte", Dest: "fw", Action: config.RuleAccept}, + } + }) + state := mustCompile(t, cfg) + var got []string + for _, r := range taggedRules(state, "input", "rule:0") { + got = append(got, describeRule(r)) + } + if want := []string{"iif=eth1"}; !reflect.DeepEqual(got, want) { + t.Errorf("all!net -> fw input rules = %v, want %v", got, want) + } + if r := taggedRules(state, "input", "rule:1"); len(r) != 0 { + t.Errorf("unknown zone compiled %d rules, want 0", len(r)) + } +} + +func TestCompile_RuleZoneExclusionIntraZoneSymmetric(t *testing.T) { + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Rules = []config.Rule{ + {Source: "lan", Dest: "all+!net", Action: config.RuleAccept}, + {Source: "lan", Dest: "all!net", Action: config.RuleAccept}, + } + }) + state := mustCompile(t, cfg) + for i, want := range []bool{true, false} { + got := false + for _, r := range taggedRules(state, "forward", fmt.Sprintf("rule:%d", i)) { + got = got || describeRule(r) == "iif=eth1 oif=eth1" + } + if got != want { + t.Errorf("rule:%d lan->lan forward = %v, want %v", i, got, want) + } + } +} + +func TestCompile_ConntrackHelperZones(t *testing.T) { + tests := []struct { + name string + ct config.ConntrackRule + want map[string][]string + wantErr string + }{ + { + name: "non-fw source is prerouting only", + ct: config.ConntrackRule{Source: "net", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {"iif=eth0"}}, + }, + { + name: "fw source is output only with dest oif", + ct: config.ConntrackRule{Source: "fw", Dest: "lan", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_output": {"oif=eth1"}}, + }, + { + name: "any source is global in both chains", + ct: config.ConntrackRule{Source: "any", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {""}, "helper_output": {""}}, + }, + { + name: "dest zone with address in prerouting", + ct: config.ConntrackRule{Source: "net", Dest: "lan:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"21"}}, + want: map[string][]string{"helper_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, + }, + { + name: "dest zone without address is rejected in prerouting", + ct: config.ConntrackRule{Source: "net", Dest: "lan"}, + wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, + }, + { + name: "chain output with non-fw source is rejected", + ct: config.ConntrackRule{Source: "net", Chain: config.ConntrackOutput}, + wantErr: "chain output needs SOURCE fw", + }, + { + name: "interface-less zone fails closed", + ct: config.ConntrackRule{Source: "dmz"}, + want: map[string][]string{}, + }, + { + name: "unknown zone fails closed", + ct: config.ConntrackRule{Source: "nte"}, + want: map[string][]string{}, + }, + { + name: "exclusion rejected", + ct: config.ConntrackRule{Source: "all!net"}, + wantErr: "zone exclusions are not supported in conntrack entries", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.ct.Action, tt.ct.Helper = config.ConntrackHelper, "ftp" + cfg := listCfg(func(cfg *config.Config) { + cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} + cfg.Zones["dmz"] = config.Zone{Type: config.ZoneIP} + cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) + cfg.Conntrack = []config.ConntrackRule{tt.ct} + }) + if tt.wantErr != "" { + if _, err := NewCompiler(cfg).Compile(); err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("Compile() error = %v, want %q", err, tt.wantErr) + } + return + } + state := mustCompile(t, cfg) + got := map[string][]string{} + for _, chain := range []string{"helper_prerouting", "helper_output"} { + for _, r := range taggedRules(state, chain, "conntrack:0:"+chain) { + got[chain] = append(got[chain], describeRule(r)) + } + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index b783e6a..c6781c6 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -84,6 +84,22 @@ func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftable Hooknum: nftables.ChainHookOutput, Priority: nftables.ChainPriorityMangle, }, + "raw_prerouting": { + Name: "raw_prerouting", + Table: table, + Type: nftables.ChainTypeFilter, + Hooknum: nftables.ChainHookPrerouting, + Priority: nftables.ChainPriorityRaw, + Policy: policyPtr(nftables.ChainPolicyAccept), + }, + "raw_output": { + Name: "raw_output", + Table: table, + Type: nftables.ChainTypeFilter, + Hooknum: nftables.ChainHookOutput, + Priority: nftables.ChainPriorityRaw, + Policy: policyPtr(nftables.ChainPolicyAccept), + }, } for name, chain := range chains { diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go index 3098119..babe6ba 100644 --- a/internal/nftables/snapshot_test.go +++ b/internal/nftables/snapshot_test.go @@ -243,3 +243,17 @@ func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, e } return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn} } + +func TestEnsureChainsRawPriority(t *testing.T) { + e := testEngine(t, nil) + chains := e.ensureChains(e.ensureTable(), nil) + for name, hook := range map[string]*nftables.ChainHook{"raw_prerouting": nftables.ChainHookPrerouting, "raw_output": nftables.ChainHookOutput} { + c, ok := chains[name] + if !ok { + t.Fatalf("%s chain not declared", name) + } + if *c.Priority != *nftables.ChainPriorityRaw || *c.Hooknum != *hook || c.Type != nftables.ChainTypeFilter || *c.Policy != nftables.ChainPolicyAccept { + t.Errorf("%s: got type %s hook %d prio %d", name, c.Type, *c.Hooknum, *c.Priority) + } + } +} diff --git a/tomswall.example.yaml b/tomswall.example.yaml index 452f0d1..b48acf4 100644 --- a/tomswall.example.yaml +++ b/tomswall.example.yaml @@ -191,13 +191,12 @@ snat: # conntrack: # - action: notrack # source: net -# dest: fw +# dest: fw:203.0.113.1 # proto: udp # dport: [53] # comment: "Skip conntrack for DNS" # - action: helper # source: loc -# dest: net # proto: tcp # dport: [21] # helper: ftp