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 {
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,41 +632,66 @@ 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
}
}
if len(writeBatch) > 0 {
if err := r.dispatcher.writeback.WriteReturnPackets(writeBatch); err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
}
return unconsumed
}
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 {
unconsumed = append(unconsumed, raw)
continue
return returnPass
}
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 parsed.isICMPError() && returnICMPError(natList, revMap, &parsed) {
return returnWrite
}
return returnPass
}
flow := findReverseFlow(natList, revMap, parsed.flowKey())
if flow == nil {
unconsumed = append(unconsumed, raw)
continue
return returnPass
}
if flow.closed.Load() {
continue
return returnDrop
}
if flow.tracker != nil {
flow.tracker.CountReverse(len(raw) - headroom)
@ -637,18 +704,24 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
} else {
applyRewrite(&parsed, &flow.reverseRule)
}
writeBatch = append(writeBatch, raw)
}
if len(writeBatch) > 0 {
err := r.dispatcher.writeback.WriteReturnPackets(writeBatch)
if err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
}
return unconsumed
return returnWrite
}
func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
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 {

View file

@ -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))

View file

@ -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)

View file

@ -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] {

View file

@ -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

View file

@ -42,16 +42,22 @@ func (c *gLazyConn) HandshakeContext(ctx context.Context) error {
wq waiter.Queue
endpoint tcpip.Endpoint
)
handshakeCtx, cancel := context.WithCancel(ctx)
var cancel context.CancelFunc
if parentDone := c.parentCtx.Done(); parentDone != nil {
var handshakeCtx context.Context
handshakeCtx, cancel = context.WithCancel(ctx)
go func() {
select {
case <-c.parentCtx.Done():
case <-parentDone:
wq.Notify(wq.Events())
case <-handshakeCtx.Done():
}
}()
}
endpoint, err := c.request.CreateEndpoint(&wq)
if cancel != nil {
cancel()
}
if err != nil {
gErr := gonet.TranslateNetstackError(err)
c.handshakeErr = gErr

View file

@ -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
)

View file

@ -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
)

View file

@ -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
}

View file

@ -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")

View file

@ -54,19 +54,37 @@ 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 {
n.portAccess.RLock()
@ -91,6 +109,27 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint1
return port
}
n.addrAccess.Lock()
defer n.addrAccess.Unlock()
if port, loaded = n.addrMap[key]; loaded {
return port
}
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.addrMap[key] = nextPort
return nextPort
}
func (n *TCPNat) allocatePortLocked() (uint16, bool) {
for range 65535 - 10000 + 1 {
nextPort := n.portIndex
if nextPort == 0 {
nextPort = 10000
@ -98,14 +137,9 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint1
} else {
n.portIndex++
}
n.addrMap[key] = nextPort
n.addrAccess.Unlock()
n.portAccess.Lock()
n.portMap[nextPort] = &TCPSession{
Source: source,
Destination: destination,
LastActive: time.Now(),
if _, occupied := n.portMap[nextPort]; !occupied {
return nextPort, true
}
n.portAccess.Unlock()
return nextPort
}
return 0, false
}

View file

@ -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)

View file

@ -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
}