diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 416ca06..7e300c6 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,12 +417,16 @@ 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) dstIfaces := c.resolveZoneInterfaces(dstZone) 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" { + return fmt.Errorf("origdest: ORIGDEST on forwarded rules is not supported yet") + } for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { @@ -427,6 +434,15 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, if err != nil { return err } + 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 +471,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 +492,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 +525,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 +1411,30 @@ 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) { + var proto byte + for i, a := range strings.Split(strings.TrimPrefix(addr, "!"), ",") { + a, _, _ = strings.Cut(a, "/") + p := byte(unix.NFPROTO_IPV6) + if ip := net.ParseIP(a); ip != nil && ip.To4() != nil { + p = unix.NFPROTO_IPV4 + } + if i > 0 && p != proto { + return nil, fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr) + } + proto = p + } + 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..890cfc6 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3,6 +3,7 @@ package nftables import ( "encoding/binary" "fmt" + "net" "reflect" "strings" "testing" @@ -1781,12 +1782,12 @@ func describeRule(r ManagedRule) string { parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) } case *expr.Payload: - if m.Base == expr.PayloadBaseNetworkHeader && m.Len == 4 { - name := map[uint32]string{12: "saddr", 16: "daddr"}[m.Offset] + if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) { + name := map[uint32]string{12: "saddr", 16: "daddr", 8: "saddr", 24: "daddr"}[m.Offset] if cmp.Op == expr.CmpOpNeq { name = "!" + name } - parts = append(parts, fmt.Sprintf("%s=%d.%d.%d.%d", name, cmp.Data[0], cmp.Data[1], cmp.Data[2], cmp.Data[3])) + parts = append(parts, name+"="+net.IP(cmp.Data).String()) } } } @@ -1850,6 +1851,31 @@ 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: "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"}}, + }, { name: "blrule zone list", blrule: &config.BlruleRule{Action: config.BlruleDrop, Source: "net,anycast", Dest: "fw"}, @@ -2255,3 +2281,36 @@ 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]) + } + } + if _, err := matchOrigDest("!203.0.113.5,2001:db8::5"); err == nil { + t.Error("mixed IPv4/IPv6 origdest: want error") + } +} + +func TestCompile_OrigDestForwardRejected(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}, "svr": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "svr", Interface: "eth2"}}, + Rules: []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"}}, + PortGroups: make(map[string]config.PortGroup), + } + _, err := NewCompiler(cfg).Compile() + if err == nil || !strings.Contains(err.Error(), "not supported yet") { + t.Fatalf("Compile() error = %v, want forwarded ORIGDEST rejection", err) + } +}