diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 416ca06..52ad4eb 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -205,7 +205,7 @@ func (c *Compiler) compileBlrules(state *FirewallState) error { } if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, rule.Proto, rule.DPort, rule.SPort, - action, rule.Log, "", fwZone, ""); err != nil { + action, rule.Log, "", "", fwZone, ""); err != nil { return fmt.Errorf("blrule[%d]: %w", i, err) } } @@ -276,14 +276,14 @@ 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.Action) > 1 { + if len(matches)*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) } } if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, proto, dports, sport, - rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil { + rule.Action, rule.Log, rule.Dest, rule.OrigDest, fwZone, rule.Section); err != nil { return fmt.Errorf("rule[%d]: %w", i, err) } @@ -360,22 +360,24 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, proto string, dports, sports config.PortSpec, action config.RuleAction, logLevel string, - dnatDest string, fwZone string, section config.RuleSection) error { + dnatDest, origDest string, fwZone string, section config.RuleSection) error { for _, src := range zoneSpecs(srcSpec) { for _, srcAddr := range splitAddrs(src.Addr) { - if action == config.RuleDNAT || action == config.RuleRedirect { - if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, dstSpec, proto, dports, action, logLevel); err != nil { - return err - } - continue - } - for _, dst := range zoneSpecs(dstSpec) { - for _, dstAddr := range splitAddrs(dst.Addr) { - if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, proto, - dports, sports, action, logLevel, fwZone, section); err != nil { + for _, od := range splitAddrs(origDest) { + if action == config.RuleDNAT || action == config.RuleRedirect { + if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, od, dstSpec, proto, dports, action, logLevel); err != nil { return err } + continue + } + for _, dst := range zoneSpecs(dstSpec) { + 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 { + return err + } + } } } } @@ -384,17 +386,18 @@ 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 string, action config.RuleAction) int { +func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int { count := func(spec string) (n int) { for _, z := range zoneSpecs(spec) { n += len(splitAddrs(z.Addr)) } return n } + n := count(srcSpec) * len(splitAddrs(origDest)) if action == config.RuleDNAT || action == config.RuleRedirect { - return count(srcSpec) + return n } - return count(srcSpec) * count(dstSpec) + return n * count(dstSpec) } // zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. @@ -414,7 +417,7 @@ func splitAddrs(addr string) []string { return strings.Split(addr, ",") } -func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, proto 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) @@ -427,6 +430,16 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, if err != nil { return err } + // ponytail: matches daddr, so a forward rule never sees a pre-DNAT address; needs `ct original daddr` (google/nftables Ct.Direction, >v0.2.0). + if origDest != "" { + od, err := matchOrigDest(origDest) + if err != nil { + return fmt.Errorf("origdest: %w", err) + } + for i := range matches { + matches[i].exprs = append(matches[i].exprs, od...) + } + } for _, m := range matches { exprs := m.exprs @@ -455,7 +468,7 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, return nil } -func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dstSpec, proto string, +func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, origDest, dstSpec, proto string, dports config.PortSpec, action config.RuleAction, logLevel string) error { chain := "prerouting" @@ -476,6 +489,14 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, srcIfaces := c.resolveZoneInterfaces(srcZone) + var odExprs []expr.Any + if origDest != "" { + var err error + if odExprs, err = matchOrigDest(origDest); err != nil { + return fmt.Errorf("origdest: %w", err) + } + } + matches, err := l4Matches(proto, dports, nil) if err != nil { return err @@ -501,6 +522,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, } exprs = append(exprs, src...) } + exprs = append(exprs, odExprs...) exprs = append(exprs, m.exprs...) @@ -1386,6 +1408,24 @@ func matchDestCIDR(cidr string) ([]expr.Any, error) { return matchAddrCIDR(cidr, false) } +// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table. +func matchOrigDest(addr string) ([]expr.Any, error) { + first, _, _ := strings.Cut(strings.TrimPrefix(addr, "!"), ",") + first, _, _ = strings.Cut(first, "/") + proto := byte(unix.NFPROTO_IPV6) + if ip := net.ParseIP(first); ip != nil && ip.To4() != nil { + proto = unix.NFPROTO_IPV4 + } + dst, err := matchDestCIDR(addr) + if err != nil { + return nil, err + } + return append([]expr.Any{ + &expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1}, + &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}}, + }, dst...), nil +} + func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) { negated := false if strings.HasPrefix(cidr, "!") { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 623b4c3..973917f 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1850,6 +1850,26 @@ func TestCompile_CommaZoneLists(t *testing.T) { rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"}, want: map[string][]string{"forward": {"iif=eth1"}}, }, + { + 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"}}, + }, + { + 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"}}, + }, + { + 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"}}, + }, + { + 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"}}, + }, { name: "blrule zone list", blrule: &config.BlruleRule{Action: config.BlruleDrop, Source: "net,anycast", Dest: "fw"}, @@ -2255,3 +2275,19 @@ func TestCompile_ColonRanges(t *testing.T) { }) } } + +func TestMatchOrigDest_FamilyGuard(t *testing.T) { + for addr, want := range map[string]byte{ + "203.0.113.5": unix.NFPROTO_IPV4, + "!203.0.113.0/24,192.0.2.1": unix.NFPROTO_IPV4, + "2001:db8::5": unix.NFPROTO_IPV6, + } { + e, err := matchOrigDest(addr) + if err != nil { + t.Fatalf("%s: %v", addr, err) + } + if m, ok := e[0].(*expr.Meta); !ok || m.Key != expr.MetaKeyNFPROTO || e[1].(*expr.Cmp).Data[0] != want { + t.Errorf("%s: missing nfproto %d guard: %v", addr, want, e[:2]) + } + } +}