Minor fixes
This commit is contained in:
parent
14c8f75f7a
commit
c17af6ee8c
13 changed files with 357 additions and 181 deletions
160
flow_dispatch.go
160
flow_dispatch.go
|
|
@ -106,6 +106,7 @@ type ForwardDispatcher struct {
|
||||||
lastSweep int64
|
lastSweep int64
|
||||||
ports map[Port]*portNAT
|
ports map[Port]*portNAT
|
||||||
natList atomic.Pointer[[]*portNAT]
|
natList atomic.Pointer[[]*portNAT]
|
||||||
|
revNAT atomic.Pointer[map[netip.Addr]*portNAT]
|
||||||
|
|
||||||
activeNATs []*portNAT
|
activeNATs []*portNAT
|
||||||
writebackBatch [][]byte
|
writebackBatch [][]byte
|
||||||
|
|
@ -113,6 +114,14 @@ type ForwardDispatcher struct {
|
||||||
|
|
||||||
segmentBuffers [][]byte
|
segmentBuffers [][]byte
|
||||||
segmentSizes []int
|
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 {
|
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)
|
d.forwardToPort(flow, packet, raw)
|
||||||
return true
|
return true
|
||||||
case ActionAccept:
|
case ActionAccept:
|
||||||
entry.deadline = now + int64(entry.idle)
|
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
||||||
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
|
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||||
d.removeEntry(key, entry, FlowCloseReset)
|
d.removeEntry(key, entry, FlowCloseReset)
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
|
if packet.tcpFlags&header.TCPFlagSyn == 0 {
|
||||||
|
entry.idle = tcpEstablishedTimeout
|
||||||
|
}
|
||||||
|
}
|
||||||
|
entry.deadline = now + int64(entry.idle)
|
||||||
return false
|
return false
|
||||||
case ActionReject:
|
case ActionReject:
|
||||||
entry.deadline = now + int64(entry.idle)
|
entry.deadline = now + int64(entry.idle)
|
||||||
|
|
@ -330,24 +345,24 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic
|
||||||
dnatPort: serverPort != clientDestinationPort && !isICMP,
|
dnatPort: serverPort != clientDestinationPort && !isICMP,
|
||||||
}
|
}
|
||||||
flow.forwardRule = rewriteRule{
|
flow.forwardRule = rewriteRule{
|
||||||
sourceAddress: tcpip.AddrFromSlice(portAddress.AsSlice()),
|
sourceAddress: addrToTCPIP(portAddress),
|
||||||
sourcePort: selector,
|
sourcePort: selector,
|
||||||
rewriteSourcePort: true,
|
rewriteSourcePort: true,
|
||||||
}
|
}
|
||||||
if flow.dnatAddress {
|
if flow.dnatAddress {
|
||||||
flow.forwardRule.destinationAddress = tcpip.AddrFromSlice(serverAddress.AsSlice())
|
flow.forwardRule.destinationAddress = addrToTCPIP(serverAddress)
|
||||||
}
|
}
|
||||||
if flow.dnatPort {
|
if flow.dnatPort {
|
||||||
flow.forwardRule.destinationPort = serverPort
|
flow.forwardRule.destinationPort = serverPort
|
||||||
flow.forwardRule.rewriteDestinationPort = true
|
flow.forwardRule.rewriteDestinationPort = true
|
||||||
}
|
}
|
||||||
flow.reverseRule = rewriteRule{
|
flow.reverseRule = rewriteRule{
|
||||||
destinationAddress: tcpip.AddrFromSlice(flow.clientAddress.AsSlice()),
|
destinationAddress: addrToTCPIP(flow.clientAddress),
|
||||||
destinationPort: flow.clientSelector,
|
destinationPort: flow.clientSelector,
|
||||||
rewriteDestinationPort: true,
|
rewriteDestinationPort: true,
|
||||||
}
|
}
|
||||||
if flow.dnatAddress {
|
if flow.dnatAddress {
|
||||||
flow.reverseRule.sourceAddress = tcpip.AddrFromSlice(clientDestinationAddress.AsSlice())
|
flow.reverseRule.sourceAddress = addrToTCPIP(clientDestinationAddress)
|
||||||
}
|
}
|
||||||
if flow.dnatPort {
|
if flow.dnatPort {
|
||||||
flow.reverseRule.sourcePort = clientDestinationPort
|
flow.reverseRule.sourcePort = clientDestinationPort
|
||||||
|
|
@ -371,7 +386,6 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
||||||
err := port.AttachReturn(&d.returnPath)
|
err := port.AttachReturn(&d.returnPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
d.logger.Trace(E.Cause(err, "attach return path"))
|
d.logger.Trace(E.Cause(err, "attach return path"))
|
||||||
d.ports[port] = nil
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
nat = newPortNAT(port)
|
nat = newPortNAT(port)
|
||||||
|
|
@ -383,6 +397,20 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
||||||
}
|
}
|
||||||
natList = append(natList, nat)
|
natList = append(natList, nat)
|
||||||
d.natList.Store(&natList)
|
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
|
return nat
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -473,6 +501,12 @@ func (d *ForwardDispatcher) Flush() {
|
||||||
d.flushPort(nat)
|
d.flushPort(nat)
|
||||||
}
|
}
|
||||||
d.activeNATs = d.activeNATs[:0]
|
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 {
|
if len(d.writebackBatch) > 0 {
|
||||||
err := d.writeback.WriteReturnPackets(d.writebackBatch)
|
err := d.writeback.WriteReturnPackets(d.writebackBatch)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -581,6 +615,14 @@ func (r *forwardReturn) ReturnHeadroom() int {
|
||||||
return r.dispatcher.writeback.ReturnHeadroom()
|
return r.dispatcher.writeback.ReturnHeadroom()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type returnDecision uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
returnPass returnDecision = iota
|
||||||
|
returnWrite
|
||||||
|
returnDrop
|
||||||
|
)
|
||||||
|
|
||||||
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||||
if r.closed.Load() {
|
if r.closed.Load() {
|
||||||
return packets
|
return packets
|
||||||
|
|
@ -590,41 +632,66 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||||
return packets
|
return packets
|
||||||
}
|
}
|
||||||
natList := *natListPtr
|
natList := *natListPtr
|
||||||
|
var revMap map[netip.Addr]*portNAT
|
||||||
|
if revPtr := r.dispatcher.revNAT.Load(); revPtr != nil {
|
||||||
|
revMap = *revPtr
|
||||||
|
}
|
||||||
headroom := r.dispatcher.writeback.ReturnHeadroom()
|
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]
|
unconsumed := packets[:0]
|
||||||
var writeBatch [][]byte
|
var writeBatch [][]byte
|
||||||
now := r.dispatcher.now()
|
|
||||||
for _, raw := range packets {
|
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)
|
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:])
|
parsed, ok := parseForwardPacket(raw[headroom:])
|
||||||
if !ok || parsed.fragment {
|
if !ok || parsed.fragment {
|
||||||
unconsumed = append(unconsumed, raw)
|
return returnPass
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
if !parsed.hasFlow {
|
if !parsed.hasFlow {
|
||||||
if parsed.isICMPError() && returnICMPError(natList, &parsed) {
|
if parsed.isICMPError() && returnICMPError(natList, revMap, &parsed) {
|
||||||
writeBatch = append(writeBatch, raw)
|
return returnWrite
|
||||||
} else {
|
|
||||||
unconsumed = append(unconsumed, raw)
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
var flow *forwardFlow
|
|
||||||
for _, nat := range natList {
|
|
||||||
flow = nat.lookup(parsed.flowKey())
|
|
||||||
if flow != nil {
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
|
return returnPass
|
||||||
}
|
}
|
||||||
|
flow := findReverseFlow(natList, revMap, parsed.flowKey())
|
||||||
if flow == nil {
|
if flow == nil {
|
||||||
unconsumed = append(unconsumed, raw)
|
return returnPass
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
if flow.closed.Load() {
|
if flow.closed.Load() {
|
||||||
continue
|
return returnDrop
|
||||||
}
|
}
|
||||||
if flow.tracker != nil {
|
if flow.tracker != nil {
|
||||||
flow.tracker.CountReverse(len(raw) - headroom)
|
flow.tracker.CountReverse(len(raw) - headroom)
|
||||||
|
|
@ -637,18 +704,24 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||||
} else {
|
} else {
|
||||||
applyRewrite(&parsed, &flow.reverseRule)
|
applyRewrite(&parsed, &flow.reverseRule)
|
||||||
}
|
}
|
||||||
writeBatch = append(writeBatch, raw)
|
return returnWrite
|
||||||
}
|
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
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()
|
inner, ok := parsed.icmpErrorInner()
|
||||||
if !ok {
|
if !ok {
|
||||||
return false
|
return false
|
||||||
|
|
@ -657,20 +730,13 @@ func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
|
||||||
if !parsedInner {
|
if !parsedInner {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
key := embedded.flowKey().reversed()
|
flow := findReverseFlow(natList, revMap, embedded.flowKey().reversed())
|
||||||
var flow *forwardFlow
|
|
||||||
for _, nat := range natList {
|
|
||||||
flow = nat.lookup(key)
|
|
||||||
if flow != nil {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if flow == nil || flow.closed.Load() {
|
if flow == nil || flow.closed.Load() {
|
||||||
return false
|
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 {
|
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)
|
parsed.network.SetDestinationAddr(flow.clientAddress)
|
||||||
if parsed.network.SourceAddr() == flow.serverAddress {
|
if parsed.network.SourceAddr() == flow.serverAddress {
|
||||||
|
|
|
||||||
42
flow_mtu.go
42
flow_mtu.go
|
|
@ -5,7 +5,9 @@ import (
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
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
|
// 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).
|
// (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
|
gsoType = GSOTCPv6
|
||||||
}
|
}
|
||||||
neededSegments := max((len(raw)-totalHeaderLength+segmentSize-1)/segmentSize, 1)
|
neededSegments := max((len(raw)-totalHeaderLength+segmentSize-1)/segmentSize, 1)
|
||||||
if d.segmentBuffers == nil || len(d.segmentBuffers) < neededSegments || len(d.segmentBuffers[0]) < int(flow.effectiveMTU) {
|
bufs, sizes := d.reserveSegments(neededSegments, 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)
|
|
||||||
}
|
|
||||||
n, err := GSOSplit(raw, GSOOptions{
|
n, err := GSOSplit(raw, GSOOptions{
|
||||||
GSOType: gsoType,
|
GSOType: gsoType,
|
||||||
HdrLen: uint16(totalHeaderLength),
|
HdrLen: uint16(totalHeaderLength),
|
||||||
CsumStart: uint16(headerLength),
|
CsumStart: uint16(headerLength),
|
||||||
CsumOffset: header.TCPChecksumOffset,
|
CsumOffset: header.TCPChecksumOffset,
|
||||||
GSOSize: uint16(segmentSize),
|
GSOSize: uint16(segmentSize),
|
||||||
}, d.segmentBuffers, d.segmentSizes, 0)
|
}, bufs, sizes, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
d.logger.Trace(E.Cause(err, "resegment packet"))
|
d.logger.Trace(E.Cause(err, "resegment packet"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
for i := range n {
|
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
|
const synthesizedTTL = 64
|
||||||
|
|
@ -79,7 +87,7 @@ func fragmentIPv4Packet(packet header.IPv4, effectiveMTU uint32) ([][]byte, bool
|
||||||
baseOffset := packet.FragmentOffset()
|
baseOffset := packet.FragmentOffset()
|
||||||
originalMore := packet.Flags()&header.IPv4FlagMoreFragments != 0
|
originalMore := packet.Flags()&header.IPv4FlagMoreFragments != 0
|
||||||
baseFlags := packet.Flags() &^ header.IPv4FlagMoreFragments
|
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 {
|
for start := 0; start < len(payload); start += maxFragmentPayload {
|
||||||
end := min(start+maxFragmentPayload, len(payload))
|
end := min(start+maxFragmentPayload, len(payload))
|
||||||
fragment := header.IPv4(make([]byte, headerLength+end-start))
|
fragment := header.IPv4(make([]byte, headerLength+end-start))
|
||||||
|
|
|
||||||
|
|
@ -11,7 +11,7 @@ func (m *defaultInterfaceMonitor) checkUpdate() error {
|
||||||
return E.Cause(err, "list rules")
|
return E.Cause(err, "list rules")
|
||||||
}
|
}
|
||||||
|
|
||||||
oldVPNEnabled := m.androidVPNEnabled
|
oldVPNEnabled := m.androidVPNEnabled.Load()
|
||||||
var defaultTableIndex int
|
var defaultTableIndex int
|
||||||
var vpnEnabled bool
|
var vpnEnabled bool
|
||||||
for _, rule := range ruleList {
|
for _, rule := range ruleList {
|
||||||
|
|
@ -30,7 +30,7 @@ func (m *defaultInterfaceMonitor) checkUpdate() error {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
m.androidVPNEnabled = vpnEnabled
|
m.androidVPNEnabled.Store(vpnEnabled)
|
||||||
|
|
||||||
if defaultTableIndex == 0 {
|
if defaultTableIndex == 0 {
|
||||||
return ErrNoRoute
|
return ErrNoRoute
|
||||||
|
|
@ -56,11 +56,11 @@ func (m *defaultInterfaceMonitor) checkUpdate() error {
|
||||||
return E.Cause(err, "find updated interface: ", link.Attrs().Name)
|
return E.Cause(err, "find updated interface: ", link.Attrs().Name)
|
||||||
}
|
}
|
||||||
oldInterface := m.defaultInterface.Swap(newInterface)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
var flags int
|
var flags int
|
||||||
if oldVPNEnabled != m.androidVPNEnabled {
|
if oldVPNEnabled != m.androidVPNEnabled.Load() {
|
||||||
flags = FlagAndroidVPNUpdate
|
flags = FlagAndroidVPNUpdate
|
||||||
}
|
}
|
||||||
m.emit(newInterface, flags)
|
m.emit(newInterface, flags)
|
||||||
|
|
|
||||||
|
|
@ -39,8 +39,8 @@ type defaultInterfaceMonitor struct {
|
||||||
overrideAndroidVPN bool
|
overrideAndroidVPN bool
|
||||||
underNetworkExtension bool
|
underNetworkExtension bool
|
||||||
defaultInterface atomic.Pointer[control.Interface]
|
defaultInterface atomic.Pointer[control.Interface]
|
||||||
androidVPNEnabled bool
|
androidVPNEnabled atomic.Bool
|
||||||
noRoute bool
|
noRoute atomic.Bool
|
||||||
networkMonitor NetworkUpdateMonitor
|
networkMonitor NetworkUpdateMonitor
|
||||||
logger logger.Logger
|
logger logger.Logger
|
||||||
checkUpdateTimer *time.Timer
|
checkUpdateTimer *time.Timer
|
||||||
|
|
@ -67,6 +67,8 @@ func (m *defaultInterfaceMonitor) Start() error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *defaultInterfaceMonitor) delayCheckUpdate() {
|
func (m *defaultInterfaceMonitor) delayCheckUpdate() {
|
||||||
|
m.access.Lock()
|
||||||
|
defer m.access.Unlock()
|
||||||
if m.checkUpdateTimer == nil {
|
if m.checkUpdateTimer == nil {
|
||||||
m.checkUpdateTimer = time.AfterFunc(time.Second, m.postCheckUpdate)
|
m.checkUpdateTimer = time.AfterFunc(time.Second, m.postCheckUpdate)
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -82,15 +84,15 @@ func (m *defaultInterfaceMonitor) postCheckUpdate() {
|
||||||
}
|
}
|
||||||
err = m.checkUpdate()
|
err = m.checkUpdate()
|
||||||
if errors.Is(err, ErrNoRoute) {
|
if errors.Is(err, ErrNoRoute) {
|
||||||
if !m.noRoute {
|
if !m.noRoute.Load() {
|
||||||
m.noRoute = true
|
m.noRoute.Store(true)
|
||||||
m.defaultInterface.Store(nil)
|
m.defaultInterface.Store(nil)
|
||||||
m.emit(nil, 0)
|
m.emit(nil, 0)
|
||||||
}
|
}
|
||||||
} else if err != nil {
|
} else if err != nil {
|
||||||
m.logger.Error("check interface: ", err)
|
m.logger.Error("check interface: ", err)
|
||||||
} else {
|
} else {
|
||||||
m.noRoute = false
|
m.noRoute.Store(false)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -110,7 +112,7 @@ func (m *defaultInterfaceMonitor) OverrideAndroidVPN() bool {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *defaultInterfaceMonitor) AndroidVPNEnabled() bool {
|
func (m *defaultInterfaceMonitor) AndroidVPNEnabled() bool {
|
||||||
return m.androidVPNEnabled
|
return m.androidVPNEnabled.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *defaultInterfaceMonitor) RegisterCallback(callback DefaultInterfaceUpdateCallback) *list.Element[DefaultInterfaceUpdateCallback] {
|
func (m *defaultInterfaceMonitor) RegisterCallback(callback DefaultInterfaceUpdateCallback) *list.Element[DefaultInterfaceUpdateCallback] {
|
||||||
|
|
|
||||||
|
|
@ -133,7 +133,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
DefaultNIC,
|
DefaultNIC,
|
||||||
id.LocalAddress,
|
id.LocalAddress,
|
||||||
id.RemoteAddress,
|
id.RemoteAddress,
|
||||||
header.IPv6ProtocolNumber,
|
header.IPv4ProtocolNumber,
|
||||||
false,
|
false,
|
||||||
)
|
)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
|
|
@ -184,7 +184,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
PayloadCsum: pkt.Data().Checksum(),
|
PayloadCsum: pkt.Data().Checksum(),
|
||||||
PayloadLen: pkt.Data().Size(),
|
PayloadLen: pkt.Data().Size(),
|
||||||
}))
|
}))
|
||||||
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber)
|
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv6ProtocolNumber)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv6 network endpoint"))
|
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv6 network endpoint"))
|
||||||
return true
|
return true
|
||||||
|
|
|
||||||
|
|
@ -42,16 +42,22 @@ func (c *gLazyConn) HandshakeContext(ctx context.Context) error {
|
||||||
wq waiter.Queue
|
wq waiter.Queue
|
||||||
endpoint tcpip.Endpoint
|
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() {
|
go func() {
|
||||||
select {
|
select {
|
||||||
case <-c.parentCtx.Done():
|
case <-parentDone:
|
||||||
wq.Notify(wq.Events())
|
wq.Notify(wq.Events())
|
||||||
case <-handshakeCtx.Done():
|
case <-handshakeCtx.Done():
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
}
|
||||||
endpoint, err := c.request.CreateEndpoint(&wq)
|
endpoint, err := c.request.CreateEndpoint(&wq)
|
||||||
|
if cancel != nil {
|
||||||
cancel()
|
cancel()
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
gErr := gonet.TranslateNetstackError(err)
|
gErr := gonet.TranslateNetstackError(err)
|
||||||
c.handshakeErr = gErr
|
c.handshakeErr = gErr
|
||||||
|
|
|
||||||
|
|
@ -9,10 +9,10 @@ import "github.com/sagernet/gvisor/pkg/tcpip/transport/tcp"
|
||||||
|
|
||||||
const (
|
const (
|
||||||
tcpRXBufMinSize = tcp.MinBufferSize
|
tcpRXBufMinSize = tcp.MinBufferSize
|
||||||
tcpRXBufDefSize = tcp.DefaultSendBufferSize
|
tcpRXBufDefSize = tcp.DefaultReceiveBufferSize
|
||||||
tcpRXBufMaxSize = 8 << 20 // 8MiB
|
tcpRXBufMaxSize = 8 << 20 // 8MiB
|
||||||
|
|
||||||
tcpTXBufMinSize = tcp.MinBufferSize
|
tcpTXBufMinSize = tcp.MinBufferSize
|
||||||
tcpTXBufDefSize = tcp.DefaultReceiveBufferSize
|
tcpTXBufDefSize = tcp.DefaultSendBufferSize
|
||||||
tcpTXBufMaxSize = 6 << 20 // 6MiB
|
tcpTXBufMaxSize = 6 << 20 // 6MiB
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -12,10 +12,10 @@ const (
|
||||||
// unchanged on iOS for now as to not increase pressure towards the
|
// unchanged on iOS for now as to not increase pressure towards the
|
||||||
// NetworkExtension memory limit.
|
// NetworkExtension memory limit.
|
||||||
tcpRXBufMinSize = tcp.MinBufferSize
|
tcpRXBufMinSize = tcp.MinBufferSize
|
||||||
tcpRXBufDefSize = tcp.DefaultSendBufferSize
|
tcpRXBufDefSize = tcp.DefaultReceiveBufferSize
|
||||||
tcpRXBufMaxSize = tcp.MaxBufferSize
|
tcpRXBufMaxSize = tcp.MaxBufferSize
|
||||||
|
|
||||||
tcpTXBufMinSize = tcp.MinBufferSize
|
tcpTXBufMinSize = tcp.MinBufferSize
|
||||||
tcpTXBufDefSize = tcp.DefaultReceiveBufferSize
|
tcpTXBufDefSize = tcp.DefaultSendBufferSize
|
||||||
tcpTXBufMaxSize = tcp.MaxBufferSize
|
tcpTXBufMaxSize = tcp.MaxBufferSize
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||||
source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort)
|
source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort)
|
||||||
destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort)
|
destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort)
|
||||||
bufferRange := pkt.Data().AsRange()
|
data := pkt.Data()
|
||||||
var bufferSlices [][]byte
|
payload, _ := data.PullUp(data.Size())
|
||||||
rangeIterate(bufferRange, func(view *buffer.View) {
|
f.udpNat.NewPacket([][]byte{payload}, source, destination, pkt)
|
||||||
bufferSlices = append(bufferSlices, view.AsSlice())
|
|
||||||
})
|
|
||||||
f.udpNat.NewPacket(bufferSlices, source, destination, pkt)
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
134
stack_system.go
134
stack_system.go
|
|
@ -9,6 +9,7 @@ import (
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/sagernet/sing-tun/gtcpip"
|
||||||
"github.com/sagernet/sing-tun/gtcpip/checksum"
|
"github.com/sagernet/sing-tun/gtcpip/checksum"
|
||||||
"github.com/sagernet/sing-tun/gtcpip/header"
|
"github.com/sagernet/sing-tun/gtcpip/header"
|
||||||
"github.com/sagernet/sing/common"
|
"github.com/sagernet/sing/common"
|
||||||
|
|
@ -448,36 +449,30 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
|
||||||
if session == nil {
|
if session == nil {
|
||||||
return false, E.New("ipv4: tcp: session not found: ", destination.Port())
|
return false, E.New("ipv4: tcp: session not found: ", destination.Port())
|
||||||
}
|
}
|
||||||
ipHdr.SetSourceAddr(session.Destination.Addr())
|
rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload,
|
||||||
tcpHdr.SetSourcePort(session.Destination.Port())
|
session.Destination.Addr(), session.Destination.Port(), true,
|
||||||
ipHdr.SetDestinationAddr(session.Source.Addr())
|
session.Source.Addr(), session.Source.Port(), true)
|
||||||
tcpHdr.SetDestinationPort(session.Source.Port())
|
|
||||||
} else {
|
} else {
|
||||||
var loopback bool
|
var loopback bool
|
||||||
for _, inet4LoopbackAddress := range s.inet4LoopbackAddress {
|
for _, inet4LoopbackAddress := range s.inet4LoopbackAddress {
|
||||||
if destination.Addr() == inet4LoopbackAddress {
|
if destination.Addr() == inet4LoopbackAddress {
|
||||||
ipHdr.SetDestinationAddr(ipHdr.SourceAddr())
|
rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload,
|
||||||
ipHdr.SetSourceAddr(inet4LoopbackAddress)
|
inet4LoopbackAddress, 0, false,
|
||||||
|
source.Addr(), 0, false)
|
||||||
loopback = true
|
loopback = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !loopback {
|
if !loopback {
|
||||||
natPort := s.tcpNat.Lookup(source, destination)
|
natPort := s.tcpNat.Lookup(source, destination)
|
||||||
ipHdr.SetSourceAddr(s.inet4NextAddress)
|
if natPort == 0 {
|
||||||
tcpHdr.SetSourcePort(natPort)
|
return false, E.New("ipv4: tcp: NAT port space exhausted")
|
||||||
ipHdr.SetDestinationAddr(s.inet4Address)
|
}
|
||||||
tcpHdr.SetDestinationPort(s.tcpPort)
|
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
|
return true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -491,38 +486,109 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
|
||||||
if session == nil {
|
if session == nil {
|
||||||
return false, E.New("ipv6: tcp: session not found: ", destination.Port())
|
return false, E.New("ipv6: tcp: session not found: ", destination.Port())
|
||||||
}
|
}
|
||||||
ipHdr.SetSourceAddr(session.Destination.Addr())
|
rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload,
|
||||||
tcpHdr.SetSourcePort(session.Destination.Port())
|
session.Destination.Addr(), session.Destination.Port(), true,
|
||||||
ipHdr.SetDestinationAddr(session.Source.Addr())
|
session.Source.Addr(), session.Source.Port(), true)
|
||||||
tcpHdr.SetDestinationPort(session.Source.Port())
|
|
||||||
} else {
|
} else {
|
||||||
var loopback bool
|
var loopback bool
|
||||||
for _, inet6LoopbackAddress := range s.inet6LoopbackAddress {
|
for _, inet6LoopbackAddress := range s.inet6LoopbackAddress {
|
||||||
if destination.Addr() == inet6LoopbackAddress {
|
if destination.Addr() == inet6LoopbackAddress {
|
||||||
ipHdr.SetDestinationAddr(ipHdr.SourceAddr())
|
rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload,
|
||||||
ipHdr.SetSourceAddr(inet6LoopbackAddress)
|
inet6LoopbackAddress, 0, false,
|
||||||
|
source.Addr(), 0, false)
|
||||||
loopback = true
|
loopback = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !loopback {
|
if !loopback {
|
||||||
natPort := s.tcpNat.Lookup(source, destination)
|
natPort := s.tcpNat.Lookup(source, destination)
|
||||||
ipHdr.SetSourceAddr(s.inet6NextAddress)
|
if natPort == 0 {
|
||||||
tcpHdr.SetSourcePort(natPort)
|
return false, E.New("ipv6: tcp: NAT port space exhausted")
|
||||||
ipHdr.SetDestinationAddr(s.inet6Address)
|
|
||||||
tcpHdr.SetDestinationPort(s.tcpPort6)
|
|
||||||
}
|
}
|
||||||
|
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
|
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 {
|
func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error {
|
||||||
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
|
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
|
||||||
return E.New("ipv4: fragment dropped")
|
return E.New("ipv4: fragment dropped")
|
||||||
|
|
|
||||||
|
|
@ -54,18 +54,36 @@ func (n *TCPNat) loopCheckTimeout(ctx context.Context) {
|
||||||
|
|
||||||
func (n *TCPNat) checkTimeout() {
|
func (n *TCPNat) checkTimeout() {
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
n.portAccess.Lock()
|
type expiredSession struct {
|
||||||
defer n.portAccess.Unlock()
|
port uint16
|
||||||
n.addrAccess.Lock()
|
session *TCPSession
|
||||||
defer n.addrAccess.Unlock()
|
}
|
||||||
|
var expired []expiredSession
|
||||||
|
n.portAccess.RLock()
|
||||||
for natPort, session := range n.portMap {
|
for natPort, session := range n.portMap {
|
||||||
session.Lock()
|
session.Lock()
|
||||||
if now.Sub(session.LastActive) > n.timeout {
|
timedOut := now.Sub(session.LastActive) > n.timeout
|
||||||
delete(n.addrMap, tcpNatKey{Source: session.Source, Destination: session.Destination})
|
|
||||||
delete(n.portMap, natPort)
|
|
||||||
}
|
|
||||||
session.Unlock()
|
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 {
|
func (n *TCPNat) LookupBack(port uint16) *TCPSession {
|
||||||
|
|
@ -91,6 +109,27 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint1
|
||||||
return port
|
return port
|
||||||
}
|
}
|
||||||
n.addrAccess.Lock()
|
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
|
nextPort := n.portIndex
|
||||||
if nextPort == 0 {
|
if nextPort == 0 {
|
||||||
nextPort = 10000
|
nextPort = 10000
|
||||||
|
|
@ -98,14 +137,9 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint1
|
||||||
} else {
|
} else {
|
||||||
n.portIndex++
|
n.portIndex++
|
||||||
}
|
}
|
||||||
n.addrMap[key] = nextPort
|
if _, occupied := n.portMap[nextPort]; !occupied {
|
||||||
n.addrAccess.Unlock()
|
return nextPort, true
|
||||||
n.portAccess.Lock()
|
|
||||||
n.portMap[nextPort] = &TCPSession{
|
|
||||||
Source: source,
|
|
||||||
Destination: destination,
|
|
||||||
LastActive: time.Now(),
|
|
||||||
}
|
}
|
||||||
n.portAccess.Unlock()
|
}
|
||||||
return nextPort
|
return 0, false
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -156,6 +156,7 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO
|
||||||
} else {
|
} else {
|
||||||
protocol = ipProtoUDP
|
protocol = ipProtoUDP
|
||||||
}
|
}
|
||||||
|
pseudoSumBase := header.PseudoHeaderChecksum(tcpip.TransportProtocolNumber(protocol), in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], 0)
|
||||||
nextSegmentDataAt := int(options.HdrLen)
|
nextSegmentDataAt := int(options.HdrLen)
|
||||||
i := 0
|
i := 0
|
||||||
for ; nextSegmentDataAt < len(in); i++ {
|
for ; nextSegmentDataAt < len(in); i++ {
|
||||||
|
|
@ -168,7 +169,7 @@ func GSOSplit(in []byte, options GSOOptions, outBufs [][]byte, sizes []int, outO
|
||||||
sizes[i] = totalLen
|
sizes[i] = totalLen
|
||||||
out := outBufs[i][outOffset:]
|
out := outBufs[i][outOffset:]
|
||||||
|
|
||||||
copy(out, in[:iphLen])
|
copy(out[:options.HdrLen], in[:options.HdrLen])
|
||||||
if ipVersion == 4 {
|
if ipVersion == 4 {
|
||||||
// For IPv4 we are responsible for incrementing the ID field,
|
// For IPv4 we are responsible for incrementing the ID field,
|
||||||
// updating the total len field, and recalculating the header
|
// 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))
|
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 {
|
if protocol == ipProtoTCP {
|
||||||
// set TCP seq and adjust TCP flags
|
// set TCP seq and adjust TCP flags
|
||||||
tcpSeq := firstTCPSeqNum + uint32(options.GSOSize*uint16(i))
|
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
|
out[transportCsumAt], out[transportCsumAt+1] = 0, 0 // clear tcp/udp checksum
|
||||||
transportHeaderLen := int(options.HdrLen - options.CsumStart)
|
transportHeaderLen := int(options.HdrLen - options.CsumStart)
|
||||||
lenForPseudo := uint16(transportHeaderLen + segmentDataLen)
|
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)
|
transportCSum = ^checksum.Checksum(out[options.CsumStart:totalLen], transportCSum)
|
||||||
binary.BigEndian.PutUint16(out[options.CsumStart+options.CsumOffset:], transportCSum)
|
binary.BigEndian.PutUint16(out[options.CsumStart+options.CsumOffset:], transportCSum)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -129,14 +129,12 @@ func (t *tcpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, t
|
||||||
if ok {
|
if ok {
|
||||||
return items, ok
|
return items, ok
|
||||||
}
|
}
|
||||||
// TODO: insert() performs another map lookup. This could be rearranged to avoid.
|
t.insert(key, pkt, tcphOffset, tcphLen, bufsIndex)
|
||||||
t.insert(pkt, srcAddrOffset, dstAddrOffset, tcphOffset, tcphLen, bufsIndex)
|
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// insert an item in the table for the provided packet and packet metadata.
|
// 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) {
|
func (t *tcpGROTable) insert(key tcpFlowKey, pkt []byte, tcphOffset, tcphLen, bufsIndex int) {
|
||||||
key := newTCPFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset)
|
|
||||||
item := tcpGROItem{
|
item := tcpGROItem{
|
||||||
key: key,
|
key: key,
|
||||||
bufsIndex: uint16(bufsIndex),
|
bufsIndex: uint16(bufsIndex),
|
||||||
|
|
@ -236,14 +234,12 @@ func (u *udpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, u
|
||||||
if ok {
|
if ok {
|
||||||
return items, ok
|
return items, ok
|
||||||
}
|
}
|
||||||
// TODO: insert() performs another map lookup. This could be rearranged to avoid.
|
u.insert(key, pkt, udphOffset, bufsIndex, false)
|
||||||
u.insert(pkt, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex, false)
|
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// insert an item in the table for the provided packet and packet metadata.
|
// 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) {
|
func (u *udpGROTable) insert(key udpFlowKey, pkt []byte, udphOffset, bufsIndex int, cSumKnownInvalid bool) {
|
||||||
key := newUDPFlowKey(pkt, srcAddrOffset, dstAddrOffset, udphOffset)
|
|
||||||
item := udpGROItem{
|
item := udpGROItem{
|
||||||
key: key,
|
key: key,
|
||||||
bufsIndex: uint16(bufsIndex),
|
bufsIndex: uint16(bufsIndex),
|
||||||
|
|
@ -456,7 +452,8 @@ func coalesceUDPPackets(pkt []byte, item *udpGROItem, bufs [][]byte, bufsOffset
|
||||||
return coalescePktInvalidCSum
|
return coalescePktInvalidCSum
|
||||||
}
|
}
|
||||||
extendBy := len(pkt) - int(headersLen)
|
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:])
|
copy(bufs[item.bufsIndex][bufsOffset+len(pktHead):], pkt[headersLen:])
|
||||||
|
|
||||||
item.numMerged++
|
item.numMerged++
|
||||||
|
|
@ -493,7 +490,8 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize
|
||||||
}
|
}
|
||||||
item.sentSeq = seq
|
item.sentSeq = seq
|
||||||
extendBy := coalescedLen - len(pktHead)
|
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):])
|
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
|
// Flip the slice headers in bufs as part of prepend. The index of item
|
||||||
// is already being tracked for writing.
|
// is already being tracked for writing.
|
||||||
|
|
@ -519,7 +517,8 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize
|
||||||
pktHead[item.iphLen+tcpFlagsOffset] |= tcpFlagPSH
|
pktHead[item.iphLen+tcpFlagsOffset] |= tcpFlagPSH
|
||||||
}
|
}
|
||||||
extendBy := len(pkt) - int(headersLen)
|
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:])
|
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
|
// 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
|
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
|
// 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
|
return groResultTableInsert
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue