From 15b67423c332dfa3c322192d63dee0fb38a48839 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 8 Jul 2026 17:05:34 +0800 Subject: [PATCH] Add flow PortWithSelectorRange --- flow.go | 5 +++++ flow_nat.go | 29 ++++++++++++++++++++++------- 2 files changed, 27 insertions(+), 7 deletions(-) diff --git a/flow.go b/flow.go index 4fbf801..0954601 100644 --- a/flow.go +++ b/flow.go @@ -73,6 +73,11 @@ type Port interface { WritePackets(packets [][]byte) error } +type PortWithSelectorRange interface { + Port + PortSelectorRange() (start uint16, count uint16) +} + type Return interface { ReturnHeadroom() int ReturnPackets(packets [][]byte) [][]byte diff --git a/flow_nat.go b/flow_nat.go index fe4ea1a..4d2bca3 100644 --- a/flow_nat.go +++ b/flow_nat.go @@ -15,10 +15,12 @@ const ( ) type portNAT struct { - port Port - hasher maphash.Hasher[flowKey] - shardMask uint32 - shards []natShard + port Port + hasher maphash.Hasher[flowKey] + shardMask uint32 + shards []natShard + selectorStart uint16 + selectorCount uint16 counter uint32 pending [][]byte @@ -40,6 +42,9 @@ func newPortNAT(port Port) *portNAT { shardMask: uint32(shardCount - 1), shards: make([]natShard, shardCount), } + if rangedPort, isRanged := port.(PortWithSelectorRange); isRanged { + nat.selectorStart, nat.selectorCount = rangedPort.PortSelectorRange() + } for i := range nat.shards { nat.shards[i].flows = make(map[flowKey]*forwardFlow) } @@ -87,16 +92,26 @@ func (n *portNAT) reverseKeyFor(protocol uint8, portAddress, serverAddress netip } } +func (n *portNAT) selectorRange(protocol uint8) (uint16, uint32) { + if n.selectorCount == 0 || + protocol == uint8(header.ICMPv4ProtocolNumber) || protocol == uint8(header.ICMPv6ProtocolNumber) { + return natSelectorMin, natSelectorMax - natSelectorMin + 1 + } + return n.selectorStart, uint32(n.selectorCount) +} + func (n *portNAT) allocateSelector(protocol uint8, portAddress, serverAddress netip.Addr, serverPort, clientSelector uint16) (uint16, flowKey, bool) { - if clientSelector != 0 { + rangeStart, rangeCount := n.selectorRange(protocol) + if clientSelector != 0 && + clientSelector >= rangeStart && uint32(clientSelector-rangeStart) < rangeCount { key := n.reverseKeyFor(protocol, portAddress, serverAddress, serverPort, clientSelector) if n.lookup(key) == nil { return clientSelector, key, true } } - for range natSelectorMax - natSelectorMin + 1 { + for range rangeCount { n.counter++ - candidate := uint16(natSelectorMin + n.counter%(natSelectorMax-natSelectorMin+1)) + candidate := rangeStart + uint16(n.counter%rangeCount) key := n.reverseKeyFor(protocol, portAddress, serverAddress, serverPort, candidate) if n.lookup(key) == nil { return candidate, key, true