diff --git a/flow.go b/flow.go index c735fc9..1b1f34d 100644 --- a/flow.go +++ b/flow.go @@ -21,7 +21,6 @@ const ( ActionReject ActionDrop ActionBypass - ActionHijackDNS ) type FlowTracker interface { diff --git a/flow_dispatch.go b/flow_dispatch.go index e872177..3f71751 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -3,7 +3,6 @@ package tun import ( "maps" "net/netip" - "sync" "sync/atomic" "time" @@ -114,14 +113,12 @@ 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] + table map[flowKey]*flowEntry + lastSweep int64 + ports map[Port]*portNAT + natList atomic.Pointer[[]*portNAT] + revNAT atomic.Pointer[map[netip.Addr]*portNAT] activeNATs []*portNAT writebackBatch [][]byte @@ -170,41 +167,26 @@ 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 { - flows = append(flows, entry.flow) + entry.flow.close(FlowCloseReset) } } - ports := make([]Port, 0, len(d.ports)) for port, nat := range d.ports { if nat != nil { - ports = append(ports, port) + port.DetachReturn(&d.returnPath) } } - 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 || d.returnPath.closed.Load() { + if d == nil { 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] @@ -213,16 +195,13 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool { loaded = false } if loaded { - handled := d.handleHit(key, entry, &parsed, packet, now) - d.access.RUnlock() - return handled + return d.handleHit(key, entry, &parsed, packet, now) } - 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) + return d.judgeAndInstall(key, &parsed, packet, now) } func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool { @@ -276,18 +255,12 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for } } -func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte) bool { +func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) 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 { @@ -318,13 +291,6 @@ 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 @@ -571,27 +537,10 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) { } } -func (d *ForwardDispatcher) ResetNetwork() { +func (d *ForwardDispatcher) Flush() { 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) } diff --git a/flow_dns.go b/flow_dns.go deleted file mode 100644 index ee9a936..0000000 --- a/flow_dns.go +++ /dev/null @@ -1,89 +0,0 @@ -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}) -} diff --git a/go.mod b/go.mod index fbba8e2..8bab948 100644 --- a/go.mod +++ b/go.mod @@ -1,32 +1,32 @@ module github.com/sagernet/sing-tun -go 1.25.0 +go 1.24.7 require ( - github.com/florianl/go-nfqueue/v2 v2.1.0 + github.com/florianl/go-nfqueue/v2 v2.0.2 github.com/go-ole/go-ole v1.3.0 github.com/google/btree v1.1.3 - 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/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/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a - github.com/sagernet/nftables v0.3.0-mod.4 - github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8 + github.com/sagernet/nftables v0.3.0-mod.2 + github.com/sagernet/sing v0.8.0 github.com/stretchr/testify v1.11.1 go4.org/netipx v0.0.0-20231129151722-fdeea329fbba - golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc - golang.org/x/net v0.57.0 - golang.org/x/sys v0.47.0 + golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 + golang.org/x/net v0.50.0 + golang.org/x/sys v0.41.0 ) require ( github.com/davecgh/go-spew v1.1.1 // indirect - github.com/fsnotify/fsnotify v1.9.0 // indirect + github.com/fsnotify/fsnotify v1.7.0 // indirect github.com/google/go-cmp v0.7.0 // indirect - github.com/mdlayher/socket v0.6.0 // indirect + github.com/mdlayher/socket v0.5.1 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/vishvananda/netns v0.0.4 // indirect - golang.org/x/sync v0.20.0 // indirect - golang.org/x/time v0.15.0 // indirect + golang.org/x/sync v0.7.0 // indirect + golang.org/x/time v0.7.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index 0c72db9..54ed5fd 100644 --- a/go.sum +++ b/go.sum @@ -1,50 +1,48 @@ 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.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/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/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/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/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/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.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/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/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.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/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/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-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/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/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -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= +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= 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= diff --git a/gtcpip/header/ipv4.go b/gtcpip/header/ipv4.go index 3041eae..d5ffbf1 100644 --- a/gtcpip/header/ipv4.go +++ b/gtcpip/header/ipv4.go @@ -22,6 +22,7 @@ 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 @@ -334,7 +335,7 @@ func (b IPv4) FragmentOffset() uint16 { } func (b IPv4) FragmentOffsetDarwinRaw() uint16 { - return binary.NativeEndian.Uint16(b[flagsFO:]) << 3 + return common.NativeEndian.Uint16(b[flagsFO:]) << 3 } // TotalLength returns the "total length" field of the IPv4 header. @@ -343,7 +344,7 @@ func (b IPv4) TotalLength() uint16 { } func (b IPv4) TotalLengthDarwinRaw() uint16 { - return binary.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength()) + return common.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength()) } // Checksum returns the checksum field of the IPv4 header. @@ -440,7 +441,7 @@ func (b IPv4) SetTotalLength(totalLength uint16) { } func (b IPv4) SetTotalLengthDarwinRaw(totalLength uint16) { - binary.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength) + common.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength) } // SetChecksum sets the checksum field of the IPv4 header. @@ -457,7 +458,7 @@ func (b IPv4) SetFlagsFragmentOffset(flags uint8, offset uint16) { func (b IPv4) SetFlagsFragmentOffsetDarwinRaw(flags uint8, offset uint16) { v := (uint16(flags) << 13) | (offset >> 3) - binary.NativeEndian.PutUint16(b[flagsFO:], v) + common.NativeEndian.PutUint16(b[flagsFO:], v) } // SetID sets the identification field. @@ -1178,7 +1179,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] - clear(b) + common.ClearArray(b) return padded } diff --git a/gtcpip/header/ipv6_extension_headers.go b/gtcpip/header/ipv6_extension_headers.go index 1ab7c9d..6c48b1b 100644 --- a/gtcpip/header/ipv6_extension_headers.go +++ b/gtcpip/header/ipv6_extension_headers.go @@ -21,6 +21,7 @@ import ( "math" "github.com/sagernet/sing-tun/gtcpip" + "github.com/sagernet/sing/common" ) // IPv6ExtensionHeaderIdentifier is an IPv6 extension header identifier. @@ -128,7 +129,7 @@ func padIPv6Option(b []byte) { b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6Pad1ExtHdrOptionIdentifier) default: // Pad with PadN. s := b[ipv6ExtHdrOptionPayloadOffset:] - clear(s) + common.ClearArray(s) b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6PadNExtHdrOptionIdentifier) b[ipv6ExtHdrOptionLengthOffset] = uint8(len(s)) } diff --git a/gtcpip/header/ndp_options.go b/gtcpip/header/ndp_options.go index ca1c6cb..c545120 100644 --- a/gtcpip/header/ndp_options.go +++ b/gtcpip/header/ndp_options.go @@ -24,6 +24,7 @@ import ( "time" "github.com/sagernet/sing-tun/gtcpip" + "github.com/sagernet/sing/common" ) // ndpOptionIdentifier is an NDP option type identifier. @@ -340,7 +341,7 @@ func (b NDPOptions) Serialize(s NDPOptionsSerializer) int { // Zero out remaining (padding) bytes, if any exists. if used+2 < l { - clear(b[used+2 : l]) + common.ClearArray(b[used+2 : l]) } b = b[l:] @@ -566,7 +567,7 @@ func (o NDPPrefixInformation) serializeInto(b []byte) int { // Zero out the Reserved2 field. reserved2 := b[ndpPrefixInformationReserved2Offset:][:ndpPrefixInformationReserved2Length] - clear(reserved2) + common.ClearArray(reserved2) return used } @@ -685,7 +686,7 @@ func (o NDPRecursiveDNSServer) serializeInto(b []byte) int { used := copy(b, o) // Zero out the reserved bytes that are before the Lifetime field. - clear(b[0:ndpRecursiveDNSServerLifetimeOffset]) + common.ClearArray(b[0:ndpRecursiveDNSServerLifetimeOffset]) return used } @@ -778,7 +779,7 @@ func (o NDPDNSSearchList) serializeInto(b []byte) int { used := copy(b, o) // Zero out the reserved bytes that are before the Lifetime field. - clear(b[0:ndpDNSSearchListLifetimeOffset]) + common.ClearArray(b[0:ndpDNSSearchListLifetimeOffset]) return used } diff --git a/internal/fdbased_darwin/endpoint.go b/internal/fdbased_darwin/endpoint.go index b3292a2..05371e7 100644 --- a/internal/fdbased_darwin/endpoint.go +++ b/internal/fdbased_darwin/endpoint.go @@ -50,6 +50,7 @@ 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" ) @@ -199,6 +200,10 @@ 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 @@ -252,6 +257,10 @@ 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") } @@ -292,7 +301,7 @@ func New(opts *Options) (stack.LinkEndpoint, error) { e.fds = append(e.fds, fdInfo{fd: fd, isSocket: true}) if opts.ProcessorsPerChannel == 0 { - opts.ProcessorsPerChannel = max(1, runtime.GOMAXPROCS(0)/len(opts.FDs)) + opts.ProcessorsPerChannel = common.Max(1, runtime.GOMAXPROCS(0)/len(opts.FDs)) } inboundDispatcher, err := newRecvMMsgDispatcher(fd, e, opts) diff --git a/internal/fdbased_darwin/processors.go b/internal/fdbased_darwin/processors.go index eb4cb2c..1ceca83 100644 --- a/internal/fdbased_darwin/processors.go +++ b/internal/fdbased_darwin/processors.go @@ -214,47 +214,34 @@ 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 - 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() + var tcpHdr header.TCP + if tcpip.TransportProtocolNumber(ipHdr.NextHeader()) == header.TCPProtocolNumber { + tcpHdr = header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen]) } 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 - } - // 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() + 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() + cid.srcPort = tcpHdr.SourcePort() + cid.dstPort = tcpHdr.DestinationPort() + cid.proto = header.IPv6ProtocolNumber default: return cid, true } diff --git a/netns_linux.go b/netns_linux.go deleted file mode 100644 index e2b5706..0000000 --- a/netns_linux.go +++ /dev/null @@ -1,73 +0,0 @@ -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 -} diff --git a/netns_other.go b/netns_other.go deleted file mode 100644 index ab5a3a2..0000000 --- a/netns_other.go +++ /dev/null @@ -1,12 +0,0 @@ -//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) -} diff --git a/ping/cmsg_windows.go b/ping/cmsg_windows.go index be5be9b..07c322c 100644 --- a/ping/cmsg_windows.go +++ b/ping/cmsg_windows.go @@ -1,10 +1,11 @@ package ping import ( - "encoding/binary" "fmt" "unsafe" + "github.com/sagernet/sing/common" + "golang.org/x/net/ipv6" "golang.org/x/sys/windows" ) @@ -36,9 +37,9 @@ func parseIPv6ControlMessage(cmsg []byte) (*ipv6.ControlMessage, error) { } switch cmsghdr.Type { case IPV6_TCLASS: - controlMessage.TrafficClass = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) + controlMessage.TrafficClass = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) case IPV6_HOPLIMIT: - controlMessage.HopLimit = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) + controlMessage.HopLimit = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) } cmsg = cmsg[msgSize:] } diff --git a/ping/socket_linux_unprivileged.go b/ping/socket_linux_unprivileged.go index 1ad1548..f709684 100644 --- a/ping/socket_linux_unprivileged.go +++ b/ping/socket_linux_unprivileged.go @@ -9,6 +9,7 @@ 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" @@ -174,7 +175,7 @@ func (c *UnprivilegedConn) Close() error { for _, conn := range c.mapping { _ = conn.Close() } - clear(c.mapping) + common.ClearMap(c.mapping) return nil } diff --git a/redirect_linux.go b/redirect_linux.go index f08d0f7..04a1fee 100644 --- a/redirect_linux.go +++ b/redirect_linux.go @@ -26,7 +26,6 @@ type autoRedirect struct { logger logger.Logger tableName string networkMonitor NetworkUpdateMonitor - ownedNetworkMonitor bool networkListener *list.Element[NetworkUpdateCallback] interfaceFinder control.InterfaceFinder localAddresses []netip.Prefix @@ -52,7 +51,7 @@ type autoRedirect struct { } func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) { - r := &autoRedirect{ + return &autoRedirect{ tunOptions: options.TunOptions, ctx: options.Context, handler: options.Handler, @@ -64,11 +63,7 @@ func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) { customRedirectPortFunc: options.CustomRedirectPort, routeAddressSet: options.RouteAddressSet, routeExcludeAddressSet: options.RouteExcludeAddressSet, - } - if options.TunOptions.NetNs != "" { - r.interfaceFinder = &networkNamespaceInterfaceFinder{control.NewDefaultInterfaceFinder(), options.TunOptions} - } - return r, nil + }, nil } func (r *autoRedirect) Start() error { @@ -94,11 +89,8 @@ 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 = runInNetworkNamespace(r.tunOptions.NetNs, r.initializeNFTables) + err = r.initializeNFTables() if err != nil { return E.Cause(err, "missing nftables support") } @@ -140,7 +132,7 @@ func (r *autoRedirect) Start() error { listenAddr = netip.IPv4Unspecified() } server := newRedirectServer(r.ctx, r.handler, r.logger, listenAddr) - err = runInNetworkNamespace(r.tunOptions.NetNs, server.Start) + err = server.Start() if err != nil { return E.Cause(err, "start redirect server") } @@ -159,43 +151,24 @@ 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 = runInNetworkNamespace(r.tunOptions.NetNs, handler.Start); err != nil { + } else if err = 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 } } - if r.tunOptions.NetNs != "" { - var monitor NetworkUpdateMonitor - monitor, err = NewNetworkUpdateMonitor(r.logger) - if err != nil { - 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 { - setupErr = r.setupRedirectRoutes() - if setupErr != nil { - r.cleanupNFTables() - return E.Cause(setupErr, "setup redirect routes") - } - } - return nil - }) + r.cleanupNFTables() + err = r.setupNFTables() if err != nil { - return err + return E.Cause(err, "setup nftables") + } + if r.tunOptions.AutoRedirectMarkMode { + err = r.setupRedirectRoutes() + if err != nil { + r.cleanupNFTables() + return E.Cause(err, "setup redirect routes") + } } } else { r.cleanupIPTables() @@ -212,14 +185,8 @@ 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() - } + r.cleanupNFTables() + r.cleanupRedirectRoutes() } else { r.cleanupIPTables() } @@ -230,7 +197,7 @@ func (r *autoRedirect) Close() error { func (r *autoRedirect) UpdateRouteAddressSet() { if r.useNFTables { - err := runInNetworkNamespace(r.tunOptions.NetNs, r.nftablesUpdateRouteAddressSet) + err := r.nftablesUpdateRouteAddressSet() if err != nil { r.logger.Error("update route address set: ", err) } diff --git a/redirect_nftables.go b/redirect_nftables.go index c71a770..5944e4e 100644 --- a/redirect_nftables.go +++ b/redirect_nftables.go @@ -299,38 +299,27 @@ 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.startDockerFirewallMonitor() + err = r.configureDockerFirewall(false) + if err != nil && r.logger != nil { + r.logger.Warn("configure docker firewall: ", err) } r.networkListener = r.networkMonitor.RegisterCallback(func() { - updateErr := runInNetworkNamespace(r.tunOptions.NetNs, r.updateNetworkAddresses) - if updateErr != nil { - r.logger.Error(updateErr) + 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) + } } }) 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() @@ -387,7 +376,6 @@ 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() @@ -401,11 +389,9 @@ 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) - } + err = r.configureDockerFirewall(true) + if err != nil && r.logger != nil { + r.logger.Warn("cleanup docker firewall: ", err) } } diff --git a/stack.go b/stack.go index 613e45d..eaf2405 100644 --- a/stack.go +++ b/stack.go @@ -14,7 +14,6 @@ import ( type Stack interface { Start() error - ResetNetwork() Close() error } @@ -24,9 +23,6 @@ type StackOptions struct { TunOptions Options UDPTimeout time.Duration ICMPTimeout time.Duration - UDPMapping NATMapping - UDPFiltering NATFiltering - UDPNATMax uint32 Handler Handler Logger logger.Logger ForwarderBindInterface bool diff --git a/stack_gvisor.go b/stack_gvisor.go index c226d05..03b2873 100644 --- a/stack_gvisor.go +++ b/stack_gvisor.go @@ -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,7 +44,6 @@ type GVisor struct { endpoint stack.LinkEndpoint dispatcher *ForwardDispatcher icmpForwarder *ICMPForwarder - udpForwarder *UDPForwarder } type GVisorTun interface { @@ -79,18 +78,11 @@ 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, + broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address), + handler: options.Handler, + logger: options.Logger, } return gStack, nil } @@ -101,7 +93,7 @@ func (t *GVisor) Start() error { return err } if t.handler != nil { - t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpNATOptions.Timeout, t.icmpTimeout) + t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout) } linkEndpoint = &LinkEndpointFilter{ LinkEndpoint: linkEndpoint, @@ -118,13 +110,7 @@ func (t *GVisor) Start() error { return err } ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).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 + ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket) icmpForwarder := NewICMPForwarder(ipStack, t.handler, t.logger) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket) @@ -134,24 +120,11 @@ 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 } diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 70e27ec..11e82af 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -72,15 +72,6 @@ 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() @@ -155,15 +146,9 @@ 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 { + if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 { 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, diff --git a/stack_gvisor_udp.go b/stack_gvisor_udp.go index 5cd0c93..2ae54cf 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -8,6 +8,7 @@ import ( "net/netip" "os" "sync" + "time" _ "unsafe" "github.com/sagernet/gvisor/pkg/buffer" @@ -20,35 +21,26 @@ 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 + udpNat *udpnat.Service } -func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, options UDPNatOptions) *UDPForwarder { +func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, timeout time.Duration) *UDPForwarder { forwarder := &UDPForwarder{ ctx: ctx, stack: stack, handler: handler, } - options.Handler = handler - options.Prepare = forwarder.PreparePacketConnection - forwarder.udpNat = NewUDPNat(options) + forwarder.udpNat = udpnat.New(handler, forwarder.PreparePacketConnection, timeout, false) 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) @@ -71,26 +63,18 @@ 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 - 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 + } + var sourceNetwork tcpip.NetworkProtocolNumber + if source.Addr.Is4() { + sourceNetwork = header.IPv4ProtocolNumber + } else { + sourceNetwork = header.IPv6ProtocolNumber } writer := &UDPBackWriter{ stack: f.stack, diff --git a/stack_mixed.go b/stack_mixed.go index 69c8b27..4680380 100644 --- a/stack_mixed.go +++ b/stack_mixed.go @@ -19,10 +19,9 @@ import ( type Mixed struct { *System - tun GVisorTun - stack *stack.Stack - endpoint *channel.Endpoint - udpForwarder *UDPForwarder + tun GVisorTun + stack *stack.Stack + endpoint *channel.Endpoint } func NewMixed( @@ -48,13 +47,7 @@ func (m *Mixed) Start() error { if err != nil { return err } - 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 + ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpTimeout).HandlePacket) m.stack = ipStack m.endpoint = endpoint go m.tunLoop() @@ -62,20 +55,10 @@ 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() { diff --git a/stack_system.go b/stack_system.go index 41644bd..3cb0cb0 100644 --- a/stack_system.go +++ b/stack_system.go @@ -5,10 +5,7 @@ import ( "errors" "net" "net/netip" - "os" "slices" - "sync" - "sync/atomic" "syscall" "time" @@ -22,6 +19,7 @@ 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") @@ -30,7 +28,6 @@ type System struct { ctx context.Context tun Tun tunName string - netNs string mtu int handler Handler logger logger.Logger @@ -47,25 +44,16 @@ type System struct { icmpTimeout time.Duration tcpListener net.Listener tcpListener6 net.Listener - // 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 - udpNATOptions UDPNatOptions - dispatcher *ForwardDispatcher - bindInterface bool - interfaceFinder control.InterfaceFinder - frontHeadroom int - txChecksumOffload bool - multiPendingPackets bool + tcpPort uint16 + tcpPort6 uint16 + tcpNat *TCPNat + udpNat *udpnat.Service + dispatcher *ForwardDispatcher + bindInterface bool + interfaceFinder control.InterfaceFinder + frontHeadroom int + txChecksumOffload bool + multiPendingPackets bool } type Session struct { @@ -80,7 +68,6 @@ 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, @@ -91,17 +78,9 @@ 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, + bindInterface: options.ForwarderBindInterface, + interfaceFinder: options.InterfaceFinder, + multiPendingPackets: options.TunOptions.EXP_MultiPendingPackets, } if len(options.TunOptions.Inet4Address) > 0 { if !HasNextAddress(options.TunOptions.Inet4Address[0], 1) { @@ -123,26 +102,8 @@ 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, @@ -158,10 +119,8 @@ func (s *System) Start() error { return nil } -// 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) { +func (s *System) start() error { + _ = fixWindowsFirewall() var listener net.ListenConfig if s.bindInterface { listener.Control = control.Append(listener.Control, func(network, address string, conn syscall.RawConn) error { @@ -172,60 +131,40 @@ func (s *System) listenTCP(isIPv6 bool) (net.Listener, error) { return nil }) } - network := "tcp4" - address := s.inet4Address - if isIPv6 { - network = "tcp6" - address = s.inet6Address - } - var ( - tcpListener net.Listener - err error - ) - for range 3 { - 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) + for range 3 { + tcpListener, err = listener.Listen(s.ctx, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0")) + if !retryableListenError(err) { + break + } + time.Sleep(time.Second) + } if err != nil { return err } s.tcpListener = tcpListener - s.tcpPort.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) - go s.acceptLoop(tcpListener, false) + s.tcpPort = M.SocksaddrFromNet(tcpListener.Addr()).Port + go s.acceptLoop(tcpListener) } if s.inet6NextAddress.IsValid() { - tcpListener, err = s.listenTCP(true) + for range 3 { + tcpListener, err = listener.Listen(s.ctx, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0")) + if !retryableListenError(err) { + break + } + time.Sleep(time.Second) + } if err != nil { return err } s.tcpListener6 = tcpListener - s.tcpPort6.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) - go s.acceptLoop(tcpListener, true) + s.tcpPort6 = M.SocksaddrFromNet(tcpListener.Addr()).Port + go s.acceptLoop(tcpListener) } s.tcpNat = NewNat(s.ctx, s.udpTimeout) - udpNATOptions := s.udpNATOptions - udpNATOptions.Handler = s.handler - udpNATOptions.Prepare = s.preparePacketConnection - s.udpNat = NewUDPNat(udpNATOptions) - err = s.udpNat.Start() - if err != nil { - return err - } + s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false) if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN { s.frontHeadroom = linuxTUN.FrontHeadroom() s.txChecksumOffload = linuxTUN.TXChecksumOffload() @@ -397,28 +336,11 @@ func (s *System) processPacket(packet []byte) bool { return writeBack } -func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) { +func (s *System) acceptLoop(listener net.Listener) { 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 + return } connPort := M.SocksaddrFromNet(conn.RemoteAddr()).Port session := s.tcpNat.LookupBack(connPort) @@ -430,47 +352,6 @@ func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) { } } -// 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: @@ -480,7 +361,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() == uint16(s.tcpPort.Load()) { + header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort { return false } case header.ICMPv4ProtocolNumber: @@ -499,7 +380,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() == uint16(s.tcpPort6.Load()) { + header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 { return false } case header.ICMPv6ProtocolNumber: @@ -563,7 +444,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() == uint16(s.tcpPort.Load()) { + } else if source.Addr() == s.inet4Address && source.Port() == s.tcpPort { session := s.tcpNat.LookupBack(destination.Port()) if session == nil { return false, E.New("ipv4: tcp: session not found: ", destination.Port()) @@ -589,7 +470,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err } rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, s.inet4NextAddress, natPort, true, - s.inet4Address, uint16(s.tcpPort.Load()), true) + s.inet4Address, s.tcpPort, true) } } return true, nil @@ -600,7 +481,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() == uint16(s.tcpPort6.Load()) { + } else if source.Addr() == s.inet6Address && source.Port() == s.tcpPort6 { session := s.tcpNat.LookupBack(destination.Port()) if session == nil { return false, E.New("ipv6: tcp: session not found: ", destination.Port()) @@ -626,7 +507,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err } rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, s.inet6NextAddress, natPort, true, - s.inet6Address, uint16(s.tcpPort6.Load()), true) + s.inet6Address, s.tcpPort6, true) } } return true, nil @@ -801,22 +682,20 @@ type systemUDPPacketWriter4 struct { txChecksumOffload bool } -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())) +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())) ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetSourceAddr(destination.Addr) udpHdr := header.UDP(ipHdr.Payload()) udpHdr.SetDestinationPort(udpHdr.SourcePort()) udpHdr.SetSourcePort(destination.Port) - udpHdr.SetLength(uint16(payloadLen + header.UDPMinimumSize)) + udpHdr.SetLength(uint16(buffer.Len() + header.UDPMinimumSize)) if !w.txChecksumOffload { udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum( header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()), @@ -825,61 +704,12 @@ func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M 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(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 + PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version) + } else { + newPacket.Advance(-w.frontHeadroom) } + return common.Error(w.tun.Write(newPacket.Bytes())) } type systemUDPPacketWriter6 struct { @@ -890,16 +720,14 @@ type systemUDPPacketWriter6 struct { txChecksumOffload bool } -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) +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()) ipHdr.SetPayloadLength(udpLen) ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetSourceAddr(destination.Addr) @@ -914,61 +742,12 @@ func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M } 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(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 + PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version) + } else { + newPacket.Advance(-w.frontHeadroom) } + return common.Error(w.tun.Write(newPacket.Bytes())) } func newSystemWriteback(tunInterface Tun, frontHeadroom int) ForwardWriteback { diff --git a/stack_system_nat.go b/stack_system_nat.go index 1dd5377..2fec29c 100644 --- a/stack_system_nat.go +++ b/stack_system_nat.go @@ -86,15 +86,6 @@ 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] diff --git a/stack_system_packet.go b/stack_system_packet.go index d00b95d..a8f8076 100644 --- a/stack_system_packet.go +++ b/stack_system_packet.go @@ -5,6 +5,7 @@ import ( "syscall" "github.com/sagernet/sing-tun/gtcpip/header" + "github.com/sagernet/sing/common" ) func PacketIPVersion(packet []byte) int { @@ -13,7 +14,7 @@ func PacketIPVersion(packet []byte) int { func PacketFillHeader(packet []byte, ipVersion int) { if PacketOffset > 0 { - clear(packet[:3]) + common.ClearArray(packet[:3]) switch ipVersion { case header.IPv4Version: packet[3] = syscall.AF_INET diff --git a/stack_system_selfheal_test.go b/stack_system_selfheal_test.go deleted file mode 100644 index ab063c2..0000000 --- a/stack_system_selfheal_test.go +++ /dev/null @@ -1,117 +0,0 @@ -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) - } -} diff --git a/tun.go b/tun.go index e770122..c6518f4 100644 --- a/tun.go +++ b/tun.go @@ -14,14 +14,12 @@ 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 } @@ -68,7 +66,6 @@ const ( type Options struct { Name string - NetNs string Inet4Address []netip.Prefix Inet6Address []netip.Prefix MTU uint32 diff --git a/tun_linux.go b/tun_linux.go index 41d4dc4..487051d 100644 --- a/tun_linux.go +++ b/tun_linux.go @@ -51,38 +51,37 @@ 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") - } - tunLink, err := netlink.LinkByName(options.Name) - if err != nil { - return nil, E.Errors(err, unix.Close(tunFd)) - } - nativeTun := &NativeTun{ - tunFd: tunFd, - tunFile: os.NewFile(uintptr(tunFd), "tun"), - options: options, - } - err = nativeTun.configure(tunLink) - if err != nil { - return nil, E.Errors(err, unix.Close(tunFd)) - } - return nativeTun, nil - }) - } - nativeTun := &NativeTun{ - tunFd: options.FileDescriptor, - tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"), - options: options, - } - if options.GSO { - err := nativeTun.enableGSO() + tunFd, err := open(options.Name, options.GSO) if err != nil { - if options.Logger != nil { - options.Logger.Warn(err) + return nil, E.Cause(err, "open tun") + } + tunLink, err := netlink.LinkByName(options.Name) + if err != nil { + return nil, E.Errors(err, unix.Close(tunFd)) + } + nativeTun = &NativeTun{ + tunFd: tunFd, + tunFile: os.NewFile(uintptr(tunFd), "tun"), + options: options, + } + err = nativeTun.configure(tunLink) + if err != nil { + return nil, E.Errors(err, unix.Close(tunFd)) + } + } else { + nativeTun = &NativeTun{ + tunFd: options.FileDescriptor, + tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"), + options: options, + } + if options.GSO { + err := nativeTun.enableGSO() + if err != nil { + if options.Logger != nil { + options.Logger.Warn(err) + } } } } @@ -291,10 +290,10 @@ func (t *NativeTun) Name() (string, error) { func (t *NativeTun) Start() error { if t.options.FileDescriptor == 0 { - if !t.options.EXP_ExternalConfiguration && t.options.NetNs == "" { + if !t.options.EXP_ExternalConfiguration { t.options.InterfaceMonitor.RegisterMyInterface(t.options.Name) } - err := runInNetworkNamespace(t.options.NetNs, t.start) + err := t.start() if err != nil { return err } @@ -355,7 +354,7 @@ func (t *NativeTun) start() error { return E.Cause(err, "set rules") } - if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { + if t.options.DNSMode != DNSModeDisabled { err = t.setSearchDomainForSystemdResolved() if err != nil { return E.Cause(err, "set search domain") @@ -375,13 +374,11 @@ func (t *NativeTun) Close() error { if t.options.EXP_ExternalConfiguration { return common.Close(common.PtrOrNil(t.tunFile)) } - if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { + if t.options.DNSMode != DNSModeDisabled { 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))) + t.unsetAddresses() + return E.Errors(t.unsetRoute(), t.unsetRules(), common.Close(common.PtrOrNil(t.tunFile))) } func (t *NativeTun) Read(p []byte) (n int, err error) { @@ -628,18 +625,16 @@ 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") - } - err = t.unsetRoute0(tunLink) - if err != nil { - return E.Cause(err, "unset old routes") - } - t.options = tunOptions - return t.setRoute(tunLink) - }) + tunLink, err := netlink.LinkByName(t.options.Name) + if err != nil { + return E.Cause(err, "find tun interface") + } + err = t.unsetRoute0(tunLink) + if err != nil { + return E.Cause(err, "unset old routes") + } + t.options = tunOptions + return t.setRoute(tunLink) } func (t *NativeTun) routes(tunLink netlink.Link) ([]netlink.Route, error) { diff --git a/udp_egress.go b/udp_egress.go deleted file mode 100644 index 0ace53e..0000000 --- a/udp_egress.go +++ /dev/null @@ -1,287 +0,0 @@ -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() - } - } -} diff --git a/udp_egress_conn.go b/udp_egress_conn.go deleted file mode 100644 index 93cd5d4..0000000 --- a/udp_egress_conn.go +++ /dev/null @@ -1,124 +0,0 @@ -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 -} diff --git a/udp_nat.go b/udp_nat.go deleted file mode 100644 index 6d5dd94..0000000 --- a/udp_nat.go +++ /dev/null @@ -1,789 +0,0 @@ -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 -} diff --git a/udp_nat_cleanup.go b/udp_nat_cleanup.go deleted file mode 100644 index c63b8af..0000000 --- a/udp_nat_cleanup.go +++ /dev/null @@ -1,219 +0,0 @@ -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: - } - } -}