Compare commits

...

10 commits

Author SHA1 Message Date
Leadaxe
d31d20ba58 system stack: self-heal the TCP forwarder accept loop (sing-box-lx SPEC 040)
Upstream acceptLoop treats any Accept error as terminal and silently
returns, leaving the stack alive but every new TCP SYN NAT-rewritten onto
a dead port (instant RST) until a full restart. When the listener fd is
closed out from under the stack (a stray close on a reused fd number from
another runtime in the same process), all new TCP dies forever while
UDP/QUIC/DNS keep working.

- System.Close() now marks a deliberate shutdown first; acceptLoop still
  exits quietly on it.
- Any other Accept error is logged (the errno names the killer path),
  the listener is recreated on the same address, the forwarder port is
  republished atomically, and the loop keeps serving.
- If the rebind fails, the loop logs an error and gives up - no worse
  than upstream.
- acceptRecoveries counter is kept as telemetry.

tcpPort/tcpPort6 become atomic (written by the heal path, read from the
tunLoop dispatch/NAT path); listener replacement is serialized against
Close() with a mutex.
2026-08-05 16:57:11 +03:00
世界
da24acaf4d
Update gvisor to 20260727.0 2026-08-05 08:12:00 +08:00
世界
2d9b8aed5f
Fix deadlock between flow judgement and close 2026-07-29 13:45:28 +08:00
世界
e5c21070ae
Fix flow close race 2026-07-27 23:11:49 +08:00
世界
b59636919c
Add Stack.ResetNetwork 2026-07-27 23:11:49 +08:00
世界
79084fa798
Add stateless DNS hijack 2026-07-27 23:11:49 +08:00
世界
1ba7d79118
Add UDPEgressPool 2026-07-27 23:11:49 +08:00
世界
95bc107a1c
refactor: New udpnat 2026-07-27 23:11:49 +08:00
世界
d1af8aaf7e
Fix lint errors 2026-07-27 23:11:48 +08:00
世界
994d6ccdbf
Add netns support 2026-07-27 23:11:48 +08:00
31 changed files with 2390 additions and 264 deletions

View file

@ -21,6 +21,7 @@ const (
ActionReject
ActionDrop
ActionBypass
ActionHijackDNS
)
type FlowTracker interface {

View file

@ -3,6 +3,7 @@ package tun
import (
"maps"
"net/netip"
"sync"
"sync/atomic"
"time"
@ -113,9 +114,11 @@ type ForwardDispatcher struct {
logger logger.Logger
udpTimeout time.Duration
icmpTimeout time.Duration
access sync.RWMutex
table map[flowKey]*flowEntry
lastSweep int64
resetPending atomic.Bool
ports map[Port]*portNAT
natList atomic.Pointer[[]*portNAT]
revNAT atomic.Pointer[map[netip.Addr]*portNAT]
@ -167,26 +170,41 @@ func (d *ForwardDispatcher) Close() {
return
}
d.returnPath.closed.Store(true)
d.access.Lock()
flows := make([]*forwardFlow, 0, len(d.table))
for _, entry := range d.table {
if entry.flow != nil {
entry.flow.close(FlowCloseReset)
flows = append(flows, entry.flow)
}
}
ports := make([]Port, 0, len(d.ports))
for port, nat := range d.ports {
if nat != nil {
port.DetachReturn(&d.returnPath)
ports = append(ports, port)
}
}
d.access.Unlock()
for _, flow := range flows {
flow.close(FlowCloseReset)
}
for _, port := range ports {
port.DetachReturn(&d.returnPath)
}
}
func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
if d == nil {
if d == nil || d.returnPath.closed.Load() {
return false
}
parsed, ok := parseForwardPacket(packet)
if !ok || parsed.fragment || !parsed.hasFlow {
return false
}
d.access.RLock()
if d.returnPath.closed.Load() {
d.access.RUnlock()
return false
}
key := parsed.flowKey()
now := d.now()
entry, loaded := d.table[key]
@ -195,13 +213,16 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
loaded = false
}
if loaded {
return d.handleHit(key, entry, &parsed, packet, now)
handled := d.handleHit(key, entry, &parsed, packet, now)
d.access.RUnlock()
return handled
}
d.access.RUnlock()
if parsed.protocol == uint8(header.TCPProtocolNumber) &&
(parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) {
return false
}
return d.judgeAndInstall(key, &parsed, packet, now)
return d.judgeAndInstall(key, &parsed, packet)
}
func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool {
@ -255,12 +276,18 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
}
}
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte) bool {
var firstPacket []byte
if packet.protocol == uint8(header.UDPProtocolNumber) {
firstPacket = header.UDP(packet.transport).Payload()
}
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination, firstPacket)
d.access.RLock()
defer d.access.RUnlock()
if d.returnPath.closed.Load() {
return false
}
now := d.now()
switch verdict.Action {
case ActionFlow:
if verdict.Port != nil {
@ -291,6 +318,13 @@ func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket,
case ActionDrop:
d.installSimple(key, ActionDrop, packet.protocol, now)
return true
case ActionHijackDNS:
if packet.protocol == uint8(header.UDPProtocolNumber) {
d.hijackDNSPacket(packet)
return true
}
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
default:
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
@ -537,10 +571,27 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) {
}
}
func (d *ForwardDispatcher) Flush() {
func (d *ForwardDispatcher) ResetNetwork() {
if d == nil {
return
}
d.resetPending.Store(true)
}
func (d *ForwardDispatcher) Flush() {
if d == nil || d.returnPath.closed.Load() {
return
}
d.access.RLock()
defer d.access.RUnlock()
if d.returnPath.closed.Load() {
return
}
if d.resetPending.Swap(false) {
for key, entry := range d.table {
d.removeEntry(key, entry, FlowCloseReset)
}
}
for _, nat := range d.activeNATs {
d.flushPort(nat)
}

89
flow_dns.go Normal file
View file

@ -0,0 +1,89 @@
package tun
import (
"net/netip"
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
func (d *ForwardDispatcher) hijackDNSPacket(packet *forwardPacket) {
writer := &dnsResponseWriter{
writeback: d.writeback,
source: packet.source,
}
d.handler.NewDNSPacket(header.UDP(packet.transport).Payload(), M.SocksaddrFromNetIP(packet.source), M.SocksaddrFromNetIP(packet.destination), writer)
}
var _ N.PacketWriter = (*dnsResponseWriter)(nil)
type dnsResponseWriter struct {
writeback ForwardWriteback
source netip.AddrPort
}
func (w *dnsResponseWriter) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
defer buffer.Release()
if !destination.IsIP() {
return E.New("invalid destination: ", destination)
}
sourceAddr := w.source.Addr().Unmap()
destinationAddr := destination.Addr.Unmap()
headroom := w.writeback.ReturnHeadroom()
udpLen := header.UDPMinimumSize + buffer.Len()
var (
packet []byte
udpHdr header.UDP
ipHdr header.Network
)
if sourceAddr.Is4() {
if !destinationAddr.Is4() {
return E.New("send IPv6 packet to IPv4 connection")
}
size := header.IPv4MinimumSize + udpLen
packet = make([]byte, headroom+size)
inet4Hdr := header.IPv4(packet[headroom:])
inet4Hdr.Encode(&header.IPv4Fields{
TotalLength: uint16(size),
TTL: synthesizedTTL,
Protocol: uint8(header.UDPProtocolNumber),
SrcAddr: destinationAddr,
DstAddr: sourceAddr,
})
udpHdr = header.UDP(inet4Hdr.Payload())
ipHdr = inet4Hdr
} else {
if destinationAddr.Is4() {
destinationAddr = netip.AddrFrom16(destinationAddr.As16())
}
size := header.IPv6MinimumSize + udpLen
packet = make([]byte, headroom+size)
inet6Hdr := header.IPv6(packet[headroom:])
inet6Hdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(udpLen),
TransportProtocol: header.UDPProtocolNumber,
HopLimit: synthesizedTTL,
SrcAddr: destinationAddr,
DstAddr: sourceAddr,
})
udpHdr = header.UDP(inet6Hdr.Payload())
ipHdr = inet6Hdr
}
udpHdr.Encode(&header.UDPFields{
SrcPort: destination.Port,
DstPort: w.source.Port(),
Length: uint16(udpLen),
})
copy(udpHdr.Payload(), buffer.Bytes())
udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum(
header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), uint16(udpLen)),
)))
if inet4Hdr, isInet4 := ipHdr.(header.IPv4); isInet4 {
inet4Hdr.SetChecksum(^inet4Hdr.CalculateChecksum())
}
return w.writeback.WriteReturnPackets([][]byte{packet})
}

28
go.mod
View file

@ -1,32 +1,32 @@
module github.com/sagernet/sing-tun
go 1.24.7
go 1.25.0
require (
github.com/florianl/go-nfqueue/v2 v2.0.2
github.com/florianl/go-nfqueue/v2 v2.1.0
github.com/go-ole/go-ole v1.3.0
github.com/google/btree v1.1.3
github.com/mdlayher/netlink v1.9.0
github.com/sagernet/fswatch v0.1.1
github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1
github.com/mdlayher/netlink v1.11.2
github.com/sagernet/fswatch v0.1.2
github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1
github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a
github.com/sagernet/nftables v0.3.0-mod.2
github.com/sagernet/sing v0.8.0
github.com/sagernet/nftables v0.3.0-mod.4
github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8
github.com/stretchr/testify v1.11.1
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8
golang.org/x/net v0.50.0
golang.org/x/sys v0.41.0
golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc
golang.org/x/net v0.57.0
golang.org/x/sys v0.47.0
)
require (
github.com/davecgh/go-spew v1.1.1 // indirect
github.com/fsnotify/fsnotify v1.7.0 // indirect
github.com/fsnotify/fsnotify v1.9.0 // indirect
github.com/google/go-cmp v0.7.0 // indirect
github.com/mdlayher/socket v0.5.1 // indirect
github.com/mdlayher/socket v0.6.0 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/vishvananda/netns v0.0.4 // indirect
golang.org/x/sync v0.7.0 // indirect
golang.org/x/time v0.7.0 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/time v0.15.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
)

54
go.sum
View file

@ -1,48 +1,50 @@
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/florianl/go-nfqueue/v2 v2.0.2 h1:FL5lQTeetgpCvac1TRwSfgaXUn0YSO7WzGvWNIp3JPE=
github.com/florianl/go-nfqueue/v2 v2.0.2/go.mod h1:VA09+iPOT43OMoCKNfXHyzujQUty2xmzyCRkBOlmabc=
github.com/fsnotify/fsnotify v1.7.0 h1:8JEhPFa5W2WU7YfeZzPNqzMP6Lwt7L2715Ggo0nosvA=
github.com/fsnotify/fsnotify v1.7.0/go.mod h1:40Bi/Hjc2AVfZrqy+aj+yEI+/bRxZnMJyTJwOpGvigM=
github.com/florianl/go-nfqueue/v2 v2.1.0 h1:Fywt30TY/evxyDySpXjxQ1jsRW7nQbLpOhELqpr4068=
github.com/florianl/go-nfqueue/v2 v2.1.0/go.mod h1:8PKUM5rYoVFO5IZV1bifx4/b0jHAglKkHXr9PRwzi4Y=
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE=
github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/mdlayher/netlink v1.9.0 h1:G8+GLq2x3v4D4MVIqDdNUhTUC7TKiCy/6MDkmItfKco=
github.com/mdlayher/netlink v1.9.0/go.mod h1:YBnl5BXsCoRuwBjKKlZ+aYmEoq0r12FDA/3JC+94KDg=
github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos=
github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ=
github.com/jsimonetti/rtnetlink/v2 v2.2.0 h1:/KfZ310gOAFrXXol5VwnFEt+ucldD/0dsSRZwpHCP9w=
github.com/jsimonetti/rtnetlink/v2 v2.2.0/go.mod h1:lbjDHxC+5RJ08lzPeA90Ls2pEoId3F08MoEMlhfHxeI=
github.com/mdlayher/netlink v1.11.2 h1:HKh2jqe+omdSWcQ88nrT7INE61B0NXfiSPFdgL4YbNI=
github.com/mdlayher/netlink v1.11.2/go.mod h1:uT2Yc/QLaZubzDpZIBi9d4GoeLwtp3x1AMeqSRrK2sA=
github.com/mdlayher/socket v0.6.0 h1:ScZPaAGyO1icQnbFrhPM8mnXyMu9qukC1K4ZoM2IQKU=
github.com/mdlayher/socket v0.6.0/go.mod h1:q7vozUAnxSqnjHc12Fik5yUKIzfZ8ITCfMkhOtE9z18=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/sagernet/fswatch v0.1.1 h1:YqID+93B7VRfqIH3PArW/XpJv5H4OLEVWDfProGoRQs=
github.com/sagernet/fswatch v0.1.1/go.mod h1:nz85laH0mkQqJfaOrqPpkwtU1znMFNVTpT/5oRsVz/o=
github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 h1:AzCE2RhBjLJ4WIWc/GejpNh+z30d5H1hwaB0nD9eY3o=
github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1/go.mod h1:NJKBtm9nVEK3iyOYWsUlrDQuoGh4zJ4KOPhSYVidvQ4=
github.com/sagernet/fswatch v0.1.2 h1:/TT7k4mkce1qFPxamLO842WjqBgbTBiXP2mlUjp9PFk=
github.com/sagernet/fswatch v0.1.2/go.mod h1:5BpGmpUQVd3Mc5r313HRpvADHRg3/rKn5QbwFteB880=
github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1 h1:IdQ7yTKkB2wv8txwshxUroPlO4npOYAV71xb7xQ7Lys=
github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1/go.mod h1:9O3SQskYuCfdHNvHEsWuEAgoyKEF74PiWp4NsNUia8g=
github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZNjr6sGeT00J8uU7JF4cNUdb44/Duis=
github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a/go.mod h1:xLnfdiJbSp8rNqYEdIW/6eDO4mVoogml14Bh2hSiFpM=
github.com/sagernet/nftables v0.3.0-mod.2 h1:ck2KMU02OxL1eDFgGaWYglMDpoOZ7OHzxje+vW5Q0OQ=
github.com/sagernet/nftables v0.3.0-mod.2/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ=
github.com/sagernet/sing v0.8.0 h1:OwLEwbcYfZHvu4olZVljxxC1XRicBqJ1HfiFr6F2WEE=
github.com/sagernet/sing v0.8.0/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
github.com/sagernet/nftables v0.3.0-mod.4 h1:vnOtcDYeSXv2e5RoRuGH0lrpttQFJ8iC4ICS2nhlDSo=
github.com/sagernet/nftables v0.3.0-mod.4/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ=
github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8 h1:dyRIj+MZ2rc9JVzJoG04jxu+MpvHrLIZLJr0QjNAMGg=
github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 h1:yixxcjnhBmY0nkL253HFVIm0JsFHwrHdT3Yh6szTnfY=
golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8/go.mod h1:jj3sYF3dwk5D+ghuXyeI3r5MFf+NT2An6/9dOA95KSI=
golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
golang.org/x/sync v0.7.0 h1:YsImfSBoP9QPYL0xyKJPq0gcaJdG3rInoqxTWbfQu9M=
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc h1:TS73t7x3KarrNd5qAipmspBDS1rkMcgVG/fS1aRb4Rc=
golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc/go.mod h1:A+z0yzpGtvnG90cToK5n2tu8UJVP2XUATh+r+sfOOOc=
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ=
golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=

View file

@ -22,7 +22,6 @@ import (
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing/common"
)
// RFC 971 defines the fields of the IPv4 header on page 11 using the following
@ -335,7 +334,7 @@ func (b IPv4) FragmentOffset() uint16 {
}
func (b IPv4) FragmentOffsetDarwinRaw() uint16 {
return common.NativeEndian.Uint16(b[flagsFO:]) << 3
return binary.NativeEndian.Uint16(b[flagsFO:]) << 3
}
// TotalLength returns the "total length" field of the IPv4 header.
@ -344,7 +343,7 @@ func (b IPv4) TotalLength() uint16 {
}
func (b IPv4) TotalLengthDarwinRaw() uint16 {
return common.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength())
return binary.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength())
}
// Checksum returns the checksum field of the IPv4 header.
@ -441,7 +440,7 @@ func (b IPv4) SetTotalLength(totalLength uint16) {
}
func (b IPv4) SetTotalLengthDarwinRaw(totalLength uint16) {
common.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength)
binary.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength)
}
// SetChecksum sets the checksum field of the IPv4 header.
@ -458,7 +457,7 @@ func (b IPv4) SetFlagsFragmentOffset(flags uint8, offset uint16) {
func (b IPv4) SetFlagsFragmentOffsetDarwinRaw(flags uint8, offset uint16) {
v := (uint16(flags) << 13) | (offset >> 3)
common.NativeEndian.PutUint16(b[flagsFO:], v)
binary.NativeEndian.PutUint16(b[flagsFO:], v)
}
// SetID sets the identification field.
@ -1179,7 +1178,7 @@ func (s IPv4OptionsSerializer) Serialize(b []byte) uint8 {
// header ends on a 32 bit boundary. The padding is zero.
padded := padIPv4OptionsLength(total)
b = b[:padded-total]
common.ClearArray(b)
clear(b)
return padded
}

View file

@ -21,7 +21,6 @@ import (
"math"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing/common"
)
// IPv6ExtensionHeaderIdentifier is an IPv6 extension header identifier.
@ -129,7 +128,7 @@ func padIPv6Option(b []byte) {
b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6Pad1ExtHdrOptionIdentifier)
default: // Pad with PadN.
s := b[ipv6ExtHdrOptionPayloadOffset:]
common.ClearArray(s)
clear(s)
b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6PadNExtHdrOptionIdentifier)
b[ipv6ExtHdrOptionLengthOffset] = uint8(len(s))
}

View file

@ -24,7 +24,6 @@ import (
"time"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing/common"
)
// ndpOptionIdentifier is an NDP option type identifier.
@ -341,7 +340,7 @@ func (b NDPOptions) Serialize(s NDPOptionsSerializer) int {
// Zero out remaining (padding) bytes, if any exists.
if used+2 < l {
common.ClearArray(b[used+2 : l])
clear(b[used+2 : l])
}
b = b[l:]
@ -567,7 +566,7 @@ func (o NDPPrefixInformation) serializeInto(b []byte) int {
// Zero out the Reserved2 field.
reserved2 := b[ndpPrefixInformationReserved2Offset:][:ndpPrefixInformationReserved2Length]
common.ClearArray(reserved2)
clear(reserved2)
return used
}
@ -686,7 +685,7 @@ func (o NDPRecursiveDNSServer) serializeInto(b []byte) int {
used := copy(b, o)
// Zero out the reserved bytes that are before the Lifetime field.
common.ClearArray(b[0:ndpRecursiveDNSServerLifetimeOffset])
clear(b[0:ndpRecursiveDNSServerLifetimeOffset])
return used
}
@ -779,7 +778,7 @@ func (o NDPDNSSearchList) serializeInto(b []byte) int {
used := copy(b, o)
// Zero out the reserved bytes that are before the Lifetime field.
common.ClearArray(b[0:ndpDNSSearchListLifetimeOffset])
clear(b[0:ndpDNSSearchListLifetimeOffset])
return used
}

View file

@ -50,7 +50,6 @@ import (
"github.com/sagernet/gvisor/pkg/tcpip/header"
"github.com/sagernet/gvisor/pkg/tcpip/stack"
rawfile "github.com/sagernet/sing-tun/internal/rawfile_darwin"
"github.com/sagernet/sing/common"
"golang.org/x/sys/unix"
)
@ -200,10 +199,6 @@ type Options struct {
// include CapabilitySaveRestore
SaveRestore bool
// DisconnectOk if true, indicates that this NIC capability set should
// include CapabilityDisconnectOk.
DisconnectOk bool
// PacketDispatchMode specifies the type of inbound dispatcher to be
// used for this endpoint.
PacketDispatchMode PacketDispatchMode
@ -257,10 +252,6 @@ func New(opts *Options) (stack.LinkEndpoint, error) {
caps |= stack.CapabilitySaveRestore
}
if opts.DisconnectOk {
caps |= stack.CapabilityDisconnectOk
}
if len(opts.FDs) == 0 {
return nil, fmt.Errorf("opts.FD is empty, at least one FD must be specified")
}
@ -301,7 +292,7 @@ func New(opts *Options) (stack.LinkEndpoint, error) {
e.fds = append(e.fds, fdInfo{fd: fd, isSocket: true})
if opts.ProcessorsPerChannel == 0 {
opts.ProcessorsPerChannel = common.Max(1, runtime.GOMAXPROCS(0)/len(opts.FDs))
opts.ProcessorsPerChannel = max(1, runtime.GOMAXPROCS(0)/len(opts.FDs))
}
inboundDispatcher, err := newRecvMMsgDispatcher(fd, e, opts)

View file

@ -214,34 +214,47 @@ func tcpipConnectionID(pkt *stack.PacketBuffer) (connectionID, bool) {
return cid, true
}
ipHdr := header.IPv6(h)
cid.srcAddr = ipHdr.SourceAddressSlice()
cid.dstAddr = ipHdr.DestinationAddressSlice()
cid.proto = header.IPv6ProtocolNumber
var tcpHdr header.TCP
if tcpip.TransportProtocolNumber(ipHdr.NextHeader()) == header.TCPProtocolNumber {
tcpHdr = header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen])
if !header.IsExtensionHeader(ipHdr.NextHeader()) {
// Known transport protocols(not just TCP) store the src and dst ports
// in the first 4 bytes after the IPv6 fixed header.
tcpHdr := header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen])
cid.srcPort = tcpHdr.SourcePort()
cid.dstPort = tcpHdr.DestinationPort()
} else {
// Slow path for IPv6 extension headers :(.
dataBuf := pkt.Data().ToBuffer()
dataBuf.TrimFront(header.IPv6MinimumSize)
it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(ipHdr.NextHeader()), dataBuf)
defer it.Release()
// All fragment packets need to be processed by the same goroutine, so
// only record the ports if this is not a fragment packet.
var isFragment bool
for {
hdr, done, err := it.Next()
if done || err != nil {
break
}
if fh, ok := hdr.(header.IPv6FragmentExtHdr); ok && !fh.IsAtomic() {
isFragment = true
}
hdr.Release()
}
if !isFragment {
h, ok = pkt.Data().PullUp(int(it.HeaderOffset()) + tcpSrcDstPortLen)
if !ok {
return cid, true
}
tcpHdr = header.TCP(h[it.HeaderOffset():][:tcpSrcDstPortLen])
}
cid.srcAddr = ipHdr.SourceAddressSlice()
cid.dstAddr = ipHdr.DestinationAddressSlice()
// Known transport protocols store the src and dst ports
// in the first 4 bytes after the IPv6 fixed header.
tcpHdr := header.TCP(h[it.HeaderOffset():][:tcpSrcDstPortLen])
cid.srcPort = tcpHdr.SourcePort()
cid.dstPort = tcpHdr.DestinationPort()
cid.proto = header.IPv6ProtocolNumber
}
}
default:
return cid, true
}

73
netns_linux.go Normal file
View file

@ -0,0 +1,73 @@
package tun
import (
"context"
"net"
"runtime"
"strings"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"golang.org/x/sys/unix"
)
func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) {
return execInNetworkNamespace(nameOrPath, func() (net.Listener, error) {
return config.Listen(ctx, network, address)
})
}
type networkNamespaceInterfaceFinder struct {
control.InterfaceFinder
options *Options
}
func (f *networkNamespaceInterfaceFinder) Update() error {
return runInNetworkNamespace(f.options.NetNs, f.InterfaceFinder.Update)
}
func execInNetworkNamespace[T any](nameOrPath string, block func() (T, error)) (T, error) {
if nameOrPath == "" {
return block()
}
type blockResult struct {
value T
err error
}
resultChannel := make(chan blockResult, 1)
go func() {
runtime.LockOSThread()
value, err := execInNetworkNamespaceThread(nameOrPath, block)
resultChannel <- blockResult{value, err}
}()
result := <-resultChannel
return result.value, result.err
}
func execInNetworkNamespaceThread[T any](nameOrPath string, block func() (T, error)) (T, error) {
var defaultValue T
var path string
if strings.HasPrefix(nameOrPath, "/") {
path = nameOrPath
} else {
path = "/run/netns/" + nameOrPath
}
targetFd, err := unix.Open(path, unix.O_RDONLY|unix.O_CLOEXEC, 0)
if err != nil {
return defaultValue, E.Cause(err, "open netns ", nameOrPath)
}
defer unix.Close(targetFd)
err = unix.Setns(targetFd, unix.CLONE_NEWNET)
if err != nil {
return defaultValue, E.Cause(err, "set netns to ", nameOrPath)
}
return block()
}
func runInNetworkNamespace(nameOrPath string, block func() error) error {
_, err := execInNetworkNamespace(nameOrPath, func() (struct{}, error) {
return struct{}{}, block()
})
return err
}

12
netns_other.go Normal file
View file

@ -0,0 +1,12 @@
//go:build !linux
package tun
import (
"context"
"net"
)
func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) {
return config.Listen(ctx, network, address)
}

View file

@ -1,11 +1,10 @@
package ping
import (
"encoding/binary"
"fmt"
"unsafe"
"github.com/sagernet/sing/common"
"golang.org/x/net/ipv6"
"golang.org/x/sys/windows"
)
@ -37,9 +36,9 @@ func parseIPv6ControlMessage(cmsg []byte) (*ipv6.ControlMessage, error) {
}
switch cmsghdr.Type {
case IPV6_TCLASS:
controlMessage.TrafficClass = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
controlMessage.TrafficClass = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
case IPV6_HOPLIMIT:
controlMessage.HopLimit = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
controlMessage.HopLimit = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
}
cmsg = cmsg[msgSize:]
}

View file

@ -9,7 +9,6 @@ import (
"time"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
M "github.com/sagernet/sing/common/metadata"
@ -175,7 +174,7 @@ func (c *UnprivilegedConn) Close() error {
for _, conn := range c.mapping {
_ = conn.Close()
}
common.ClearMap(c.mapping)
clear(c.mapping)
return nil
}

View file

@ -26,6 +26,7 @@ type autoRedirect struct {
logger logger.Logger
tableName string
networkMonitor NetworkUpdateMonitor
ownedNetworkMonitor bool
networkListener *list.Element[NetworkUpdateCallback]
interfaceFinder control.InterfaceFinder
localAddresses []netip.Prefix
@ -51,7 +52,7 @@ type autoRedirect struct {
}
func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) {
return &autoRedirect{
r := &autoRedirect{
tunOptions: options.TunOptions,
ctx: options.Context,
handler: options.Handler,
@ -63,7 +64,11 @@ func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) {
customRedirectPortFunc: options.CustomRedirectPort,
routeAddressSet: options.RouteAddressSet,
routeExcludeAddressSet: options.RouteExcludeAddressSet,
}, nil
}
if options.TunOptions.NetNs != "" {
r.interfaceFinder = &networkNamespaceInterfaceFinder{control.NewDefaultInterfaceFinder(), options.TunOptions}
}
return r, nil
}
func (r *autoRedirect) Start() error {
@ -89,8 +94,11 @@ func (r *autoRedirect) Start() error {
}
}
} else {
if r.tunOptions.NetNs != "" && !r.useNFTables {
return E.New("auto_redirect in network namespace requires nftables")
}
if r.useNFTables {
err = r.initializeNFTables()
err = runInNetworkNamespace(r.tunOptions.NetNs, r.initializeNFTables)
if err != nil {
return E.Cause(err, "missing nftables support")
}
@ -132,7 +140,7 @@ func (r *autoRedirect) Start() error {
listenAddr = netip.IPv4Unspecified()
}
server := newRedirectServer(r.ctx, r.handler, r.logger, listenAddr)
err = server.Start()
err = runInNetworkNamespace(r.tunOptions.NetNs, server.Start)
if err != nil {
return E.Cause(err, "start redirect server")
}
@ -151,25 +159,44 @@ func (r *autoRedirect) Start() error {
})
if err != nil {
r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
} else if err = handler.Start(); err != nil {
} else if err = runInNetworkNamespace(r.tunOptions.NetNs, handler.Start); err != nil {
r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
} else {
r.nfqueueHandler = handler
r.nfqueueEnabled = true
}
}
r.cleanupNFTables()
err = r.setupNFTables()
if r.tunOptions.NetNs != "" {
var monitor NetworkUpdateMonitor
monitor, err = NewNetworkUpdateMonitor(r.logger)
if err != nil {
return E.Cause(err, "setup nftables")
return E.Cause(err, "create netns network monitor")
}
err = runInNetworkNamespace(r.tunOptions.NetNs, monitor.Start)
if err != nil {
return E.Cause(err, "start netns network monitor")
}
r.networkMonitor = monitor
r.ownedNetworkMonitor = true
}
err = runInNetworkNamespace(r.tunOptions.NetNs, func() error {
r.cleanupNFTables()
setupErr := r.setupNFTables()
if setupErr != nil {
return E.Cause(setupErr, "setup nftables")
}
if r.tunOptions.AutoRedirectMarkMode {
err = r.setupRedirectRoutes()
if err != nil {
setupErr = r.setupRedirectRoutes()
if setupErr != nil {
r.cleanupNFTables()
return E.Cause(err, "setup redirect routes")
return E.Cause(setupErr, "setup redirect routes")
}
}
return nil
})
if err != nil {
return err
}
} else {
r.cleanupIPTables()
err = r.setupIPTables()
@ -185,8 +212,14 @@ func (r *autoRedirect) Close() error {
r.nfqueueHandler.Close()
}
if r.useNFTables {
_ = runInNetworkNamespace(r.tunOptions.NetNs, func() error {
r.cleanupNFTables()
r.cleanupRedirectRoutes()
return nil
})
if r.ownedNetworkMonitor {
_ = r.networkMonitor.Close()
}
} else {
r.cleanupIPTables()
}
@ -197,7 +230,7 @@ func (r *autoRedirect) Close() error {
func (r *autoRedirect) UpdateRouteAddressSet() {
if r.useNFTables {
err := r.nftablesUpdateRouteAddressSet()
err := runInNetworkNamespace(r.tunOptions.NetNs, r.nftablesUpdateRouteAddressSet)
if err != nil {
r.logger.Error("update route address set: ", err)
}

View file

@ -299,27 +299,38 @@ func (r *autoRedirect) setupNFTables() error {
if err != nil {
return E.Cause(err, "flush nftables")
}
if r.tunOptions.NetNs == "" {
r.startDockerFirewallMonitor()
err = r.configureDockerFirewall(false)
if err != nil && r.logger != nil {
r.logger.Warn("configure docker firewall: ", err)
}
}
r.networkListener = r.networkMonitor.RegisterCallback(func() {
err = r.nftablesUpdateLocalAddressSet()
if err != nil {
r.logger.Error("update local address set: ", err)
}
if r.tunOptions.AutoRedirectMarkMode {
err = r.updateRedirectRoutes()
if err != nil {
r.logger.Error("update redirect routes: ", err)
}
updateErr := runInNetworkNamespace(r.tunOptions.NetNs, r.updateNetworkAddresses)
if updateErr != nil {
r.logger.Error(updateErr)
}
})
return nil
}
func (r *autoRedirect) updateNetworkAddresses() error {
err := r.nftablesUpdateLocalAddressSet()
if err != nil {
err = E.Cause(err, "update local address set")
}
if r.tunOptions.AutoRedirectMarkMode {
routeErr := r.updateRedirectRoutes()
if routeErr != nil {
routeErr = E.Cause(routeErr, "update redirect routes")
}
err = E.Errors(err, routeErr)
}
return err
}
// TODO: test if this works
func (r *autoRedirect) nftablesUpdateLocalAddressSet() error {
err := r.interfaceFinder.Update()
@ -376,6 +387,7 @@ func (r *autoRedirect) nftablesUpdateRouteAddressSet() error {
func (r *autoRedirect) cleanupNFTables() {
if r.networkListener != nil {
r.networkMonitor.UnregisterCallback(r.networkListener)
r.networkListener = nil
}
r.stopDockerFirewallMonitor()
nft, err := nftables.New()
@ -389,10 +401,12 @@ func (r *autoRedirect) cleanupNFTables() {
_ = r.configureOpenWRTFirewall4(nft, true)
_ = nft.Flush()
_ = nft.CloseLasting()
if r.tunOptions.NetNs == "" {
err = r.configureDockerFirewall(true)
if err != nil && r.logger != nil {
r.logger.Warn("cleanup docker firewall: ", err)
}
}
}
func (r *autoRedirect) nftablesCreatePreMatchChains(nft *nftables.Conn, table *nftables.Table) error {

View file

@ -14,6 +14,7 @@ import (
type Stack interface {
Start() error
ResetNetwork()
Close() error
}
@ -23,6 +24,9 @@ type StackOptions struct {
TunOptions Options
UDPTimeout time.Duration
ICMPTimeout time.Duration
UDPMapping NATMapping
UDPFiltering NATFiltering
UDPNATMax uint32
Handler Handler
Logger logger.Logger
ForwarderBindInterface bool

View file

@ -35,8 +35,8 @@ type GVisor struct {
inet6Address netip.Addr
inet4LoopbackAddress []netip.Addr
inet6LoopbackAddress []netip.Addr
udpTimeout time.Duration
icmpTimeout time.Duration
udpNATOptions UDPNatOptions
broadcastAddr netip.Addr
handler Handler
logger logger.Logger
@ -44,6 +44,7 @@ type GVisor struct {
endpoint stack.LinkEndpoint
dispatcher *ForwardDispatcher
icmpForwarder *ICMPForwarder
udpForwarder *UDPForwarder
}
type GVisorTun interface {
@ -78,8 +79,15 @@ func NewGVisor(
inet6Address: inet6Address,
inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress,
inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress,
udpTimeout: options.UDPTimeout,
icmpTimeout: options.ICMPTimeout,
udpNATOptions: UDPNatOptions{
Timeout: options.UDPTimeout,
Mapping: options.UDPMapping,
Filtering: options.UDPFiltering,
MaxSize: options.UDPNATMax,
InterfaceFinder: options.InterfaceFinder,
ExcludeInterface: []string{options.TunOptions.Name},
},
broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address),
handler: options.Handler,
logger: options.Logger,
@ -93,7 +101,7 @@ func (t *GVisor) Start() error {
return err
}
if t.handler != nil {
t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout)
t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpNATOptions.Timeout, t.icmpTimeout)
}
linkEndpoint = &LinkEndpointFilter{
LinkEndpoint: linkEndpoint,
@ -110,7 +118,13 @@ func (t *GVisor) Start() error {
return err
}
ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket)
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket)
udpForwarder := NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpNATOptions)
err = udpForwarder.Start()
if err != nil {
return err
}
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, udpForwarder.HandlePacket)
t.udpForwarder = udpForwarder
icmpForwarder := NewICMPForwarder(ipStack, t.handler, t.logger)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
@ -120,11 +134,24 @@ func (t *GVisor) Start() error {
return nil
}
func (t *GVisor) ResetNetwork() {
if t.udpForwarder != nil {
t.udpForwarder.udpNat.Purge()
}
if t.icmpForwarder != nil {
t.icmpForwarder.Purge()
}
t.dispatcher.ResetNetwork()
}
func (t *GVisor) Close() error {
t.dispatcher.Close()
if t.icmpForwarder != nil {
t.icmpForwarder.Close()
}
if t.udpForwarder != nil {
t.udpForwarder.Close()
}
if t.stack == nil {
return nil
}

View file

@ -72,6 +72,15 @@ func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger)
return forwarder
}
func (f *ICMPForwarder) Purge() {
f.flowAccess.Lock()
for key, flow := range f.flows {
flow.close(FlowCloseReset)
delete(f.flows, key)
}
f.flowAccess.Unlock()
}
func (f *ICMPForwarder) Close() error {
f.returnPath.closed.Store(true)
f.flowAccess.Lock()
@ -146,9 +155,15 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
} else {
ipHdr := header.IPv6(pkt.NetworkHeader().Slice())
icmpHdr := header.ICMPv6(pkt.TransportHeader().Slice())
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
if icmpHdr.Type() != header.ICMPv6EchoRequest {
return false
}
if icmpHdr.Code() != 0 {
// The IPv6 built-in echo reply path lacks the LocalAddressTemporary
// check its IPv4 sibling has, so returning false would make the stack
// reply on behalf of arbitrary forwarded destinations.
return true
}
identifier := icmpHdr.Ident()
key := icmpFlowKey{
v6: true,

View file

@ -8,7 +8,6 @@ import (
"net/netip"
"os"
"sync"
"time"
_ "unsafe"
"github.com/sagernet/gvisor/pkg/buffer"
@ -21,26 +20,35 @@ import (
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/udpnat2"
)
type UDPForwarder struct {
ctx context.Context
stack *stack.Stack
handler Handler
udpNat *udpnat.Service
udpNat *UDPNat
}
func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, timeout time.Duration) *UDPForwarder {
func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, options UDPNatOptions) *UDPForwarder {
forwarder := &UDPForwarder{
ctx: ctx,
stack: stack,
handler: handler,
}
forwarder.udpNat = udpnat.New(handler, forwarder.PreparePacketConnection, timeout, false)
options.Handler = handler
options.Prepare = forwarder.PreparePacketConnection
forwarder.udpNat = NewUDPNat(options)
return forwarder
}
func (f *UDPForwarder) Start() error {
return f.udpNat.Start()
}
func (f *UDPForwarder) Close() error {
return f.udpNat.Close()
}
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)
@ -63,18 +71,26 @@ func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M
firstPacket = append(firstPacket[:len(firstPacket):len(firstPacket)], view.AsSlice()...)
}
})
var sourceNetwork tcpip.NetworkProtocolNumber
if source.Addr.Is4() {
sourceNetwork = header.IPv4ProtocolNumber
} else {
sourceNetwork = header.IPv6ProtocolNumber
}
switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort(), firstPacket).Action {
case ActionReject:
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
return false, nil, nil, nil
case ActionDrop:
return false, nil, nil, nil
}
var sourceNetwork tcpip.NetworkProtocolNumber
if source.Addr.Is4() {
sourceNetwork = header.IPv4ProtocolNumber
} else {
sourceNetwork = header.IPv6ProtocolNumber
case ActionHijackDNS:
f.handler.NewDNSPacket(firstPacket, source, destination, &UDPBackWriter{
stack: f.stack,
source: AddressFromAddr(source.Addr),
sourcePort: source.Port,
sourceNetwork: sourceNetwork,
})
return false, nil, nil, nil
}
writer := &UDPBackWriter{
stack: f.stack,

View file

@ -22,6 +22,7 @@ type Mixed struct {
tun GVisorTun
stack *stack.Stack
endpoint *channel.Endpoint
udpForwarder *UDPForwarder
}
func NewMixed(
@ -47,7 +48,13 @@ func (m *Mixed) Start() error {
if err != nil {
return err
}
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpTimeout).HandlePacket)
udpForwarder := NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpNATOptions)
err = udpForwarder.Start()
if err != nil {
return err
}
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, udpForwarder.HandlePacket)
m.udpForwarder = udpForwarder
m.stack = ipStack
m.endpoint = endpoint
go m.tunLoop()
@ -55,10 +62,20 @@ func (m *Mixed) Start() error {
return nil
}
func (m *Mixed) ResetNetwork() {
m.System.ResetNetwork()
if m.udpForwarder != nil {
m.udpForwarder.udpNat.Purge()
}
}
func (m *Mixed) Close() error {
if m.stack == nil {
return nil
}
if m.udpForwarder != nil {
m.udpForwarder.Close()
}
m.endpoint.Attach(nil)
m.stack.Close()
for _, endpoint := range m.stack.CleanupEndpoints() {

View file

@ -5,7 +5,10 @@ import (
"errors"
"net"
"net/netip"
"os"
"slices"
"sync"
"sync/atomic"
"syscall"
"time"
@ -19,7 +22,6 @@ import (
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/udpnat2"
)
var ErrIncludeAllNetworks = E.New("`system` and `mixed` stack are not available when `includeAllNetworks` is enabled. See https://github.com/SagerNet/sing-tun/issues/25")
@ -28,6 +30,7 @@ type System struct {
ctx context.Context
tun Tun
tunName string
netNs string
mtu int
handler Handler
logger logger.Logger
@ -44,10 +47,19 @@ type System struct {
icmpTimeout time.Duration
tcpListener net.Listener
tcpListener6 net.Listener
tcpPort uint16
tcpPort6 uint16
// lx/040: ports are written by acceptLoop on self-heal relisten and read
// concurrently from the tunLoop path (dispatch filter + NAT rewrite) —
// they must be atomic. listenAccess serializes listener replacement
// against Close(); closing marks a deliberate shutdown so acceptLoop can
// tell it apart from the listener dying out from under the stack.
tcpPort atomic.Uint32
tcpPort6 atomic.Uint32
closing atomic.Bool
listenAccess sync.Mutex
acceptRecoveries atomic.Uint32
tcpNat *TCPNat
udpNat *udpnat.Service
udpNat *UDPNat
udpNATOptions UDPNatOptions
dispatcher *ForwardDispatcher
bindInterface bool
interfaceFinder control.InterfaceFinder
@ -68,6 +80,7 @@ func NewSystem(options StackOptions) (Stack, error) {
ctx: options.Context,
tun: options.Tun,
tunName: options.TunOptions.Name,
netNs: options.TunOptions.NetNs,
mtu: int(options.TunOptions.MTU),
inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress,
inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress,
@ -78,6 +91,14 @@ func NewSystem(options StackOptions) (Stack, error) {
inet4Prefixes: options.TunOptions.Inet4Address,
inet6Prefixes: options.TunOptions.Inet6Address,
broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address),
udpNATOptions: UDPNatOptions{
Timeout: options.UDPTimeout,
Mapping: options.UDPMapping,
Filtering: options.UDPFiltering,
MaxSize: options.UDPNATMax,
InterfaceFinder: options.InterfaceFinder,
ExcludeInterface: []string{options.TunOptions.Name},
},
bindInterface: options.ForwarderBindInterface,
interfaceFinder: options.InterfaceFinder,
multiPendingPackets: options.TunOptions.EXP_MultiPendingPackets,
@ -102,8 +123,26 @@ func NewSystem(options StackOptions) (Stack, error) {
return stack, nil
}
func (s *System) ResetNetwork() {
if s.tcpNat != nil {
s.tcpNat.Purge()
}
if s.udpNat != nil {
s.udpNat.Purge()
}
s.dispatcher.ResetNetwork()
}
func (s *System) Close() error {
// lx/040: mark the deliberate shutdown BEFORE closing the listeners so
// acceptLoop exits quietly instead of treating it as a foreign kill.
s.closing.Store(true)
s.dispatcher.Close()
if s.udpNat != nil {
s.udpNat.Close()
}
s.listenAccess.Lock()
defer s.listenAccess.Unlock()
return common.Close(
s.tcpListener,
s.tcpListener6,
@ -119,8 +158,10 @@ func (s *System) Start() error {
return nil
}
func (s *System) start() error {
_ = fixWindowsFirewall()
// lx/040: TCP forwarder bind, shared by start() and the acceptLoop self-heal
// relisten path. isIPv6 selects the address family; the bind-to-interface
// Control and the EADDRNOTAVAIL retry loop match the original start() code.
func (s *System) listenTCP(isIPv6 bool) (net.Listener, error) {
var listener net.ListenConfig
if s.bindInterface {
listener.Control = control.Append(listener.Control, func(network, address string, conn syscall.RawConn) error {
@ -131,40 +172,60 @@ func (s *System) start() error {
return nil
})
}
var tcpListener net.Listener
var err error
if s.inet4NextAddress.IsValid() {
network := "tcp4"
address := s.inet4Address
if isIPv6 {
network = "tcp6"
address = s.inet6Address
}
var (
tcpListener net.Listener
err error
)
for range 3 {
tcpListener, err = listener.Listen(s.ctx, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0"))
tcpListener, err = listenNetworkNamespace(s.ctx, s.netNs, listener, network, net.JoinHostPort(address.String(), "0"))
if !retryableListenError(err) {
break
}
time.Sleep(time.Second)
}
if err != nil {
return nil, err
}
return tcpListener, nil
}
func (s *System) start() error {
_ = fixWindowsFirewall()
var tcpListener net.Listener
var err error
if s.inet4NextAddress.IsValid() {
tcpListener, err = s.listenTCP(false)
if err != nil {
return err
}
s.tcpListener = tcpListener
s.tcpPort = M.SocksaddrFromNet(tcpListener.Addr()).Port
go s.acceptLoop(tcpListener)
s.tcpPort.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port))
go s.acceptLoop(tcpListener, false)
}
if s.inet6NextAddress.IsValid() {
for range 3 {
tcpListener, err = listener.Listen(s.ctx, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0"))
if !retryableListenError(err) {
break
}
time.Sleep(time.Second)
}
tcpListener, err = s.listenTCP(true)
if err != nil {
return err
}
s.tcpListener6 = tcpListener
s.tcpPort6 = M.SocksaddrFromNet(tcpListener.Addr()).Port
go s.acceptLoop(tcpListener)
s.tcpPort6.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port))
go s.acceptLoop(tcpListener, true)
}
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false)
udpNATOptions := s.udpNATOptions
udpNATOptions.Handler = s.handler
udpNATOptions.Prepare = s.preparePacketConnection
s.udpNat = NewUDPNat(udpNATOptions)
err = s.udpNat.Start()
if err != nil {
return err
}
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
@ -336,12 +397,29 @@ func (s *System) processPacket(packet []byte) bool {
return writeBack
}
func (s *System) acceptLoop(listener net.Listener) {
func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) {
for {
conn, err := listener.Accept()
if err != nil {
// lx/040 (SPECS/TASKS/040): upstream silently returns on ANY Accept
// error, leaving the stack alive but every new TCP SYN NAT-rewritten
// onto a dead port (instant RST) until a VPN restart — the LxBox §047
// "browser dead, QUIC alive" failure. A deliberate System.Close is the
// only quiet exit; anything else means the listener died out from
// under us (e.g. a foreign close on a reused fd number from the
// Java side of the shared Android process) — log it (the errno names
// the killer) and recreate the listener.
if s.closing.Load() {
return
}
newListener, healErr := s.healListener(listener, isIPv6, err)
if healErr != nil {
s.logger.Error("system stack: tcp", ipVersionSuffix(isIPv6), " accept loop died: ", err, "; relisten failed: ", healErr)
return
}
listener = newListener
continue
}
connPort := M.SocksaddrFromNet(conn.RemoteAddr()).Port
session := s.tcpNat.LookupBack(connPort)
if session == nil {
@ -352,6 +430,47 @@ func (s *System) acceptLoop(listener net.Listener) {
}
}
// lx/040: recreate a TCP forwarder listener that died out from under the
// stack. Returns the replacement listener after publishing it (listener field
// + atomic port) under listenAccess, or an error if the stack is closing or
// the bind failed.
func (s *System) healListener(dead net.Listener, isIPv6 bool, cause error) (net.Listener, error) {
port := &s.tcpPort
if isIPv6 {
port = &s.tcpPort6
}
oldPort := port.Load()
s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener (port ", oldPort, ") accept failed: ", cause, " — recreating listener")
_ = dead.Close() // release netpoll state; harmless if already closed
newListener, err := s.listenTCP(isIPv6)
if err != nil {
return nil, err
}
s.listenAccess.Lock()
defer s.listenAccess.Unlock()
if s.closing.Load() {
_ = newListener.Close()
return nil, net.ErrClosed
}
if isIPv6 {
s.tcpListener6 = newListener
} else {
s.tcpListener = newListener
}
newPort := uint32(M.SocksaddrFromNet(newListener.Addr()).Port)
port.Store(newPort)
recoveries := s.acceptRecoveries.Add(1)
s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener recreated (port ", oldPort, " → ", newPort, ", recoveries: ", recoveries, ")")
return newListener, nil
}
func ipVersionSuffix(isIPv6 bool) string {
if isIPv6 {
return "6"
}
return "4"
}
func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
@ -361,7 +480,7 @@ func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
if ipHdr.SourceAddr() == s.inet4Address &&
ipHdr.FragmentOffset() == 0 &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort {
header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort.Load()) {
return false
}
case header.ICMPv4ProtocolNumber:
@ -380,7 +499,7 @@ func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool {
}
if ipHdr.SourceAddr() == s.inet6Address &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 {
header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort6.Load()) {
return false
}
case header.ICMPv6ProtocolNumber:
@ -444,7 +563,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
if !destination.Addr().IsGlobalUnicast() {
return false, nil
} else if source.Addr() == s.inet4Address && source.Port() == s.tcpPort {
} else if source.Addr() == s.inet4Address && source.Port() == uint16(s.tcpPort.Load()) {
session := s.tcpNat.LookupBack(destination.Port())
if session == nil {
return false, E.New("ipv4: tcp: session not found: ", destination.Port())
@ -470,7 +589,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
}
rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload,
s.inet4NextAddress, natPort, true,
s.inet4Address, s.tcpPort, true)
s.inet4Address, uint16(s.tcpPort.Load()), true)
}
}
return true, nil
@ -481,7 +600,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
if !destination.Addr().IsGlobalUnicast() {
return false, nil
} else if source.Addr() == s.inet6Address && source.Port() == s.tcpPort6 {
} else if source.Addr() == s.inet6Address && source.Port() == uint16(s.tcpPort6.Load()) {
session := s.tcpNat.LookupBack(destination.Port())
if session == nil {
return false, E.New("ipv6: tcp: session not found: ", destination.Port())
@ -507,7 +626,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
}
rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload,
s.inet6NextAddress, natPort, true,
s.inet6Address, s.tcpPort6, true)
s.inet6Address, uint16(s.tcpPort6.Load()), true)
}
}
return true, nil
@ -682,20 +801,22 @@ type systemUDPPacketWriter4 struct {
txChecksumOffload bool
}
func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len())
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(w.header)
newPacket.Write(buffer.Bytes())
ipHdr := header.IPv4(newPacket.Bytes())
ipHdr.SetTotalLength(uint16(newPacket.Len()))
func (w *systemUDPPacketWriter4) FrontHeadroom() int {
return w.frontHeadroom + len(w.header)
}
func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
payloadLen := buffer.Len()
buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer)
copy(buffer.ExtendHeader(len(w.header)), w.header)
ipHdr := header.IPv4(buffer.Bytes())
ipHdr.SetTotalLength(uint16(buffer.Len()))
ipHdr.SetDestinationAddress(ipHdr.SourceAddress())
ipHdr.SetSourceAddr(destination.Addr)
udpHdr := header.UDP(ipHdr.Payload())
udpHdr.SetDestinationPort(udpHdr.SourcePort())
udpHdr.SetSourcePort(destination.Port)
udpHdr.SetLength(uint16(buffer.Len() + header.UDPMinimumSize))
udpHdr.SetLength(uint16(payloadLen + header.UDPMinimumSize))
if !w.txChecksumOffload {
udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum(
header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()),
@ -704,12 +825,61 @@ func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.S
udpHdr.SetChecksum(0)
}
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
return buffer
}
func (w *systemUDPPacketWriter4) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
buffer = w.preparePacket(buffer, destination)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} else {
newPacket.Advance(-w.frontHeadroom)
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv4Version)
}
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 {
buffer.Advance(-remainingHeadroom)
}
return buffer
}
func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
buffer = w.prepareWritePacket(buffer, destination)
defer buffer.Release()
return common.Error(w.tun.Write(buffer.Bytes()))
}
func (w *systemUDPPacketWriter4) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
switch w.tun.(type) {
case LinuxTUN, DarwinTUN:
return w, true
default:
return nil, false
}
}
func (w *systemUDPPacketWriter4) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error {
if len(buffers) == 0 || len(buffers) != len(destinations) {
buf.ReleaseMulti(buffers)
return os.ErrInvalid
}
defer func() {
buf.ReleaseMulti(buffers)
}()
switch tunInterface := w.tun.(type) {
case LinuxTUN:
packets := make([][]byte, len(buffers))
for index, buffer := range buffers {
buffer = w.preparePacket(buffer, destinations[index])
buffer.Advance(-w.frontHeadroom)
buffers[index] = buffer
packets[index] = buffer.Bytes()
}
return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom))
case DarwinTUN:
for index, buffer := range buffers {
buffers[index] = w.preparePacket(buffer, destinations[index])
}
return tunInterface.BatchWrite(buffers)
default:
return os.ErrInvalid
}
return common.Error(w.tun.Write(newPacket.Bytes()))
}
type systemUDPPacketWriter6 struct {
@ -720,14 +890,16 @@ type systemUDPPacketWriter6 struct {
txChecksumOffload bool
}
func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len())
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(w.header)
newPacket.Write(buffer.Bytes())
ipHdr := header.IPv6(newPacket.Bytes())
udpLen := uint16(header.UDPMinimumSize + buffer.Len())
func (w *systemUDPPacketWriter6) FrontHeadroom() int {
return w.frontHeadroom + len(w.header)
}
func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
payloadLen := buffer.Len()
buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer)
copy(buffer.ExtendHeader(len(w.header)), w.header)
ipHdr := header.IPv6(buffer.Bytes())
udpLen := uint16(header.UDPMinimumSize + payloadLen)
ipHdr.SetPayloadLength(udpLen)
ipHdr.SetDestinationAddress(ipHdr.SourceAddress())
ipHdr.SetSourceAddr(destination.Addr)
@ -742,12 +914,61 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
} else {
udpHdr.SetChecksum(0)
}
return buffer
}
func (w *systemUDPPacketWriter6) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
buffer = w.preparePacket(buffer, destination)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-w.frontHeadroom)
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv6Version)
}
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 {
buffer.Advance(-remainingHeadroom)
}
return buffer
}
func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
buffer = w.prepareWritePacket(buffer, destination)
defer buffer.Release()
return common.Error(w.tun.Write(buffer.Bytes()))
}
func (w *systemUDPPacketWriter6) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
switch w.tun.(type) {
case LinuxTUN, DarwinTUN:
return w, true
default:
return nil, false
}
}
func (w *systemUDPPacketWriter6) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error {
if len(buffers) == 0 || len(buffers) != len(destinations) {
buf.ReleaseMulti(buffers)
return os.ErrInvalid
}
defer func() {
buf.ReleaseMulti(buffers)
}()
switch tunInterface := w.tun.(type) {
case LinuxTUN:
packets := make([][]byte, len(buffers))
for index, buffer := range buffers {
buffer = w.preparePacket(buffer, destinations[index])
buffer.Advance(-w.frontHeadroom)
buffers[index] = buffer
packets[index] = buffer.Bytes()
}
return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom))
case DarwinTUN:
for index, buffer := range buffers {
buffers[index] = w.preparePacket(buffer, destinations[index])
}
return tunInterface.BatchWrite(buffers)
default:
return os.ErrInvalid
}
return common.Error(w.tun.Write(newPacket.Bytes()))
}
func newSystemWriteback(tunInterface Tun, frontHeadroom int) ForwardWriteback {

View file

@ -86,6 +86,15 @@ func (n *TCPNat) checkTimeout() {
n.addrAccess.Unlock()
}
func (n *TCPNat) Purge() {
n.addrAccess.Lock()
n.portAccess.Lock()
clear(n.addrMap)
clear(n.portMap)
n.portAccess.Unlock()
n.addrAccess.Unlock()
}
func (n *TCPNat) LookupBack(port uint16) *TCPSession {
n.portAccess.RLock()
session := n.portMap[port]

View file

@ -5,7 +5,6 @@ import (
"syscall"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common"
)
func PacketIPVersion(packet []byte) int {
@ -14,7 +13,7 @@ func PacketIPVersion(packet []byte) int {
func PacketFillHeader(packet []byte, ipVersion int) {
if PacketOffset > 0 {
common.ClearArray(packet[:3])
clear(packet[:3])
switch ipVersion {
case header.IPv4Version:
packet[3] = syscall.AF_INET

View file

@ -0,0 +1,117 @@
package tun
// lx/040 (SPECS/TASKS/040-SINGTUN_ACCEPTLOOP_SELFHEAL): acceptLoop self-heal.
//
// Red/green против апстрима 2d9b8aed5fe2: там acceptLoop(listener) при любой
// ошибке Accept молча выходит навсегда — восстановления нет, порт не меняется,
// новый connect вечно бьётся в мёртвый сокет. Для red-прогона на чистом
// апстрим-чекауте достаточно адаптировать хелперы ниже (currentTCPPort →
// s.tcpPort, spawnAcceptLoop → go s.acceptLoop(ln)): тест упадёт по таймауту
// ожидания восстановления.
import (
"context"
"fmt"
"net"
"net/netip"
"testing"
"time"
"github.com/sagernet/sing/common/logger"
)
func newSelfHealTestSystem(t *testing.T) *System {
t.Helper()
s := &System{
ctx: context.Background(),
logger: logger.NOP(),
inet4Address: netip.MustParseAddr("127.0.0.1"),
udpTimeout: time.Minute,
}
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
ln, err := s.listenTCP(false)
if err != nil {
t.Fatalf("listenTCP: %v", err)
}
s.tcpListener = ln
s.tcpPort.Store(uint32(ln.Addr().(*net.TCPAddr).Port))
spawnAcceptLoop(s, ln)
return s
}
func currentTCPPort(s *System) uint32 {
return s.tcpPort.Load()
}
func spawnAcceptLoop(s *System, ln net.Listener) {
go s.acceptLoop(ln, false)
}
func dialForwarder(t *testing.T, port uint32) error {
t.Helper()
conn, err := net.DialTimeout("tcp4", fmt.Sprintf("127.0.0.1:%d", port), time.Second)
if err == nil {
_ = conn.Close()
}
return err
}
// Убийство listener'а мимо System.Close (эмуляция чужого close по
// переиспользованному fd-номеру) должно приводить к пересозданию listener'а
// и продолжению приёма TCP, а не к вечной смерти петли.
func TestSystemAcceptLoopSelfHeal(t *testing.T) {
s := newSelfHealTestSystem(t)
oldPort := currentTCPPort(s)
if err := dialForwarder(t, oldPort); err != nil {
t.Fatalf("healthy listener refused connect: %v", err)
}
// Убить listener из-под стека: closing НЕ выставлен.
_ = s.tcpListener.Close()
deadline := time.Now().Add(5 * time.Second)
healed := false
for time.Now().Before(deadline) {
if s.acceptRecoveries.Load() > 0 {
healed = true
break
}
time.Sleep(10 * time.Millisecond)
}
if !healed {
t.Fatalf("acceptLoop did not recover within 5s (upstream behavior: silent permanent death)")
}
newPort := currentTCPPort(s)
if newPort == oldPort {
t.Fatalf("recovered port equals dead port %d — relisten did not publish a new port", oldPort)
}
if err := dialForwarder(t, newPort); err != nil {
t.Fatalf("connect to recreated listener (port %d) failed: %v", newPort, err)
}
if got := s.acceptRecoveries.Load(); got != 1 {
t.Fatalf("acceptRecoveries = %d, want 1", got)
}
s.closing.Store(true)
_ = s.tcpListener.Close()
}
// Штатное закрытие (closing выставлен, как это делает System.Close) обязано
// оставаться тихим: без пересозданий и без роста счётчика.
func TestSystemAcceptLoopQuietOnClose(t *testing.T) {
s := newSelfHealTestSystem(t)
oldPort := currentTCPPort(s)
s.closing.Store(true)
_ = s.tcpListener.Close()
time.Sleep(300 * time.Millisecond)
if got := s.acceptRecoveries.Load(); got != 0 {
t.Fatalf("deliberate close triggered %d recoveries, want 0", got)
}
if port := currentTCPPort(s); port != oldPort {
t.Fatalf("deliberate close changed port %d → %d", oldPort, port)
}
}

3
tun.go
View file

@ -14,12 +14,14 @@ import (
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/ranges"
)
type Handler interface {
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort, firstPacket []byte) FlowVerdict
NewDNSPacket(payload []byte, source M.Socksaddr, destination M.Socksaddr, writer N.PacketWriter)
N.TCPConnectionHandlerEx
N.UDPConnectionHandlerEx
}
@ -66,6 +68,7 @@ const (
type Options struct {
Name string
NetNs string
Inet4Address []netip.Prefix
Inet6Address []netip.Prefix
MTU uint32

View file

@ -51,8 +51,8 @@ type NativeTun struct {
}
func New(options Options) (Tun, error) {
var nativeTun *NativeTun
if options.FileDescriptor == 0 {
return execInNetworkNamespace(options.NetNs, func() (Tun, error) {
tunFd, err := open(options.Name, options.GSO)
if err != nil {
return nil, E.Cause(err, "open tun")
@ -61,7 +61,7 @@ func New(options Options) (Tun, error) {
if err != nil {
return nil, E.Errors(err, unix.Close(tunFd))
}
nativeTun = &NativeTun{
nativeTun := &NativeTun{
tunFd: tunFd,
tunFile: os.NewFile(uintptr(tunFd), "tun"),
options: options,
@ -70,8 +70,10 @@ func New(options Options) (Tun, error) {
if err != nil {
return nil, E.Errors(err, unix.Close(tunFd))
}
} else {
nativeTun = &NativeTun{
return nativeTun, nil
})
}
nativeTun := &NativeTun{
tunFd: options.FileDescriptor,
tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"),
options: options,
@ -84,7 +86,6 @@ func New(options Options) (Tun, error) {
}
}
}
}
return nativeTun, nil
}
@ -290,10 +291,10 @@ func (t *NativeTun) Name() (string, error) {
func (t *NativeTun) Start() error {
if t.options.FileDescriptor == 0 {
if !t.options.EXP_ExternalConfiguration {
if !t.options.EXP_ExternalConfiguration && t.options.NetNs == "" {
t.options.InterfaceMonitor.RegisterMyInterface(t.options.Name)
}
err := t.start()
err := runInNetworkNamespace(t.options.NetNs, t.start)
if err != nil {
return err
}
@ -354,7 +355,7 @@ func (t *NativeTun) start() error {
return E.Cause(err, "set rules")
}
if t.options.DNSMode != DNSModeDisabled {
if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" {
err = t.setSearchDomainForSystemdResolved()
if err != nil {
return E.Cause(err, "set search domain")
@ -374,11 +375,13 @@ func (t *NativeTun) Close() error {
if t.options.EXP_ExternalConfiguration {
return common.Close(common.PtrOrNil(t.tunFile))
}
if t.options.DNSMode != DNSModeDisabled {
if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" {
t.unsetSearchDomainForSystemdResolved()
}
return E.Errors(runInNetworkNamespace(t.options.NetNs, func() error {
t.unsetAddresses()
return E.Errors(t.unsetRoute(), t.unsetRules(), common.Close(common.PtrOrNil(t.tunFile)))
return E.Errors(t.unsetRoute(), t.unsetRules())
}), common.Close(common.PtrOrNil(t.tunFile)))
}
func (t *NativeTun) Read(p []byte) (n int, err error) {
@ -625,6 +628,7 @@ func (t *NativeTun) UpdateRouteOptions(tunOptions Options) error {
t.options = tunOptions
return nil
}
return runInNetworkNamespace(t.options.NetNs, func() error {
tunLink, err := netlink.LinkByName(t.options.Name)
if err != nil {
return E.Cause(err, "find tun interface")
@ -635,6 +639,7 @@ func (t *NativeTun) UpdateRouteOptions(tunOptions Options) error {
}
t.options = tunOptions
return t.setRoute(tunLink)
})
}
func (t *NativeTun) routes(tunLink netlink.Link) ([]netlink.Route, error) {

287
udp_egress.go Normal file
View file

@ -0,0 +1,287 @@
package tun
import (
"context"
"net"
"net/netip"
"runtime"
"slices"
"sync"
"sync/atomic"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
"github.com/sagernet/sing/common/x/list"
)
const udpEgressBufferSize = 65535
type UDPEgressPoolOptions struct {
Logger logger.Logger
Network string
Control control.Func
InterfaceFinder control.InterfaceFinder
InterfaceMonitor DefaultInterfaceMonitor
ExcludeInterface string
IsExempt func() bool
}
type UDPEgressPool struct {
logger logger.Logger
network string
control control.Func
interfaceFinder control.InterfaceFinder
interfaceMonitor DefaultInterfaceMonitor
excludeInterface string
isExempt func() bool
access sync.Mutex
port uint16
anchorInterfaceIndex int
receiveDone chan struct{}
members map[udpEgressSpec]*udpEgressMember
state atomic.Pointer[[]*udpEgressMember]
packetChan chan udpEgressPacket
memberReaders sync.WaitGroup
finderElement *list.Element[control.InterfaceUpdateCallback]
}
type udpEgressSpec struct {
interfaceIndex int
interfaceName string
prefix netip.Prefix
}
type udpEgressMember struct {
prefix netip.Prefix
conn *net.UDPConn
}
type udpEgressPacket struct {
buffer *buf.Buffer
source netip.AddrPort
}
func NewUDPEgressPool(options UDPEgressPoolOptions) *UDPEgressPool {
return &UDPEgressPool{
logger: options.Logger,
network: options.Network,
control: options.Control,
interfaceFinder: options.InterfaceFinder,
interfaceMonitor: options.InterfaceMonitor,
excludeInterface: options.ExcludeInterface,
isExempt: options.IsExempt,
anchorInterfaceIndex: -1,
members: make(map[udpEgressSpec]*udpEgressMember),
packetChan: make(chan udpEgressPacket, 128),
}
}
func (p *UDPEgressPool) Close() {
p.SetEgressPort(0)
p.access.Lock()
defer p.access.Unlock()
if p.finderElement != nil {
p.interfaceFinder.UnregisterInterfaceUpdateCallback(p.finderElement)
p.finderElement = nil
}
}
func (p *UDPEgressPool) SetEgressPort(port uint16) bool {
p.access.Lock()
defer p.access.Unlock()
if p.port == port {
return p.state.Load() != nil
}
if p.receiveDone != nil {
close(p.receiveDone)
p.receiveDone = nil
}
p.port = 0
p.state.Store(nil)
for spec, member := range p.members {
delete(p.members, spec)
member.conn.Close()
}
p.memberReaders.Wait()
for {
select {
case packet := <-p.packetChan:
packet.buffer.Release()
default:
goto drained
}
}
drained:
p.anchorInterfaceIndex = -1
if port == 0 {
return false
}
p.port = port
defaultInterface := p.interfaceMonitor.DefaultInterface()
if defaultInterface != nil {
p.anchorInterfaceIndex = defaultInterface.Index
}
p.receiveDone = make(chan struct{})
if p.finderElement == nil {
p.finderElement = p.interfaceFinder.RegisterInterfaceUpdateCallback(func(interfaces []control.Interface) {
p.access.Lock()
defer p.access.Unlock()
p.rebuildLocked()
})
}
p.rebuildLocked()
return p.state.Load() != nil
}
func (p *UDPEgressPool) LookupEgress(destination netip.AddrPort) *net.UDPConn {
members := p.state.Load()
if members == nil {
return nil
}
address := destination.Addr().Unmap()
for _, member := range *members {
if member.prefix.Contains(address) {
return member.conn
}
}
return nil
}
func (p *UDPEgressPool) ReceiveEgress(buffer []byte) (int, netip.AddrPort, error) {
p.access.Lock()
receiveDone := p.receiveDone
p.access.Unlock()
if receiveDone == nil {
return 0, netip.AddrPort{}, net.ErrClosed
}
select {
case <-receiveDone:
return 0, netip.AddrPort{}, net.ErrClosed
default:
}
select {
case packet := <-p.packetChan:
copied := copy(buffer, packet.buffer.Bytes())
packet.buffer.Release()
return copied, packet.source, nil
case <-receiveDone:
return 0, netip.AddrPort{}, net.ErrClosed
}
}
func (p *UDPEgressPool) rebuildLocked() {
if p.port == 0 {
return
}
specs := make(map[udpEgressSpec]struct{})
if !p.isExempt() {
for _, networkInterface := range p.interfaceFinder.Interfaces() {
if networkInterface.Flags&net.FlagUp == 0 ||
networkInterface.Flags&net.FlagLoopback != 0 ||
networkInterface.Flags&net.FlagPointToPoint != 0 ||
networkInterface.Flags&net.FlagBroadcast == 0 ||
networkInterface.Index == p.anchorInterfaceIndex ||
networkInterface.Name == p.excludeInterface {
continue
}
for _, prefix := range networkInterface.Addresses {
if !prefix.Addr().IsGlobalUnicast() {
continue
}
if p.network == "udp4" && !prefix.Addr().Is4() {
continue
}
if p.network == "udp6" && prefix.Addr().Is4() {
continue
}
specs[udpEgressSpec{
interfaceIndex: networkInterface.Index,
interfaceName: networkInterface.Name,
prefix: prefix,
}] = struct{}{}
}
}
}
for spec, member := range p.members {
_, loaded := specs[spec]
if loaded {
continue
}
delete(p.members, spec)
member.conn.Close()
}
for spec := range specs {
_, loaded := p.members[spec]
if loaded {
continue
}
memberConn, err := p.listenMember(spec)
if err != nil {
p.logger.Warn(E.Cause(err, "listen egress member on ", spec.interfaceName, " (", spec.prefix.Addr(), ")"))
continue
}
member := &udpEgressMember{
prefix: spec.prefix.Masked(),
conn: memberConn,
}
p.members[spec] = member
p.memberReaders.Add(1)
go p.readMember(member, p.receiveDone)
}
members := make([]*udpEgressMember, 0, len(p.members))
for _, member := range p.members {
members = append(members, member)
}
slices.SortFunc(members, func(firstMember, secondMember *udpEgressMember) int {
return secondMember.prefix.Bits() - firstMember.prefix.Bits()
})
if len(members) == 0 {
p.state.Store(nil)
} else {
p.state.Store(&members)
}
}
func (p *UDPEgressPool) listenMember(spec udpEgressSpec) (*net.UDPConn, error) {
var listenConfig net.ListenConfig
if runtime.GOOS == "darwin" || runtime.GOOS == "ios" {
listenConfig.Control = control.ReuseAddrOnly()
}
listenConfig.Control = control.Append(listenConfig.Control, control.DisableUDPNetReset())
listenConfig.Control = control.Append(listenConfig.Control, control.BindToInterface(p.interfaceFinder, spec.interfaceName, spec.interfaceIndex))
listenConfig.Control = control.Append(listenConfig.Control, p.control)
var network string
if spec.prefix.Addr().Is4() {
network = "udp4"
} else {
network = "udp6"
}
packetConn, err := listenConfig.ListenPacket(context.Background(), network, netip.AddrPortFrom(spec.prefix.Addr(), p.port).String())
if err != nil {
return nil, err
}
return packetConn.(*net.UDPConn), nil
}
func (p *UDPEgressPool) readMember(member *udpEgressMember, doneChan <-chan struct{}) {
defer p.memberReaders.Done()
for {
buffer := buf.NewSize(udpEgressBufferSize)
dataLength, source, err := member.conn.ReadFromUDPAddrPort(buffer.FreeBytes())
if err != nil {
buffer.Release()
return
}
buffer.Extend(dataLength)
select {
case p.packetChan <- udpEgressPacket{buffer: buffer, source: source}:
case <-doneChan:
buffer.Release()
return
default:
buffer.Release()
}
}
}

124
udp_egress_conn.go Normal file
View file

@ -0,0 +1,124 @@
package tun
import (
"net"
"net/netip"
"sync"
"time"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
)
type UDPEgressConn struct {
anchor *net.UDPConn
pool *UDPEgressPool
packetChan chan udpEgressConnPacket
doneChan chan struct{}
closeOnce sync.Once
readWait sync.WaitGroup
}
type udpEgressConnPacket struct {
buffer *buf.Buffer
source netip.AddrPort
err error
}
func NewUDPEgressConn(anchor *net.UDPConn, pool *UDPEgressPool) *UDPEgressConn {
conn := &UDPEgressConn{
anchor: anchor,
pool: pool,
packetChan: make(chan udpEgressConnPacket, 64),
doneChan: make(chan struct{}),
}
conn.readWait.Add(2)
go conn.read(anchor.ReadFromUDPAddrPort)
go conn.read(pool.ReceiveEgress)
return conn
}
func (c *UDPEgressConn) read(readPacket func([]byte) (int, netip.AddrPort, error)) {
defer c.readWait.Done()
for {
buffer := buf.NewSize(udpEgressBufferSize)
dataLength, source, err := readPacket(buffer.FreeBytes())
if err != nil {
buffer.Release()
if E.IsClosed(err) {
return
}
select {
case c.packetChan <- udpEgressConnPacket{err: err}:
case <-c.doneChan:
return
}
continue
}
buffer.Extend(dataLength)
select {
case c.packetChan <- udpEgressConnPacket{buffer: buffer, source: source}:
case <-c.doneChan:
buffer.Release()
return
}
}
}
func (c *UDPEgressConn) ReadFromUDPAddrPort(buffer []byte) (int, netip.AddrPort, error) {
select {
case packet := <-c.packetChan:
if packet.err != nil {
return 0, netip.AddrPort{}, packet.err
}
copied := copy(buffer, packet.buffer.Bytes())
packet.buffer.Release()
return copied, packet.source, nil
case <-c.doneChan:
return 0, netip.AddrPort{}, net.ErrClosed
}
}
func (c *UDPEgressConn) WriteToUDPAddrPort(buffer []byte, destination netip.AddrPort) (int, error) {
memberConn := c.pool.LookupEgress(destination)
if memberConn != nil {
return memberConn.WriteToUDPAddrPort(buffer, destination)
}
return c.anchor.WriteToUDPAddrPort(buffer, destination)
}
func (c *UDPEgressConn) LocalAddr() net.Addr {
return c.anchor.LocalAddr()
}
func (c *UDPEgressConn) SetDeadline(t time.Time) error {
return c.anchor.SetDeadline(t)
}
func (c *UDPEgressConn) SetReadDeadline(t time.Time) error {
return c.anchor.SetReadDeadline(t)
}
func (c *UDPEgressConn) SetWriteDeadline(t time.Time) error {
return c.anchor.SetWriteDeadline(t)
}
func (c *UDPEgressConn) Close() error {
c.closeOnce.Do(func() {
close(c.doneChan)
c.anchor.Close()
c.pool.Close()
c.readWait.Wait()
for {
select {
case packet := <-c.packetChan:
if packet.buffer != nil {
packet.buffer.Release()
}
default:
return
}
}
})
return nil
}

789
udp_nat.go Normal file
View file

@ -0,0 +1,789 @@
package tun
import (
"context"
"io"
"net"
"net/netip"
"os"
"runtime"
"slices"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/canceler"
"github.com/sagernet/sing/common/control"
"github.com/sagernet/sing/common/memory"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/pipe"
"github.com/sagernet/sing/common/x/list"
"github.com/sagernet/sing/contrab/freelru"
"github.com/sagernet/sing/contrab/maphash"
)
type NATMapping uint8
const (
NATMappingEndpointIndependent NATMapping = iota
NATMappingAddressDependent
NATMappingAddressAndPortDependent
)
type NATFiltering uint8
const (
NATFilteringEndpointIndependent NATFiltering = iota
NATFilteringAddressDependent
NATFilteringAddressAndPortDependent
)
type UDPNatPrepareFunc func(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc)
type UDPNatOptions struct {
Handler N.UDPConnectionHandlerEx
Prepare UDPNatPrepareFunc
Timeout time.Duration
Shared bool
Mapping NATMapping
Filtering NATFiltering
MaxSize uint32
InterfaceFinder control.InterfaceFinder
ExcludeInterface []string
}
type udpNatSessionKey struct {
sourceAddr netip.Addr
destinationAddr netip.Addr
sourcePort uint16
destinationPort uint16
interfaceIndex uint32
}
type udpNatFilterKey struct {
sessionID uint64
peer netip.AddrPort
}
type udpNatEgressEntry struct {
prefix netip.Prefix
interfaceIndex uint32
}
const udpNatEgressLinearThreshold = 8
type udpNatEgressBuckets struct {
inet4 [256][]udpNatEgressEntry
inet6 [256][]udpNatEgressEntry
}
type udpNatEgressTable struct {
entries []udpNatEgressEntry
buckets *udpNatEgressBuckets
}
type UDPNat struct {
handler N.UDPConnectionHandlerEx
prepare UDPNatPrepareFunc
timeout time.Duration
mapping NATMapping
filtering NATFiltering
cache *freelru.Cache[udpNatSessionKey, *udpNatConn]
filterCache *freelru.Cache[udpNatFilterKey, *udpNatConn]
nextFilterSessionID atomic.Uint64
interfaceFinder control.InterfaceFinder
excludeInterface []string
interfaceElement *list.Element[control.InterfaceUpdateCallback]
egress atomic.Pointer[udpNatEgressTable]
classAccess sync.Mutex
classConns map[uint32]map[*udpNatConn]struct{}
cleanup *udpNatCleanupQueue
state atomic.Uint32
lifecycleAccess sync.Mutex
closeOnce sync.Once
cleanupDone chan struct{}
cleanupWait sync.WaitGroup
}
func NewUDPNat(options UDPNatOptions) *UDPNat {
if options.Timeout == 0 {
panic("invalid timeout")
}
maxSize := options.MaxSize
if maxSize == 0 {
if runtime.GOOS == "ios" {
maxSize = 4096
} else if totalMemory := memory.Total(); totalMemory == 0 {
maxSize = 16384
} else {
maxSize = uint32(min(max(totalMemory/16384, 4096), 16384))
}
}
hasher := maphash.NewHasher[udpNatSessionKey]()
cache := common.Must1(freelru.New[udpNatSessionKey, *udpNatConn](maxSize, hasher.Hash32, options.Shared))
var filterCache *freelru.Cache[udpNatFilterKey, *udpNatConn]
if NATMapping(options.Filtering) > options.Mapping {
filterHasher := maphash.NewHasher[udpNatFilterKey]()
filterCache = common.Must1(freelru.New[udpNatFilterKey, *udpNatConn](maxSize, filterHasher.Hash32, options.Shared))
}
service := &UDPNat{
handler: options.Handler,
prepare: options.Prepare,
timeout: options.Timeout,
mapping: options.Mapping,
filtering: options.Filtering,
cache: cache,
filterCache: filterCache,
interfaceFinder: options.InterfaceFinder,
excludeInterface: options.ExcludeInterface,
classConns: make(map[uint32]map[*udpNatConn]struct{}),
cleanupDone: make(chan struct{}),
}
service.cleanup = newUDPNatCleanupQueue(service)
cache.SetLifetime(options.Timeout)
cache.SetHealthCheck(func(_ udpNatSessionKey, conn *udpNatConn) bool {
select {
case <-conn.doneChan:
return false
default:
return true
}
})
cache.SetOnEvict(func(_ udpNatSessionKey, conn *udpNatConn) {
conn.closeFromCache()
})
if filterCache != nil {
filterCache.SetOnEvict(func(key udpNatFilterKey, conn *udpNatConn) {
conn.removeFilterPeer(key.peer)
})
}
return service
}
func (s *UDPNat) Close() error {
s.closeOnce.Do(func() {
s.lifecycleAccess.Lock()
previousState := s.state.Swap(udpNatStateClosed)
if previousState == udpNatStateStarted {
close(s.cleanupDone)
}
s.lifecycleAccess.Unlock()
if previousState == udpNatStateStarted {
s.cleanupWait.Wait()
}
if s.interfaceElement != nil {
s.interfaceFinder.UnregisterInterfaceUpdateCallback(s.interfaceElement)
s.interfaceElement = nil
}
s.cache.Purge()
if s.filterCache != nil {
s.filterCache.Purge()
}
s.cleanup.clear()
})
return nil
}
func (s *UDPNat) reloadInterfaces() {
s.updateInterfaces(s.interfaceFinder.Interfaces())
}
func (s *UDPNat) updateInterfaces(interfaces []control.Interface) {
var entries []udpNatEgressEntry
for _, networkInterface := range interfaces {
if networkInterface.Flags&net.FlagUp == 0 ||
networkInterface.Flags&net.FlagLoopback != 0 ||
networkInterface.Flags&net.FlagPointToPoint != 0 ||
networkInterface.Flags&net.FlagBroadcast == 0 {
continue
}
if slices.Contains(s.excludeInterface, networkInterface.Name) {
continue
}
for _, prefix := range networkInterface.Addresses {
if !prefix.Addr().IsGlobalUnicast() {
continue
}
entries = append(entries, udpNatEgressEntry{prefix.Masked(), uint32(networkInterface.Index)})
}
}
s.egress.Store(newUDPNatEgressTable(entries))
var closeConns []*udpNatConn
s.classAccess.Lock()
for interfaceIndex, conns := range s.classConns {
if !slices.ContainsFunc(entries, func(entry udpNatEgressEntry) bool {
return entry.interfaceIndex == interfaceIndex
}) {
for conn := range conns {
closeConns = append(closeConns, conn)
}
delete(s.classConns, interfaceIndex)
}
}
s.classAccess.Unlock()
for _, conn := range closeConns {
conn.Close()
}
}
func (s *UDPNat) classify(destination M.Socksaddr) uint32 {
table := s.egress.Load()
if table == nil || !destination.IsIP() {
return 0
}
return table.lookup(destination.Addr.Unmap())
}
func newUDPNatEgressTable(entries []udpNatEgressEntry) *udpNatEgressTable {
entries = slices.Clone(entries)
slices.SortStableFunc(entries, func(a, b udpNatEgressEntry) int {
return b.prefix.Bits() - a.prefix.Bits()
})
table := &udpNatEgressTable{entries: entries}
if len(entries) <= udpNatEgressLinearThreshold {
return table
}
buckets := new(udpNatEgressBuckets)
for _, entry := range entries {
address := entry.prefix.Addr().Unmap()
bits := entry.prefix.Bits()
var target *[256][]udpNatEgressEntry
var firstByte byte
if address.Is4() {
target = &buckets.inet4
firstByte = address.As4()[0]
} else {
target = &buckets.inet6
firstByte = address.As16()[0]
}
if bits >= 8 {
target[firstByte] = append(target[firstByte], entry)
continue
}
var mask byte
if bits > 0 {
mask = ^byte(0) << (8 - bits)
}
firstByte &= mask
for index := 0; index < 1<<(8-bits); index++ {
bucketIndex := firstByte + byte(index)
target[bucketIndex] = append(target[bucketIndex], entry)
}
}
table.buckets = buckets
return table
}
func (t *udpNatEgressTable) lookup(address netip.Addr) uint32 {
entries := t.entries
if t.buckets != nil {
if address.Is4() {
entries = t.buckets.inet4[address.As4()[0]]
} else {
entries = t.buckets.inet6[address.As16()[0]]
}
}
for _, entry := range entries {
if entry.prefix.Contains(address) {
return entry.interfaceIndex
}
}
return 0
}
func (s *UDPNat) registerClass(conn *udpNatConn) {
s.classAccess.Lock()
conns := s.classConns[conn.interfaceIndex]
if conns == nil {
conns = make(map[*udpNatConn]struct{})
s.classConns[conn.interfaceIndex] = conns
}
conns[conn] = struct{}{}
s.classAccess.Unlock()
}
func (s *UDPNat) unregisterClass(conn *udpNatConn) {
s.classAccess.Lock()
conns := s.classConns[conn.interfaceIndex]
if conns != nil {
delete(conns, conn)
if len(conns) == 0 {
delete(s.classConns, conn.interfaceIndex)
}
}
s.classAccess.Unlock()
}
func (s *UDPNat) NewPacket(bufferSlices [][]byte, source M.Socksaddr, destination M.Socksaddr, userData any) {
conn, ok := s.getOrCreateConn(source, destination, userData)
if !ok {
return
}
readWaitOptions := conn.loadReadWaitOptions()
var dataLen int
for _, bufferSlice := range bufferSlices {
dataLen += len(bufferSlice)
}
buffer := readWaitOptions.NewBufferSize(dataLen)
for _, bufferSlice := range bufferSlices {
buffer.Write(bufferSlice)
}
readWaitOptions.PostReturn(buffer)
conn.enqueue(buffer, destination)
}
func (s *UDPNat) getOrCreateConn(source M.Socksaddr, destination M.Socksaddr, userData any) (*udpNatConn, bool) {
if s.state.Load() != udpNatStateStarted {
return nil, false
}
key := udpNatSessionKey{
sourceAddr: source.Addr.Unmap(),
sourcePort: source.Port,
}
switch s.mapping {
case NATMappingEndpointIndependent:
key.interfaceIndex = s.classify(destination)
case NATMappingAddressDependent:
key.destinationAddr = destination.Addr.Unmap()
case NATMappingAddressAndPortDependent:
key.destinationAddr = destination.Addr.Unmap()
key.destinationPort = destination.Port
}
var (
newContext context.Context
newOnClose N.CloseHandlerFunc
)
conn, loaded, ok := s.cache.GetAndRefreshOrAdd(key, func() (*udpNatConn, bool) {
ok, ctx, writer, onClose := s.prepare(source, destination, userData)
if !ok {
return nil, false
}
newConn := &udpNatConn{
service: s,
key: key,
writer: writer,
localAddr: source,
packetChan: make(chan *N.PacketBuffer, 64),
doneChan: make(chan struct{}),
readDeadline: pipe.MakeDeadline(),
}
newConn.cleanupEntry = &udpNatCleanupEntry{
conn: newConn,
index: -1,
}
if s.filtering != NATFilteringEndpointIndependent {
if destination.IsIP() {
newConn.filterPeer = s.filterPeer(destination)
newConn.filterPeerValid = true
}
if s.filterCache != nil {
filterSessionID := s.nextFilterSessionID.Add(1)
if filterSessionID == 0 {
filterSessionID = s.nextFilterSessionID.Add(1)
}
newConn.filterSessionID = filterSessionID
}
}
interfaceIndex := key.interfaceIndex
if s.mapping != NATMappingEndpointIndependent {
interfaceIndex = s.classify(destination)
}
if interfaceIndex != 0 {
newConn.interfaceIndex = interfaceIndex
s.registerClass(newConn)
}
newContext = ctx
newOnClose = onClose
return newConn, true
})
if !ok {
return nil, false
}
if s.state.Load() != udpNatStateStarted {
conn.Close()
s.cache.Peek(key)
return nil, false
}
if !loaded {
s.cleanup.addOrUpdate(conn.cleanupEntry, time.Now().Add(s.timeout))
if conn.isClosed() {
return nil, false
}
go s.handler.NewPacketConnectionEx(newContext, conn, source, destination, newOnClose)
}
conn.addFilterPeer(destination)
return conn, true
}
func (c *udpNatConn) enqueue(buffer *buf.Buffer, destination M.Socksaddr) {
c.packetAccess.RLock()
select {
case <-c.doneChan:
buffer.Release()
c.packetAccess.RUnlock()
return
default:
}
packet := N.NewPacketBuffer()
*packet = N.PacketBuffer{
Buffer: buffer,
Destination: destination,
}
select {
case c.packetChan <- packet:
default:
packet.Buffer.Release()
N.PutPacketBuffer(packet)
}
c.packetAccess.RUnlock()
}
func (s *UDPNat) NewPacketBatch(buffers []*buf.Buffer, sources []M.Socksaddr, destination M.Socksaddr, userData any) {
if len(buffers) != len(sources) {
buf.ReleaseMulti(buffers)
return
}
for index, buffer := range buffers {
conn, ok := s.getOrCreateConn(sources[index], destination, userData)
if !ok {
buffer.Release()
continue
}
readWaitOptions := conn.loadReadWaitOptions()
conn.enqueue(readWaitOptions.Copy(buffer), destination)
}
}
func (s *UDPNat) filterPeer(destination M.Socksaddr) netip.AddrPort {
if s.filtering == NATFilteringAddressDependent {
return netip.AddrPortFrom(destination.Addr.Unmap(), 0)
}
return netip.AddrPortFrom(destination.Addr.Unmap(), destination.Port)
}
func (s *UDPNat) Purge() {
if s.filterCache != nil {
s.filterCache.Purge()
}
s.cache.Purge()
}
func (s *UDPNat) PurgeExpired() {
s.cache.PurgeExpired()
}
var (
_ N.PacketConn = (*udpNatConn)(nil)
_ canceler.PacketConn = (*udpNatConn)(nil)
_ N.PacketBatchReadWaitCreator = (*udpNatConn)(nil)
_ N.PacketBatchWriteCreator = (*udpNatConn)(nil)
)
type udpNatConn struct {
service *UDPNat
key udpNatSessionKey
interfaceIndex uint32
writer N.PacketWriter
localAddr M.Socksaddr
packetChan chan *N.PacketBuffer
packetAccess sync.RWMutex
closeOnce sync.Once
doneChan chan struct{}
readDeadline pipe.Deadline
readWaitOptions atomic.Pointer[N.ReadWaitOptions]
readBatch *udpNatReadBatch
cleanupEntry *udpNatCleanupEntry
filterSessionID uint64
filterPeer netip.AddrPort
filterPeerValid bool
filterAccess sync.Mutex
filterPeers map[netip.AddrPort]struct{}
}
type udpNatReadBatch struct {
buffers []*buf.Buffer
destinations []M.Socksaddr
}
func (c *udpNatConn) loadReadWaitOptions() N.ReadWaitOptions {
options := c.readWaitOptions.Load()
if options == nil {
return N.ReadWaitOptions{}
}
return *options
}
func (c *udpNatConn) addFilterPeer(destination M.Socksaddr) {
if c.filterSessionID == 0 || !destination.IsIP() {
return
}
key := udpNatFilterKey{
sessionID: c.filterSessionID,
peer: c.service.filterPeer(destination),
}
if c.filterPeerValid && c.filterPeer == key.peer {
return
}
if c.isClosed() || c.service.state.Load() != udpNatStateStarted {
return
}
c.service.filterCache.Add(key, c)
c.filterAccess.Lock()
if c.isClosed() || c.service.state.Load() != udpNatStateStarted {
c.filterAccess.Unlock()
c.service.filterCache.Remove(key)
return
}
if c.filterPeers == nil {
c.filterPeers = make(map[netip.AddrPort]struct{})
}
c.filterPeers[key.peer] = struct{}{}
c.filterAccess.Unlock()
filterConn, loaded := c.service.filterCache.Peek(key)
if !loaded || filterConn != c {
c.removeFilterPeer(key.peer)
return
}
if c.isClosed() || c.service.state.Load() != udpNatStateStarted {
c.service.filterCache.Remove(key)
}
}
func (c *udpNatConn) removeFilterPeer(peer netip.AddrPort) {
c.filterAccess.Lock()
delete(c.filterPeers, peer)
c.filterAccess.Unlock()
}
func (c *udpNatConn) clearFilterPeers() {
if c.filterSessionID == 0 {
return
}
c.filterAccess.Lock()
filterPeers := c.filterPeers
c.filterPeers = nil
c.filterAccess.Unlock()
for peer := range filterPeers {
c.service.filterCache.Remove(udpNatFilterKey{
sessionID: c.filterSessionID,
peer: peer,
})
}
}
func (c *udpNatConn) allowPeer(destination M.Socksaddr) bool {
if c.service.filtering == NATFilteringEndpointIndependent || !destination.IsIP() {
return true
}
peer := c.service.filterPeer(destination)
if c.filterPeerValid && c.filterPeer == peer {
return true
}
if c.filterSessionID == 0 {
return false
}
filterConn, loaded := c.service.filterCache.Get(udpNatFilterKey{
sessionID: c.filterSessionID,
peer: peer,
})
return loaded && filterConn == c
}
func (c *udpNatConn) ReadPacket(buffer *buf.Buffer) (addr M.Socksaddr, err error) {
select {
case p := <-c.packetChan:
_, err = buffer.ReadOnceFrom(p.Buffer)
destination := p.Destination
p.Buffer.Release()
N.PutPacketBuffer(p)
return destination, err
case <-c.doneChan:
return M.Socksaddr{}, io.ErrClosedPipe
case <-c.readDeadline.Wait():
return M.Socksaddr{}, os.ErrDeadlineExceeded
}
}
func (c *udpNatConn) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
if !c.allowPeer(destination) {
buffer.Release()
return nil
}
return c.writer.WritePacket(buffer, destination)
}
func (c *udpNatConn) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
if c.service.filtering != NATFilteringEndpointIndependent {
return nil, false
}
if creator, isCreator := c.writer.(N.PacketBatchWriteCreator); isCreator {
return creator.CreatePacketBatchWriter()
}
if writer, isWriter := c.writer.(N.PacketBatchWriter); isWriter {
return writer, true
}
return nil, false
}
func (c *udpNatConn) InitializeReadWaiter(options N.ReadWaitOptions) (needCopy bool) {
c.readWaitOptions.Store(&options)
return false
}
func (c *udpNatConn) WaitReadPacket() (buffer *buf.Buffer, destination M.Socksaddr, err error) {
return c.waitReadPacket(c.loadReadWaitOptions())
}
func (c *udpNatConn) waitReadPacket(options N.ReadWaitOptions) (buffer *buf.Buffer, destination M.Socksaddr, err error) {
select {
case packet := <-c.packetChan:
buffer = options.Copy(packet.Buffer)
destination = packet.Destination
N.PutPacketBuffer(packet)
return
case <-c.doneChan:
return nil, M.Socksaddr{}, io.ErrClosedPipe
case <-c.readDeadline.Wait():
return nil, M.Socksaddr{}, os.ErrDeadlineExceeded
}
}
func (c *udpNatConn) CreatePacketBatchReadWaiter() (N.PacketBatchReadWaiter, bool) {
return c, true
}
func (c *udpNatConn) WaitReadPackets() (buffers []*buf.Buffer, destinations []M.Socksaddr, err error) {
options := c.loadReadWaitOptions()
buffer, destination, err := c.waitReadPacket(options)
if err != nil {
return nil, nil, err
}
batchSize := options.BatchSize
if batchSize <= 0 {
batchSize = 1
}
batch := c.readBatch
if batch == nil {
batch = new(udpNatReadBatch)
c.readBatch = batch
} else {
clear(batch.buffers)
clear(batch.destinations)
}
buffers = batch.buffers[:0]
destinations = batch.destinations[:0]
defer func() {
batch.buffers = buffers
batch.destinations = destinations
}()
buffers = append(buffers, buffer)
destinations = append(destinations, destination)
for len(buffers) < batchSize {
select {
case packet := <-c.packetChan:
buffers = append(buffers, options.Copy(packet.Buffer))
destinations = append(destinations, packet.Destination)
N.PutPacketBuffer(packet)
default:
return
}
}
return
}
func (c *udpNatConn) Timeout() time.Duration {
rawConn, lifetime, loaded := c.service.cache.PeekWithLifetime(c.key)
if !loaded || rawConn != c {
return 0
}
if lifetime.UnixMilli() == 0 {
return 0
}
return time.Until(lifetime)
}
func (c *udpNatConn) SetTimeout(timeout time.Duration) bool {
updated := c.service.cache.UpdateLifetime(c.key, c, timeout)
if !updated {
return false
}
if timeout == 0 {
c.service.cleanup.remove(c.cleanupEntry)
} else {
c.service.cleanup.addOrUpdate(c.cleanupEntry, time.Now().Add(timeout))
}
return true
}
func (c *udpNatConn) Close() error {
c.close()
if c.service.state.Load() == udpNatStateStarted {
c.service.cleanup.addOrUpdate(c.cleanupEntry, time.Now())
}
return nil
}
func (c *udpNatConn) close() {
c.closeOnce.Do(func() {
c.packetAccess.Lock()
close(c.doneChan)
drained := false
for !drained {
select {
case packet := <-c.packetChan:
packet.Buffer.Release()
N.PutPacketBuffer(packet)
default:
drained = true
}
}
c.packetAccess.Unlock()
c.clearFilterPeers()
if c.interfaceIndex != 0 {
c.service.unregisterClass(c)
}
})
}
func (c *udpNatConn) closeFromCache() {
c.close()
c.service.cleanup.remove(c.cleanupEntry)
}
func (c *udpNatConn) isClosed() bool {
select {
case <-c.doneChan:
return true
default:
return false
}
}
func (c *udpNatConn) LocalAddr() net.Addr {
return c.localAddr
}
func (c *udpNatConn) RemoteAddr() net.Addr {
return M.Socksaddr{}
}
func (c *udpNatConn) SetDeadline(t time.Time) error {
return os.ErrInvalid
}
func (c *udpNatConn) SetReadDeadline(t time.Time) error {
c.readDeadline.Set(t)
return nil
}
func (c *udpNatConn) SetWriteDeadline(t time.Time) error {
return os.ErrInvalid
}
func (c *udpNatConn) Upstream() any {
return c.writer
}

219
udp_nat_cleanup.go Normal file
View file

@ -0,0 +1,219 @@
package tun
import (
"container/heap"
"os"
"sync"
"time"
)
const (
udpNatStateCreated uint32 = iota
udpNatStateStarted
udpNatStateClosed
)
func (s *UDPNat) Start() error {
s.lifecycleAccess.Lock()
defer s.lifecycleAccess.Unlock()
switch s.state.Load() {
case udpNatStateCreated:
if s.interfaceFinder != nil {
s.interfaceElement = s.interfaceFinder.RegisterInterfaceUpdateCallback(s.updateInterfaces)
s.reloadInterfaces()
}
s.state.Store(udpNatStateStarted)
s.cleanupWait.Add(1)
go s.cleanupLoop()
return nil
case udpNatStateStarted:
return nil
default:
return os.ErrClosed
}
}
type udpNatCleanupEntry struct {
conn *udpNatConn
deadline time.Time
index int
}
type udpNatCleanupQueue struct {
service *UDPNat
access sync.Mutex
wake chan struct{}
entries udpNatCleanupHeap
}
func newUDPNatCleanupQueue(service *UDPNat) *udpNatCleanupQueue {
queue := &udpNatCleanupQueue{
service: service,
wake: make(chan struct{}, 1),
}
return queue
}
func (q *udpNatCleanupQueue) notify() {
select {
case q.wake <- struct{}{}:
default:
}
}
func (q *udpNatCleanupQueue) addOrUpdate(entry *udpNatCleanupEntry, deadline time.Time) {
if entry == nil || q.service.state.Load() == udpNatStateClosed {
return
}
q.access.Lock()
now := time.Now()
if entry.conn.isClosed() && deadline.After(now) {
deadline = now
}
entry.deadline = deadline
if entry.index == -1 {
heap.Push(&q.entries, entry)
} else {
heap.Fix(&q.entries, entry.index)
}
q.access.Unlock()
q.notify()
}
func (q *udpNatCleanupQueue) remove(entry *udpNatCleanupEntry) {
if entry == nil {
return
}
q.access.Lock()
if entry.index != -1 {
heap.Remove(&q.entries, entry.index)
}
q.access.Unlock()
q.notify()
}
func (q *udpNatCleanupQueue) next() (time.Time, bool) {
q.access.Lock()
defer q.access.Unlock()
if len(q.entries) == 0 {
return time.Time{}, false
}
return q.entries[0].deadline, true
}
func (q *udpNatCleanupQueue) popDue(now time.Time) *udpNatCleanupEntry {
q.access.Lock()
defer q.access.Unlock()
if len(q.entries) == 0 || q.entries[0].deadline.After(now) {
return nil
}
return heap.Pop(&q.entries).(*udpNatCleanupEntry)
}
func (q *udpNatCleanupQueue) clear() {
q.access.Lock()
for _, entry := range q.entries {
entry.index = -1
}
clear(q.entries)
q.entries = nil
q.access.Unlock()
q.notify()
}
type udpNatCleanupHeap []*udpNatCleanupEntry
func (h udpNatCleanupHeap) Len() int {
return len(h)
}
func (h udpNatCleanupHeap) Less(i int, j int) bool {
return h[i].deadline.Before(h[j].deadline)
}
func (h udpNatCleanupHeap) Swap(i int, j int) {
h[i], h[j] = h[j], h[i]
h[i].index = i
h[j].index = j
}
func (h *udpNatCleanupHeap) Push(value any) {
entry := value.(*udpNatCleanupEntry)
entry.index = len(*h)
*h = append(*h, entry)
}
func (h *udpNatCleanupHeap) Pop() any {
oldItems := *h
lastIndex := len(oldItems) - 1
entry := oldItems[lastIndex]
oldItems[lastIndex] = nil
entry.index = -1
*h = oldItems[:lastIndex]
return entry
}
func (s *UDPNat) cleanupLoop() {
defer s.cleanupWait.Done()
timer := time.NewTimer(time.Hour)
stopUDPNatCleanupTimer(timer)
defer timer.Stop()
for {
select {
case <-s.cleanup.wake:
default:
}
deadline, loaded := s.cleanup.next()
if !loaded {
select {
case <-s.cleanupDone:
return
case <-s.cleanup.wake:
continue
}
}
waitDuration := time.Until(deadline)
if waitDuration > 0 {
timer.Reset(waitDuration)
select {
case <-s.cleanupDone:
stopUDPNatCleanupTimer(timer)
return
case <-s.cleanup.wake:
stopUDPNatCleanupTimer(timer)
continue
case <-timer.C:
}
}
for {
entry := s.cleanup.popDue(time.Now())
if entry == nil {
break
}
s.cleanupEntry(entry)
}
}
}
func (s *UDPNat) cleanupEntry(entry *udpNatCleanupEntry) {
conn, lifetime, loaded := s.cache.PeekWithLifetime(entry.conn.key)
if !loaded || conn != entry.conn {
return
}
if lifetime.UnixMilli() == 0 {
return
}
if conn.isClosed() {
lifetime = time.Now()
}
s.cleanup.addOrUpdate(entry, lifetime)
}
func stopUDPNatCleanupTimer(timer *time.Timer) {
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
}