From c17af6ee8c77719622d40169e4fb05f0fea00ef8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Tue, 7 Jul 2026 15:36:21 +0800 Subject: [PATCH] Minor fixes --- flow_dispatch.go | 192 ++++++++++++++++++++++----------- flow_mtu.go | 42 +++++--- monitor_android.go | 8 +- monitor_shared.go | 14 +-- stack_gvisor_icmp.go | 4 +- stack_gvisor_lazy.go | 24 +++-- stack_gvisor_tcpbuf_default.go | 4 +- stack_gvisor_tcpbuf_ios.go | 4 +- stack_gvisor_udp.go | 9 +- stack_system.go | 136 +++++++++++++++++------ stack_system_nat.go | 68 +++++++++--- tun_offload.go | 8 +- tun_offload_linux.go | 25 +++-- 13 files changed, 357 insertions(+), 181 deletions(-) diff --git a/flow_dispatch.go b/flow_dispatch.go index b70dc1f..e35a521 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -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 { diff --git a/flow_mtu.go b/flow_mtu.go index 26c675b..6ce275b 100644 --- a/flow_mtu.go +++ b/flow_mtu.go @@ -5,7 +5,9 @@ import ( E "github.com/sagernet/sing/common/exceptions" ) -const segmentScratchCount = 128 +// segmentRetainCount bounds how many segment buffers survive a Flush; the pool +// grows to the burst high-water mark within a batch and is trimmed afterwards. +const segmentRetainCount = 128 // Linux delivers TSO aggregates to the TUN even with IFF_VNET_HDR off // (observed on 6.x: the pre-segmentation skb is handed to the fd as-is). @@ -35,33 +37,39 @@ func (d *ForwardDispatcher) resegmentTCP(flow *forwardFlow, packet *forwardPacke gsoType = GSOTCPv6 } neededSegments := max((len(raw)-totalHeaderLength+segmentSize-1)/segmentSize, 1) - if d.segmentBuffers == nil || len(d.segmentBuffers) < neededSegments || len(d.segmentBuffers[0]) < int(flow.effectiveMTU) { - bufferSize := int(flow.effectiveMTU) - if d.segmentBuffers != nil && len(d.segmentBuffers[0]) > bufferSize { - bufferSize = len(d.segmentBuffers[0]) - } - segmentCount := max(neededSegments, segmentScratchCount, len(d.segmentBuffers)) - d.segmentBuffers = make([][]byte, segmentCount) - for i := range d.segmentBuffers { - d.segmentBuffers[i] = make([]byte, bufferSize) - } - d.segmentSizes = make([]int, segmentCount) - } + bufs, sizes := d.reserveSegments(neededSegments, int(flow.effectiveMTU)) n, err := GSOSplit(raw, GSOOptions{ GSOType: gsoType, HdrLen: uint16(totalHeaderLength), CsumStart: uint16(headerLength), CsumOffset: header.TCPChecksumOffset, GSOSize: uint16(segmentSize), - }, d.segmentBuffers, d.segmentSizes, 0) + }, bufs, sizes, 0) if err != nil { d.logger.Trace(E.Cause(err, "resegment packet")) return } for i := range n { - d.stagePort(flow.nat, d.segmentBuffers[i][:d.segmentSizes[i]]) + d.stagePort(flow.nat, bufs[i][:sizes[i]]) } - d.flushPort(flow.nat) +} + +func (d *ForwardDispatcher) reserveSegments(count, size int) ([][]byte, []int) { + start := d.segmentUsed + end := start + count + for len(d.segmentBuffers) < end { + d.segmentBuffers = append(d.segmentBuffers, make([]byte, size)) + d.segmentSizes = append(d.segmentSizes, 0) + } + for i := start; i < end; i++ { + if cap(d.segmentBuffers[i]) < size { + d.segmentBuffers[i] = make([]byte, size) + } else { + d.segmentBuffers[i] = d.segmentBuffers[i][:size] + } + } + d.segmentUsed = end + return d.segmentBuffers[start:end], d.segmentSizes[start:end] } const synthesizedTTL = 64 @@ -79,7 +87,7 @@ func fragmentIPv4Packet(packet header.IPv4, effectiveMTU uint32) ([][]byte, bool baseOffset := packet.FragmentOffset() originalMore := packet.Flags()&header.IPv4FlagMoreFragments != 0 baseFlags := packet.Flags() &^ header.IPv4FlagMoreFragments - var fragments [][]byte + fragments := make([][]byte, 0, (len(payload)+maxFragmentPayload-1)/maxFragmentPayload) for start := 0; start < len(payload); start += maxFragmentPayload { end := min(start+maxFragmentPayload, len(payload)) fragment := header.IPv4(make([]byte, headerLength+end-start)) diff --git a/monitor_android.go b/monitor_android.go index c83440d..5d0e9c6 100644 --- a/monitor_android.go +++ b/monitor_android.go @@ -11,7 +11,7 @@ func (m *defaultInterfaceMonitor) checkUpdate() error { return E.Cause(err, "list rules") } - oldVPNEnabled := m.androidVPNEnabled + oldVPNEnabled := m.androidVPNEnabled.Load() var defaultTableIndex int var vpnEnabled bool for _, rule := range ruleList { @@ -30,7 +30,7 @@ func (m *defaultInterfaceMonitor) checkUpdate() error { break } } - m.androidVPNEnabled = vpnEnabled + m.androidVPNEnabled.Store(vpnEnabled) if defaultTableIndex == 0 { return ErrNoRoute @@ -56,11 +56,11 @@ func (m *defaultInterfaceMonitor) checkUpdate() error { return E.Cause(err, "find updated interface: ", link.Attrs().Name) } oldInterface := m.defaultInterface.Swap(newInterface) - if oldInterface != nil && oldInterface.Equals(*newInterface) && oldVPNEnabled == m.androidVPNEnabled { + if oldInterface != nil && oldInterface.Equals(*newInterface) && oldVPNEnabled == m.androidVPNEnabled.Load() { return nil } var flags int - if oldVPNEnabled != m.androidVPNEnabled { + if oldVPNEnabled != m.androidVPNEnabled.Load() { flags = FlagAndroidVPNUpdate } m.emit(newInterface, flags) diff --git a/monitor_shared.go b/monitor_shared.go index 8d239d2..ad48ce0 100644 --- a/monitor_shared.go +++ b/monitor_shared.go @@ -39,8 +39,8 @@ type defaultInterfaceMonitor struct { overrideAndroidVPN bool underNetworkExtension bool defaultInterface atomic.Pointer[control.Interface] - androidVPNEnabled bool - noRoute bool + androidVPNEnabled atomic.Bool + noRoute atomic.Bool networkMonitor NetworkUpdateMonitor logger logger.Logger checkUpdateTimer *time.Timer @@ -67,6 +67,8 @@ func (m *defaultInterfaceMonitor) Start() error { } func (m *defaultInterfaceMonitor) delayCheckUpdate() { + m.access.Lock() + defer m.access.Unlock() if m.checkUpdateTimer == nil { m.checkUpdateTimer = time.AfterFunc(time.Second, m.postCheckUpdate) } else { @@ -82,15 +84,15 @@ func (m *defaultInterfaceMonitor) postCheckUpdate() { } err = m.checkUpdate() if errors.Is(err, ErrNoRoute) { - if !m.noRoute { - m.noRoute = true + if !m.noRoute.Load() { + m.noRoute.Store(true) m.defaultInterface.Store(nil) m.emit(nil, 0) } } else if err != nil { m.logger.Error("check interface: ", err) } else { - m.noRoute = false + m.noRoute.Store(false) } } @@ -110,7 +112,7 @@ func (m *defaultInterfaceMonitor) OverrideAndroidVPN() bool { } func (m *defaultInterfaceMonitor) AndroidVPNEnabled() bool { - return m.androidVPNEnabled + return m.androidVPNEnabled.Load() } func (m *defaultInterfaceMonitor) RegisterCallback(callback DefaultInterfaceUpdateCallback) *list.Element[DefaultInterfaceUpdateCallback] { diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 96466c1..9b742e5 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -133,7 +133,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa DefaultNIC, id.LocalAddress, id.RemoteAddress, - header.IPv6ProtocolNumber, + header.IPv4ProtocolNumber, false, ) if gErr != nil { @@ -184,7 +184,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa PayloadCsum: pkt.Data().Checksum(), PayloadLen: pkt.Data().Size(), })) - outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber) + outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv6ProtocolNumber) if gErr != nil { f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv6 network endpoint")) return true diff --git a/stack_gvisor_lazy.go b/stack_gvisor_lazy.go index 258392d..96d8897 100644 --- a/stack_gvisor_lazy.go +++ b/stack_gvisor_lazy.go @@ -42,16 +42,22 @@ func (c *gLazyConn) HandshakeContext(ctx context.Context) error { wq waiter.Queue endpoint tcpip.Endpoint ) - handshakeCtx, cancel := context.WithCancel(ctx) - go func() { - select { - case <-c.parentCtx.Done(): - wq.Notify(wq.Events()) - case <-handshakeCtx.Done(): - } - }() + var cancel context.CancelFunc + if parentDone := c.parentCtx.Done(); parentDone != nil { + var handshakeCtx context.Context + handshakeCtx, cancel = context.WithCancel(ctx) + go func() { + select { + case <-parentDone: + wq.Notify(wq.Events()) + case <-handshakeCtx.Done(): + } + }() + } endpoint, err := c.request.CreateEndpoint(&wq) - cancel() + if cancel != nil { + cancel() + } if err != nil { gErr := gonet.TranslateNetstackError(err) c.handshakeErr = gErr diff --git a/stack_gvisor_tcpbuf_default.go b/stack_gvisor_tcpbuf_default.go index f636d1a..fc3bc2b 100644 --- a/stack_gvisor_tcpbuf_default.go +++ b/stack_gvisor_tcpbuf_default.go @@ -9,10 +9,10 @@ import "github.com/sagernet/gvisor/pkg/tcpip/transport/tcp" const ( tcpRXBufMinSize = tcp.MinBufferSize - tcpRXBufDefSize = tcp.DefaultSendBufferSize + tcpRXBufDefSize = tcp.DefaultReceiveBufferSize tcpRXBufMaxSize = 8 << 20 // 8MiB tcpTXBufMinSize = tcp.MinBufferSize - tcpTXBufDefSize = tcp.DefaultReceiveBufferSize + tcpTXBufDefSize = tcp.DefaultSendBufferSize tcpTXBufMaxSize = 6 << 20 // 6MiB ) diff --git a/stack_gvisor_tcpbuf_ios.go b/stack_gvisor_tcpbuf_ios.go index 495e59b..6704c9d 100644 --- a/stack_gvisor_tcpbuf_ios.go +++ b/stack_gvisor_tcpbuf_ios.go @@ -12,10 +12,10 @@ const ( // unchanged on iOS for now as to not increase pressure towards the // NetworkExtension memory limit. tcpRXBufMinSize = tcp.MinBufferSize - tcpRXBufDefSize = tcp.DefaultSendBufferSize + tcpRXBufDefSize = tcp.DefaultReceiveBufferSize tcpRXBufMaxSize = tcp.MaxBufferSize tcpTXBufMinSize = tcp.MinBufferSize - tcpTXBufDefSize = tcp.DefaultReceiveBufferSize + tcpTXBufDefSize = tcp.DefaultSendBufferSize tcpTXBufMaxSize = tcp.MaxBufferSize ) diff --git a/stack_gvisor_udp.go b/stack_gvisor_udp.go index 3e4afdd..2ae54cf 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -44,12 +44,9 @@ func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, t func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool { source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort) destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort) - bufferRange := pkt.Data().AsRange() - var bufferSlices [][]byte - rangeIterate(bufferRange, func(view *buffer.View) { - bufferSlices = append(bufferSlices, view.AsSlice()) - }) - f.udpNat.NewPacket(bufferSlices, source, destination, pkt) + data := pkt.Data() + payload, _ := data.PullUp(data.Size()) + f.udpNat.NewPacket([][]byte{payload}, source, destination, pkt) return true } diff --git a/stack_system.go b/stack_system.go index 8ab9be6..46f6f87 100644 --- a/stack_system.go +++ b/stack_system.go @@ -9,6 +9,7 @@ import ( "syscall" "time" + "github.com/sagernet/sing-tun/gtcpip" "github.com/sagernet/sing-tun/gtcpip/checksum" "github.com/sagernet/sing-tun/gtcpip/header" "github.com/sagernet/sing/common" @@ -448,36 +449,30 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err if session == nil { return false, E.New("ipv4: tcp: session not found: ", destination.Port()) } - ipHdr.SetSourceAddr(session.Destination.Addr()) - tcpHdr.SetSourcePort(session.Destination.Port()) - ipHdr.SetDestinationAddr(session.Source.Addr()) - tcpHdr.SetDestinationPort(session.Source.Port()) + rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, + session.Destination.Addr(), session.Destination.Port(), true, + session.Source.Addr(), session.Source.Port(), true) } else { var loopback bool for _, inet4LoopbackAddress := range s.inet4LoopbackAddress { if destination.Addr() == inet4LoopbackAddress { - ipHdr.SetDestinationAddr(ipHdr.SourceAddr()) - ipHdr.SetSourceAddr(inet4LoopbackAddress) + rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, + inet4LoopbackAddress, 0, false, + source.Addr(), 0, false) loopback = true break } } if !loopback { natPort := s.tcpNat.Lookup(source, destination) - ipHdr.SetSourceAddr(s.inet4NextAddress) - tcpHdr.SetSourcePort(natPort) - ipHdr.SetDestinationAddr(s.inet4Address) - tcpHdr.SetDestinationPort(s.tcpPort) + if natPort == 0 { + return false, E.New("ipv4: tcp: NAT port space exhausted") + } + rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, + s.inet4NextAddress, natPort, true, + s.inet4Address, s.tcpPort, true) } } - if !s.txChecksumOffload { - tcpHdr.SetChecksum(^checksum.Checksum(tcpHdr.Payload(), tcpHdr.CalculateChecksum( - header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()), - ))) - } else { - tcpHdr.SetChecksum(0) - } - ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) return true, nil } @@ -491,38 +486,109 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err if session == nil { return false, E.New("ipv6: tcp: session not found: ", destination.Port()) } - ipHdr.SetSourceAddr(session.Destination.Addr()) - tcpHdr.SetSourcePort(session.Destination.Port()) - ipHdr.SetDestinationAddr(session.Source.Addr()) - tcpHdr.SetDestinationPort(session.Source.Port()) + rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, + session.Destination.Addr(), session.Destination.Port(), true, + session.Source.Addr(), session.Source.Port(), true) } else { var loopback bool for _, inet6LoopbackAddress := range s.inet6LoopbackAddress { if destination.Addr() == inet6LoopbackAddress { - ipHdr.SetDestinationAddr(ipHdr.SourceAddr()) - ipHdr.SetSourceAddr(inet6LoopbackAddress) + rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, + inet6LoopbackAddress, 0, false, + source.Addr(), 0, false) loopback = true break } } if !loopback { natPort := s.tcpNat.Lookup(source, destination) - ipHdr.SetSourceAddr(s.inet6NextAddress) - tcpHdr.SetSourcePort(natPort) - ipHdr.SetDestinationAddr(s.inet6Address) - tcpHdr.SetDestinationPort(s.tcpPort6) + if natPort == 0 { + return false, E.New("ipv6: tcp: NAT port space exhausted") + } + rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, + s.inet6NextAddress, natPort, true, + s.inet6Address, s.tcpPort6, true) } } - if !s.txChecksumOffload { - tcpHdr.SetChecksum(^checksum.Checksum(tcpHdr.Payload(), tcpHdr.CalculateChecksum( - header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()), - ))) - } else { - tcpHdr.SetChecksum(0) - } return true, nil } +func rewriteIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP, txChecksumOffload bool, + newSource netip.Addr, newSourcePort uint16, rewriteSourcePort bool, + newDestination netip.Addr, newDestinationPort uint16, rewriteDestinationPort bool, +) { + oldSource := ipHdr.SourceAddress() + oldDestination := ipHdr.DestinationAddress() + newSourceAddr := tcpip.AddrFrom4(newSource.As4()) + newDestinationAddr := tcpip.AddrFrom4(newDestination.As4()) + if newSourceAddr != oldSource { + ipHdr.SetSourceAddressWithChecksumUpdate(newSourceAddr) + if !txChecksumOffload { + tcpHdr.UpdateChecksumPseudoHeaderAddress(oldSource, newSourceAddr, true) + } + } + if newDestinationAddr != oldDestination { + ipHdr.SetDestinationAddressWithChecksumUpdate(newDestinationAddr) + if !txChecksumOffload { + tcpHdr.UpdateChecksumPseudoHeaderAddress(oldDestination, newDestinationAddr, true) + } + } + if txChecksumOffload { + if rewriteSourcePort { + tcpHdr.SetSourcePort(newSourcePort) + } + if rewriteDestinationPort { + tcpHdr.SetDestinationPort(newDestinationPort) + } + tcpHdr.SetChecksum(0) + } else { + if rewriteSourcePort { + tcpHdr.SetSourcePortWithChecksumUpdate(newSourcePort) + } + if rewriteDestinationPort { + tcpHdr.SetDestinationPortWithChecksumUpdate(newDestinationPort) + } + } +} + +func rewriteIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP, txChecksumOffload bool, + newSource netip.Addr, newSourcePort uint16, rewriteSourcePort bool, + newDestination netip.Addr, newDestinationPort uint16, rewriteDestinationPort bool, +) { + oldSource := ipHdr.SourceAddress() + oldDestination := ipHdr.DestinationAddress() + newSourceAddr := tcpip.AddrFrom16(newSource.As16()) + newDestinationAddr := tcpip.AddrFrom16(newDestination.As16()) + if newSourceAddr != oldSource { + ipHdr.SetSourceAddress(newSourceAddr) + if !txChecksumOffload { + tcpHdr.UpdateChecksumPseudoHeaderAddress(oldSource, newSourceAddr, true) + } + } + if newDestinationAddr != oldDestination { + ipHdr.SetDestinationAddress(newDestinationAddr) + if !txChecksumOffload { + tcpHdr.UpdateChecksumPseudoHeaderAddress(oldDestination, newDestinationAddr, true) + } + } + if txChecksumOffload { + if rewriteSourcePort { + tcpHdr.SetSourcePort(newSourcePort) + } + if rewriteDestinationPort { + tcpHdr.SetDestinationPort(newDestinationPort) + } + tcpHdr.SetChecksum(0) + } else { + if rewriteSourcePort { + tcpHdr.SetSourcePortWithChecksumUpdate(newSourcePort) + } + if rewriteDestinationPort { + tcpHdr.SetDestinationPortWithChecksumUpdate(newDestinationPort) + } + } +} + func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error { if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 { return E.New("ipv4: fragment dropped") diff --git a/stack_system_nat.go b/stack_system_nat.go index fd0e382..2fec29c 100644 --- a/stack_system_nat.go +++ b/stack_system_nat.go @@ -54,18 +54,36 @@ func (n *TCPNat) loopCheckTimeout(ctx context.Context) { func (n *TCPNat) checkTimeout() { now := time.Now() - n.portAccess.Lock() - defer n.portAccess.Unlock() - n.addrAccess.Lock() - defer n.addrAccess.Unlock() + type expiredSession struct { + port uint16 + session *TCPSession + } + var expired []expiredSession + n.portAccess.RLock() for natPort, session := range n.portMap { session.Lock() - if now.Sub(session.LastActive) > n.timeout { - delete(n.addrMap, tcpNatKey{Source: session.Source, Destination: session.Destination}) - delete(n.portMap, natPort) - } + timedOut := now.Sub(session.LastActive) > n.timeout session.Unlock() + if timedOut { + expired = append(expired, expiredSession{port: natPort, session: session}) + } } + n.portAccess.RUnlock() + if len(expired) == 0 { + return + } + n.addrAccess.Lock() + n.portAccess.Lock() + for _, e := range expired { + e.session.Lock() + if now.Sub(e.session.LastActive) > n.timeout { + delete(n.addrMap, tcpNatKey{Source: e.session.Source, Destination: e.session.Destination}) + delete(n.portMap, e.port) + } + e.session.Unlock() + } + n.portAccess.Unlock() + n.addrAccess.Unlock() } func (n *TCPNat) LookupBack(port uint16) *TCPSession { @@ -91,21 +109,37 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint1 return port } n.addrAccess.Lock() - nextPort := n.portIndex - if nextPort == 0 { - nextPort = 10000 - n.portIndex = 10001 - } else { - n.portIndex++ + defer n.addrAccess.Unlock() + if port, loaded = n.addrMap[key]; loaded { + return port } - n.addrMap[key] = nextPort - n.addrAccess.Unlock() n.portAccess.Lock() + defer n.portAccess.Unlock() + nextPort, ok := n.allocatePortLocked() + if !ok { + return 0 + } n.portMap[nextPort] = &TCPSession{ Source: source, Destination: destination, LastActive: time.Now(), } - n.portAccess.Unlock() + n.addrMap[key] = nextPort return nextPort } + +func (n *TCPNat) allocatePortLocked() (uint16, bool) { + for range 65535 - 10000 + 1 { + nextPort := n.portIndex + if nextPort == 0 { + nextPort = 10000 + n.portIndex = 10001 + } else { + n.portIndex++ + } + if _, occupied := n.portMap[nextPort]; !occupied { + return nextPort, true + } + } + return 0, false +} diff --git a/tun_offload.go b/tun_offload.go index 3bf6268..f68fbfe 100644 --- a/tun_offload.go +++ b/tun_offload.go @@ -156,6 +156,7 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO } else { protocol = ipProtoUDP } + pseudoSumBase := header.PseudoHeaderChecksum(tcpip.TransportProtocolNumber(protocol), in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], 0) nextSegmentDataAt := int(options.HdrLen) i := 0 for ; nextSegmentDataAt < len(in); i++ { @@ -168,7 +169,7 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO sizes[i] = totalLen out := outBufs[i][outOffset:] - copy(out, in[:iphLen]) + copy(out[:options.HdrLen], in[:options.HdrLen]) if ipVersion == 4 { // For IPv4 we are responsible for incrementing the ID field, // updating the total len field, and recalculating the header @@ -187,9 +188,6 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO binary.BigEndian.PutUint16(out[4:], uint16(totalLen-iphLen)) } - // copy transport header - copy(out[options.CsumStart:options.HdrLen], in[options.CsumStart:options.HdrLen]) - if protocol == ipProtoTCP { // set TCP seq and adjust TCP flags tcpSeq := firstTCPSeqNum + uint32(options.GSOSize*uint16(i)) @@ -211,7 +209,7 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO out[transportCsumAt], out[transportCsumAt+1] = 0, 0 // clear tcp/udp checksum transportHeaderLen := int(options.HdrLen - options.CsumStart) lenForPseudo := uint16(transportHeaderLen + segmentDataLen) - transportCSum := header.PseudoHeaderChecksum(tcpip.TransportProtocolNumber(protocol), in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], lenForPseudo) + transportCSum := checksum.Combine(pseudoSumBase, lenForPseudo) transportCSum = ^checksum.Checksum(out[options.CsumStart:totalLen], transportCSum) binary.BigEndian.PutUint16(out[options.CsumStart+options.CsumOffset:], transportCSum) diff --git a/tun_offload_linux.go b/tun_offload_linux.go index 57834ef..4e4dc79 100644 --- a/tun_offload_linux.go +++ b/tun_offload_linux.go @@ -129,14 +129,12 @@ func (t *tcpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, t if ok { return items, ok } - // TODO: insert() performs another map lookup. This could be rearranged to avoid. - t.insert(pkt, srcAddrOffset, dstAddrOffset, tcphOffset, tcphLen, bufsIndex) + t.insert(key, pkt, tcphOffset, tcphLen, bufsIndex) return nil, false } // insert an item in the table for the provided packet and packet metadata. -func (t *tcpGROTable) insert(pkt []byte, srcAddrOffset, dstAddrOffset, tcphOffset, tcphLen, bufsIndex int) { - key := newTCPFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset) +func (t *tcpGROTable) insert(key tcpFlowKey, pkt []byte, tcphOffset, tcphLen, bufsIndex int) { item := tcpGROItem{ key: key, bufsIndex: uint16(bufsIndex), @@ -236,14 +234,12 @@ func (u *udpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, u if ok { return items, ok } - // TODO: insert() performs another map lookup. This could be rearranged to avoid. - u.insert(pkt, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex, false) + u.insert(key, pkt, udphOffset, bufsIndex, false) return nil, false } // insert an item in the table for the provided packet and packet metadata. -func (u *udpGROTable) insert(pkt []byte, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex int, cSumKnownInvalid bool) { - key := newUDPFlowKey(pkt, srcAddrOffset, dstAddrOffset, udphOffset) +func (u *udpGROTable) insert(key udpFlowKey, pkt []byte, udphOffset, bufsIndex int, cSumKnownInvalid bool) { item := udpGROItem{ key: key, bufsIndex: uint16(bufsIndex), @@ -456,7 +452,8 @@ func coalesceUDPPackets(pkt []byte, item *udpGROItem, bufs [][]byte, bufsOffset return coalescePktInvalidCSum } extendBy := len(pkt) - int(headersLen) - bufs[item.bufsIndex] = append(bufs[item.bufsIndex], make([]byte, extendBy)...) + b := bufs[item.bufsIndex] + bufs[item.bufsIndex] = b[:len(b)+extendBy] copy(bufs[item.bufsIndex][bufsOffset+len(pktHead):], pkt[headersLen:]) item.numMerged++ @@ -493,7 +490,8 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize } item.sentSeq = seq extendBy := coalescedLen - len(pktHead) - bufs[pktBuffsIndex] = append(bufs[pktBuffsIndex], make([]byte, extendBy)...) + b := bufs[pktBuffsIndex] + bufs[pktBuffsIndex] = b[:len(b)+extendBy] copy(bufs[pktBuffsIndex][bufsOffset+len(pkt):], bufs[item.bufsIndex][bufsOffset+int(headersLen):]) // Flip the slice headers in bufs as part of prepend. The index of item // is already being tracked for writing. @@ -519,7 +517,8 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize pktHead[item.iphLen+tcpFlagsOffset] |= tcpFlagPSH } extendBy := len(pkt) - int(headersLen) - bufs[item.bufsIndex] = append(bufs[item.bufsIndex], make([]byte, extendBy)...) + b := bufs[item.bufsIndex] + bufs[item.bufsIndex] = b[:len(b)+extendBy] copy(bufs[item.bufsIndex][bufsOffset+len(pktHead):], pkt[headersLen:]) } @@ -639,7 +638,7 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool) } } // failed to coalesce with any other packets; store the item in the flow - table.insert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, tcphLen, pktI) + table.insert(newTCPFlowKey(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen), pkt, iphLen, tcphLen, pktI) return groResultTableInsert } @@ -900,7 +899,7 @@ func udpGRO(bufs [][]byte, offset int, pktI int, table *udpGROTable, isV6 bool) } } // failed to coalesce with any other packets; store the item in the flow - table.insert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, pktI, pktCSumKnownInvalid) + table.insert(newUDPFlowKey(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen), pkt, iphLen, pktI, pktCSumKnownInvalid) return groResultTableInsert }