Minor fixes

This commit is contained in:
世界 2026-07-07 15:36:21 +08:00
parent 14c8f75f7a
commit c17af6ee8c
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
13 changed files with 357 additions and 181 deletions

View file

@ -106,6 +106,7 @@ type ForwardDispatcher struct {
lastSweep int64
ports map[Port]*portNAT
natList atomic.Pointer[[]*portNAT]
revNAT atomic.Pointer[map[netip.Addr]*portNAT]
activeNATs []*portNAT
writebackBatch [][]byte
@ -113,6 +114,14 @@ type ForwardDispatcher struct {
segmentBuffers [][]byte
segmentSizes []int
segmentUsed int
}
func addrToTCPIP(addr netip.Addr) tcpip.Address {
if addr.Is4() {
return tcpip.AddrFrom4(addr.As4())
}
return tcpip.AddrFrom16(addr.As16())
}
func NewForwardDispatcher(handler Handler, writeback ForwardWriteback, logger logger.Logger, udpTimeout time.Duration, icmpTimeout time.Duration) *ForwardDispatcher {
@ -206,10 +215,16 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
d.forwardToPort(flow, packet, raw)
return true
case ActionAccept:
entry.deadline = now + int64(entry.idle)
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
d.removeEntry(key, entry, FlowCloseReset)
if packet.protocol == uint8(header.TCPProtocolNumber) {
if packet.tcpFlags&header.TCPFlagRst != 0 {
d.removeEntry(key, entry, FlowCloseReset)
return false
}
if packet.tcpFlags&header.TCPFlagSyn == 0 {
entry.idle = tcpEstablishedTimeout
}
}
entry.deadline = now + int64(entry.idle)
return false
case ActionReject:
entry.deadline = now + int64(entry.idle)
@ -330,24 +345,24 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic
dnatPort: serverPort != clientDestinationPort && !isICMP,
}
flow.forwardRule = rewriteRule{
sourceAddress: tcpip.AddrFromSlice(portAddress.AsSlice()),
sourceAddress: addrToTCPIP(portAddress),
sourcePort: selector,
rewriteSourcePort: true,
}
if flow.dnatAddress {
flow.forwardRule.destinationAddress = tcpip.AddrFromSlice(serverAddress.AsSlice())
flow.forwardRule.destinationAddress = addrToTCPIP(serverAddress)
}
if flow.dnatPort {
flow.forwardRule.destinationPort = serverPort
flow.forwardRule.rewriteDestinationPort = true
}
flow.reverseRule = rewriteRule{
destinationAddress: tcpip.AddrFromSlice(flow.clientAddress.AsSlice()),
destinationAddress: addrToTCPIP(flow.clientAddress),
destinationPort: flow.clientSelector,
rewriteDestinationPort: true,
}
if flow.dnatAddress {
flow.reverseRule.sourceAddress = tcpip.AddrFromSlice(clientDestinationAddress.AsSlice())
flow.reverseRule.sourceAddress = addrToTCPIP(clientDestinationAddress)
}
if flow.dnatPort {
flow.reverseRule.sourcePort = clientDestinationPort
@ -371,7 +386,6 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
err := port.AttachReturn(&d.returnPath)
if err != nil {
d.logger.Trace(E.Cause(err, "attach return path"))
d.ports[port] = nil
return nil
}
nat = newPortNAT(port)
@ -383,6 +397,20 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
}
natList = append(natList, nat)
d.natList.Store(&natList)
revMap := make(map[netip.Addr]*portNAT)
if currentRev := d.revNAT.Load(); currentRev != nil {
for addr, existing := range *currentRev {
revMap[addr] = existing
}
}
v4Address, v6Address := port.PortAddresses()
if v4Address.IsValid() {
revMap[v4Address] = nat
}
if v6Address.IsValid() {
revMap[v6Address] = nat
}
d.revNAT.Store(&revMap)
return nat
}
@ -473,6 +501,12 @@ func (d *ForwardDispatcher) Flush() {
d.flushPort(nat)
}
d.activeNATs = d.activeNATs[:0]
if retain := max(d.segmentUsed, segmentRetainCount); len(d.segmentBuffers) > retain {
clear(d.segmentBuffers[retain:])
d.segmentBuffers = d.segmentBuffers[:retain]
d.segmentSizes = d.segmentSizes[:retain]
}
d.segmentUsed = 0
if len(d.writebackBatch) > 0 {
err := d.writeback.WriteReturnPackets(d.writebackBatch)
if err != nil {
@ -581,6 +615,14 @@ func (r *forwardReturn) ReturnHeadroom() int {
return r.dispatcher.writeback.ReturnHeadroom()
}
type returnDecision uint8
const (
returnPass returnDecision = iota
returnWrite
returnDrop
)
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
if r.closed.Load() {
return packets
@ -590,65 +632,96 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
return packets
}
natList := *natListPtr
var revMap map[netip.Addr]*portNAT
if revPtr := r.dispatcher.revNAT.Load(); revPtr != nil {
revMap = *revPtr
}
headroom := r.dispatcher.writeback.ReturnHeadroom()
now := r.dispatcher.now()
if len(packets) == 1 {
switch r.classifyReturn(packets[0], natList, revMap, headroom, now) {
case returnWrite:
if err := r.dispatcher.writeback.WriteReturnPackets(packets[:1]); err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
return packets[:0]
case returnDrop:
return packets[:0]
default:
return packets
}
}
unconsumed := packets[:0]
var writeBatch [][]byte
now := r.dispatcher.now()
for _, raw := range packets {
if len(raw) < headroom+header.IPv4MinimumSize {
switch r.classifyReturn(raw, natList, revMap, headroom, now) {
case returnWrite:
writeBatch = append(writeBatch, raw)
case returnDrop:
default:
unconsumed = append(unconsumed, raw)
continue
}
parsed, ok := parseForwardPacket(raw[headroom:])
if !ok || parsed.fragment {
unconsumed = append(unconsumed, raw)
continue
}
if !parsed.hasFlow {
if parsed.isICMPError() && returnICMPError(natList, &parsed) {
writeBatch = append(writeBatch, raw)
} else {
unconsumed = append(unconsumed, raw)
}
continue
}
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(parsed.flowKey())
if flow != nil {
break
}
}
if flow == nil {
unconsumed = append(unconsumed, raw)
continue
}
if flow.closed.Load() {
continue
}
if flow.tracker != nil {
flow.tracker.CountReverse(len(raw) - headroom)
}
flow.observeReverse(&parsed, now)
if parsed.isTCPSyn() {
applyRewriteRaw(&parsed, &flow.reverseRule)
clampTCPMSS(&parsed, flow.effectiveMTU)
recomputeChecksums(&parsed)
} else {
applyRewrite(&parsed, &flow.reverseRule)
}
writeBatch = append(writeBatch, raw)
}
if len(writeBatch) > 0 {
err := r.dispatcher.writeback.WriteReturnPackets(writeBatch)
if err != nil {
if err := r.dispatcher.writeback.WriteReturnPackets(writeBatch); err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
}
return unconsumed
}
func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
func (r *forwardReturn) classifyReturn(raw []byte, natList []*portNAT, revMap map[netip.Addr]*portNAT, headroom int, now int64) returnDecision {
if len(raw) < headroom+header.IPv4MinimumSize {
return returnPass
}
parsed, ok := parseForwardPacket(raw[headroom:])
if !ok || parsed.fragment {
return returnPass
}
if !parsed.hasFlow {
if parsed.isICMPError() && returnICMPError(natList, revMap, &parsed) {
return returnWrite
}
return returnPass
}
flow := findReverseFlow(natList, revMap, parsed.flowKey())
if flow == nil {
return returnPass
}
if flow.closed.Load() {
return returnDrop
}
if flow.tracker != nil {
flow.tracker.CountReverse(len(raw) - headroom)
}
flow.observeReverse(&parsed, now)
if parsed.isTCPSyn() {
applyRewriteRaw(&parsed, &flow.reverseRule)
clampTCPMSS(&parsed, flow.effectiveMTU)
recomputeChecksums(&parsed)
} else {
applyRewrite(&parsed, &flow.reverseRule)
}
return returnWrite
}
func findReverseFlow(natList []*portNAT, revMap map[netip.Addr]*portNAT, key flowKey) *forwardFlow {
if nat, ok := revMap[key.destination.Addr()]; ok {
if flow := nat.lookup(key); flow != nil {
return flow
}
}
for _, nat := range natList {
if flow := nat.lookup(key); flow != nil {
return flow
}
}
return nil
}
func returnICMPError(natList []*portNAT, revMap map[netip.Addr]*portNAT, parsed *forwardPacket) bool {
inner, ok := parsed.icmpErrorInner()
if !ok {
return false
@ -657,20 +730,13 @@ func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
if !parsedInner {
return false
}
key := embedded.flowKey().reversed()
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(key)
if flow != nil {
break
}
}
flow := findReverseFlow(natList, revMap, embedded.flowKey().reversed())
if flow == nil || flow.closed.Load() {
return false
}
rewriteEmbeddedSource(&embedded, tcpip.AddrFromSlice(flow.clientAddress.AsSlice()), flow.clientSelector, true)
rewriteEmbeddedSource(&embedded, addrToTCPIP(flow.clientAddress), flow.clientSelector, true)
if flow.dnatAddress || flow.dnatPort {
rewriteEmbeddedDestination(&embedded, tcpip.AddrFromSlice(flow.clientDestinationAddress.AsSlice()), flow.clientDestinationPort, flow.dnatPort)
rewriteEmbeddedDestination(&embedded, addrToTCPIP(flow.clientDestinationAddress), flow.clientDestinationPort, flow.dnatPort)
}
parsed.network.SetDestinationAddr(flow.clientAddress)
if parsed.network.SourceAddr() == flow.serverAddress {