diff --git a/flow.go b/flow.go index 1b1f34d..c735fc9 100644 --- a/flow.go +++ b/flow.go @@ -21,6 +21,7 @@ const ( ActionReject ActionDrop ActionBypass + ActionHijackDNS ) type FlowTracker interface { diff --git a/flow_dispatch.go b/flow_dispatch.go index 3f71751..e872177 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -3,6 +3,7 @@ package tun import ( "maps" "net/netip" + "sync" "sync/atomic" "time" @@ -113,12 +114,14 @@ type ForwardDispatcher struct { logger logger.Logger udpTimeout time.Duration icmpTimeout time.Duration + access sync.RWMutex - table map[flowKey]*flowEntry - lastSweep int64 - ports map[Port]*portNAT - natList atomic.Pointer[[]*portNAT] - revNAT atomic.Pointer[map[netip.Addr]*portNAT] + table map[flowKey]*flowEntry + lastSweep int64 + resetPending atomic.Bool + ports map[Port]*portNAT + natList atomic.Pointer[[]*portNAT] + revNAT atomic.Pointer[map[netip.Addr]*portNAT] activeNATs []*portNAT writebackBatch [][]byte @@ -167,26 +170,41 @@ func (d *ForwardDispatcher) Close() { return } d.returnPath.closed.Store(true) + d.access.Lock() + flows := make([]*forwardFlow, 0, len(d.table)) for _, entry := range d.table { if entry.flow != nil { - entry.flow.close(FlowCloseReset) + flows = append(flows, entry.flow) } } + ports := make([]Port, 0, len(d.ports)) for port, nat := range d.ports { if nat != nil { - port.DetachReturn(&d.returnPath) + ports = append(ports, port) } } + d.access.Unlock() + for _, flow := range flows { + flow.close(FlowCloseReset) + } + for _, port := range ports { + port.DetachReturn(&d.returnPath) + } } func (d *ForwardDispatcher) Dispatch(packet []byte) bool { - if d == nil { + if d == nil || d.returnPath.closed.Load() { return false } parsed, ok := parseForwardPacket(packet) if !ok || parsed.fragment || !parsed.hasFlow { return false } + d.access.RLock() + if d.returnPath.closed.Load() { + d.access.RUnlock() + return false + } key := parsed.flowKey() now := d.now() entry, loaded := d.table[key] @@ -195,13 +213,16 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool { loaded = false } if loaded { - return d.handleHit(key, entry, &parsed, packet, now) + handled := d.handleHit(key, entry, &parsed, packet, now) + d.access.RUnlock() + return handled } + d.access.RUnlock() if parsed.protocol == uint8(header.TCPProtocolNumber) && (parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) { return false } - return d.judgeAndInstall(key, &parsed, packet, now) + return d.judgeAndInstall(key, &parsed, packet) } func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool { @@ -255,12 +276,18 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for } } -func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool { +func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte) bool { var firstPacket []byte if packet.protocol == uint8(header.UDPProtocolNumber) { firstPacket = header.UDP(packet.transport).Payload() } verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination, firstPacket) + d.access.RLock() + defer d.access.RUnlock() + if d.returnPath.closed.Load() { + return false + } + now := d.now() switch verdict.Action { case ActionFlow: if verdict.Port != nil { @@ -291,6 +318,13 @@ func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, case ActionDrop: d.installSimple(key, ActionDrop, packet.protocol, now) return true + case ActionHijackDNS: + if packet.protocol == uint8(header.UDPProtocolNumber) { + d.hijackDNSPacket(packet) + return true + } + d.installSimple(key, ActionAccept, packet.protocol, now) + return false default: d.installSimple(key, ActionAccept, packet.protocol, now) return false @@ -537,10 +571,27 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) { } } -func (d *ForwardDispatcher) Flush() { +func (d *ForwardDispatcher) ResetNetwork() { if d == nil { return } + d.resetPending.Store(true) +} + +func (d *ForwardDispatcher) Flush() { + if d == nil || d.returnPath.closed.Load() { + return + } + d.access.RLock() + defer d.access.RUnlock() + if d.returnPath.closed.Load() { + return + } + if d.resetPending.Swap(false) { + for key, entry := range d.table { + d.removeEntry(key, entry, FlowCloseReset) + } + } for _, nat := range d.activeNATs { d.flushPort(nat) } diff --git a/flow_dns.go b/flow_dns.go new file mode 100644 index 0000000..ee9a936 --- /dev/null +++ b/flow_dns.go @@ -0,0 +1,89 @@ +package tun + +import ( + "net/netip" + + "github.com/sagernet/sing-tun/gtcpip/checksum" + "github.com/sagernet/sing-tun/gtcpip/header" + "github.com/sagernet/sing/common/buf" + E "github.com/sagernet/sing/common/exceptions" + M "github.com/sagernet/sing/common/metadata" + N "github.com/sagernet/sing/common/network" +) + +func (d *ForwardDispatcher) hijackDNSPacket(packet *forwardPacket) { + writer := &dnsResponseWriter{ + writeback: d.writeback, + source: packet.source, + } + d.handler.NewDNSPacket(header.UDP(packet.transport).Payload(), M.SocksaddrFromNetIP(packet.source), M.SocksaddrFromNetIP(packet.destination), writer) +} + +var _ N.PacketWriter = (*dnsResponseWriter)(nil) + +type dnsResponseWriter struct { + writeback ForwardWriteback + source netip.AddrPort +} + +func (w *dnsResponseWriter) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + defer buffer.Release() + if !destination.IsIP() { + return E.New("invalid destination: ", destination) + } + sourceAddr := w.source.Addr().Unmap() + destinationAddr := destination.Addr.Unmap() + headroom := w.writeback.ReturnHeadroom() + udpLen := header.UDPMinimumSize + buffer.Len() + var ( + packet []byte + udpHdr header.UDP + ipHdr header.Network + ) + if sourceAddr.Is4() { + if !destinationAddr.Is4() { + return E.New("send IPv6 packet to IPv4 connection") + } + size := header.IPv4MinimumSize + udpLen + packet = make([]byte, headroom+size) + inet4Hdr := header.IPv4(packet[headroom:]) + inet4Hdr.Encode(&header.IPv4Fields{ + TotalLength: uint16(size), + TTL: synthesizedTTL, + Protocol: uint8(header.UDPProtocolNumber), + SrcAddr: destinationAddr, + DstAddr: sourceAddr, + }) + udpHdr = header.UDP(inet4Hdr.Payload()) + ipHdr = inet4Hdr + } else { + if destinationAddr.Is4() { + destinationAddr = netip.AddrFrom16(destinationAddr.As16()) + } + size := header.IPv6MinimumSize + udpLen + packet = make([]byte, headroom+size) + inet6Hdr := header.IPv6(packet[headroom:]) + inet6Hdr.Encode(&header.IPv6Fields{ + PayloadLength: uint16(udpLen), + TransportProtocol: header.UDPProtocolNumber, + HopLimit: synthesizedTTL, + SrcAddr: destinationAddr, + DstAddr: sourceAddr, + }) + udpHdr = header.UDP(inet6Hdr.Payload()) + ipHdr = inet6Hdr + } + udpHdr.Encode(&header.UDPFields{ + SrcPort: destination.Port, + DstPort: w.source.Port(), + Length: uint16(udpLen), + }) + copy(udpHdr.Payload(), buffer.Bytes()) + udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum( + header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), uint16(udpLen)), + ))) + if inet4Hdr, isInet4 := ipHdr.(header.IPv4); isInet4 { + inet4Hdr.SetChecksum(^inet4Hdr.CalculateChecksum()) + } + return w.writeback.WriteReturnPackets([][]byte{packet}) +} diff --git a/go.mod b/go.mod index 8bab948..fbba8e2 100644 --- a/go.mod +++ b/go.mod @@ -1,32 +1,32 @@ module github.com/sagernet/sing-tun -go 1.24.7 +go 1.25.0 require ( - github.com/florianl/go-nfqueue/v2 v2.0.2 + github.com/florianl/go-nfqueue/v2 v2.1.0 github.com/go-ole/go-ole v1.3.0 github.com/google/btree v1.1.3 - github.com/mdlayher/netlink v1.9.0 - github.com/sagernet/fswatch v0.1.1 - github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 + github.com/mdlayher/netlink v1.11.2 + github.com/sagernet/fswatch v0.1.2 + github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1 github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a - github.com/sagernet/nftables v0.3.0-mod.2 - github.com/sagernet/sing v0.8.0 + github.com/sagernet/nftables v0.3.0-mod.4 + github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8 github.com/stretchr/testify v1.11.1 go4.org/netipx v0.0.0-20231129151722-fdeea329fbba - golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 - golang.org/x/net v0.50.0 - golang.org/x/sys v0.41.0 + golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc + golang.org/x/net v0.57.0 + golang.org/x/sys v0.47.0 ) require ( github.com/davecgh/go-spew v1.1.1 // indirect - github.com/fsnotify/fsnotify v1.7.0 // indirect + github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/google/go-cmp v0.7.0 // indirect - github.com/mdlayher/socket v0.5.1 // indirect + github.com/mdlayher/socket v0.6.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/vishvananda/netns v0.0.4 // indirect - golang.org/x/sync v0.7.0 // indirect - golang.org/x/time v0.7.0 // indirect + golang.org/x/sync v0.20.0 // indirect + golang.org/x/time v0.15.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index 54ed5fd..0c72db9 100644 --- a/go.sum +++ b/go.sum @@ -1,48 +1,50 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/florianl/go-nfqueue/v2 v2.0.2 h1:FL5lQTeetgpCvac1TRwSfgaXUn0YSO7WzGvWNIp3JPE= -github.com/florianl/go-nfqueue/v2 v2.0.2/go.mod h1:VA09+iPOT43OMoCKNfXHyzujQUty2xmzyCRkBOlmabc= -github.com/fsnotify/fsnotify v1.7.0 h1:8JEhPFa5W2WU7YfeZzPNqzMP6Lwt7L2715Ggo0nosvA= -github.com/fsnotify/fsnotify v1.7.0/go.mod h1:40Bi/Hjc2AVfZrqy+aj+yEI+/bRxZnMJyTJwOpGvigM= +github.com/florianl/go-nfqueue/v2 v2.1.0 h1:Fywt30TY/evxyDySpXjxQ1jsRW7nQbLpOhELqpr4068= +github.com/florianl/go-nfqueue/v2 v2.1.0/go.mod h1:8PKUM5rYoVFO5IZV1bifx4/b0jHAglKkHXr9PRwzi4Y= +github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= +github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg= github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= -github.com/mdlayher/netlink v1.9.0 h1:G8+GLq2x3v4D4MVIqDdNUhTUC7TKiCy/6MDkmItfKco= -github.com/mdlayher/netlink v1.9.0/go.mod h1:YBnl5BXsCoRuwBjKKlZ+aYmEoq0r12FDA/3JC+94KDg= -github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos= -github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ= +github.com/jsimonetti/rtnetlink/v2 v2.2.0 h1:/KfZ310gOAFrXXol5VwnFEt+ucldD/0dsSRZwpHCP9w= +github.com/jsimonetti/rtnetlink/v2 v2.2.0/go.mod h1:lbjDHxC+5RJ08lzPeA90Ls2pEoId3F08MoEMlhfHxeI= +github.com/mdlayher/netlink v1.11.2 h1:HKh2jqe+omdSWcQ88nrT7INE61B0NXfiSPFdgL4YbNI= +github.com/mdlayher/netlink v1.11.2/go.mod h1:uT2Yc/QLaZubzDpZIBi9d4GoeLwtp3x1AMeqSRrK2sA= +github.com/mdlayher/socket v0.6.0 h1:ScZPaAGyO1icQnbFrhPM8mnXyMu9qukC1K4ZoM2IQKU= +github.com/mdlayher/socket v0.6.0/go.mod h1:q7vozUAnxSqnjHc12Fik5yUKIzfZ8ITCfMkhOtE9z18= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/sagernet/fswatch v0.1.1 h1:YqID+93B7VRfqIH3PArW/XpJv5H4OLEVWDfProGoRQs= -github.com/sagernet/fswatch v0.1.1/go.mod h1:nz85laH0mkQqJfaOrqPpkwtU1znMFNVTpT/5oRsVz/o= -github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 h1:AzCE2RhBjLJ4WIWc/GejpNh+z30d5H1hwaB0nD9eY3o= -github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1/go.mod h1:NJKBtm9nVEK3iyOYWsUlrDQuoGh4zJ4KOPhSYVidvQ4= +github.com/sagernet/fswatch v0.1.2 h1:/TT7k4mkce1qFPxamLO842WjqBgbTBiXP2mlUjp9PFk= +github.com/sagernet/fswatch v0.1.2/go.mod h1:5BpGmpUQVd3Mc5r313HRpvADHRg3/rKn5QbwFteB880= +github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1 h1:IdQ7yTKkB2wv8txwshxUroPlO4npOYAV71xb7xQ7Lys= +github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1/go.mod h1:9O3SQskYuCfdHNvHEsWuEAgoyKEF74PiWp4NsNUia8g= github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZNjr6sGeT00J8uU7JF4cNUdb44/Duis= github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a/go.mod h1:xLnfdiJbSp8rNqYEdIW/6eDO4mVoogml14Bh2hSiFpM= -github.com/sagernet/nftables v0.3.0-mod.2 h1:ck2KMU02OxL1eDFgGaWYglMDpoOZ7OHzxje+vW5Q0OQ= -github.com/sagernet/nftables v0.3.0-mod.2/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ= -github.com/sagernet/sing v0.8.0 h1:OwLEwbcYfZHvu4olZVljxxC1XRicBqJ1HfiFr6F2WEE= -github.com/sagernet/sing v0.8.0/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak= +github.com/sagernet/nftables v0.3.0-mod.4 h1:vnOtcDYeSXv2e5RoRuGH0lrpttQFJ8iC4ICS2nhlDSo= +github.com/sagernet/nftables v0.3.0-mod.4/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ= +github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8 h1:dyRIj+MZ2rc9JVzJoG04jxu+MpvHrLIZLJr0QjNAMGg= +github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8= github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM= go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M= go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y= -golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 h1:yixxcjnhBmY0nkL253HFVIm0JsFHwrHdT3Yh6szTnfY= -golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8/go.mod h1:jj3sYF3dwk5D+ghuXyeI3r5MFf+NT2An6/9dOA95KSI= -golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60= -golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM= -golang.org/x/sync v0.7.0 h1:YsImfSBoP9QPYL0xyKJPq0gcaJdG3rInoqxTWbfQu9M= -golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= +golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc h1:TS73t7x3KarrNd5qAipmspBDS1rkMcgVG/fS1aRb4Rc= +golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc/go.mod h1:A+z0yzpGtvnG90cToK5n2tu8UJVP2XUATh+r+sfOOOc= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= -golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ= -golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/gtcpip/header/ipv4.go b/gtcpip/header/ipv4.go index d5ffbf1..3041eae 100644 --- a/gtcpip/header/ipv4.go +++ b/gtcpip/header/ipv4.go @@ -22,7 +22,6 @@ import ( "github.com/sagernet/sing-tun/gtcpip" "github.com/sagernet/sing-tun/gtcpip/checksum" - "github.com/sagernet/sing/common" ) // RFC 971 defines the fields of the IPv4 header on page 11 using the following @@ -335,7 +334,7 @@ func (b IPv4) FragmentOffset() uint16 { } func (b IPv4) FragmentOffsetDarwinRaw() uint16 { - return common.NativeEndian.Uint16(b[flagsFO:]) << 3 + return binary.NativeEndian.Uint16(b[flagsFO:]) << 3 } // TotalLength returns the "total length" field of the IPv4 header. @@ -344,7 +343,7 @@ func (b IPv4) TotalLength() uint16 { } func (b IPv4) TotalLengthDarwinRaw() uint16 { - return common.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength()) + return binary.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength()) } // Checksum returns the checksum field of the IPv4 header. @@ -441,7 +440,7 @@ func (b IPv4) SetTotalLength(totalLength uint16) { } func (b IPv4) SetTotalLengthDarwinRaw(totalLength uint16) { - common.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength) + binary.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength) } // SetChecksum sets the checksum field of the IPv4 header. @@ -458,7 +457,7 @@ func (b IPv4) SetFlagsFragmentOffset(flags uint8, offset uint16) { func (b IPv4) SetFlagsFragmentOffsetDarwinRaw(flags uint8, offset uint16) { v := (uint16(flags) << 13) | (offset >> 3) - common.NativeEndian.PutUint16(b[flagsFO:], v) + binary.NativeEndian.PutUint16(b[flagsFO:], v) } // SetID sets the identification field. @@ -1179,7 +1178,7 @@ func (s IPv4OptionsSerializer) Serialize(b []byte) uint8 { // header ends on a 32 bit boundary. The padding is zero. padded := padIPv4OptionsLength(total) b = b[:padded-total] - common.ClearArray(b) + clear(b) return padded } diff --git a/gtcpip/header/ipv6_extension_headers.go b/gtcpip/header/ipv6_extension_headers.go index 6c48b1b..1ab7c9d 100644 --- a/gtcpip/header/ipv6_extension_headers.go +++ b/gtcpip/header/ipv6_extension_headers.go @@ -21,7 +21,6 @@ import ( "math" "github.com/sagernet/sing-tun/gtcpip" - "github.com/sagernet/sing/common" ) // IPv6ExtensionHeaderIdentifier is an IPv6 extension header identifier. @@ -129,7 +128,7 @@ func padIPv6Option(b []byte) { b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6Pad1ExtHdrOptionIdentifier) default: // Pad with PadN. s := b[ipv6ExtHdrOptionPayloadOffset:] - common.ClearArray(s) + clear(s) b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6PadNExtHdrOptionIdentifier) b[ipv6ExtHdrOptionLengthOffset] = uint8(len(s)) } diff --git a/gtcpip/header/ndp_options.go b/gtcpip/header/ndp_options.go index c545120..ca1c6cb 100644 --- a/gtcpip/header/ndp_options.go +++ b/gtcpip/header/ndp_options.go @@ -24,7 +24,6 @@ import ( "time" "github.com/sagernet/sing-tun/gtcpip" - "github.com/sagernet/sing/common" ) // ndpOptionIdentifier is an NDP option type identifier. @@ -341,7 +340,7 @@ func (b NDPOptions) Serialize(s NDPOptionsSerializer) int { // Zero out remaining (padding) bytes, if any exists. if used+2 < l { - common.ClearArray(b[used+2 : l]) + clear(b[used+2 : l]) } b = b[l:] @@ -567,7 +566,7 @@ func (o NDPPrefixInformation) serializeInto(b []byte) int { // Zero out the Reserved2 field. reserved2 := b[ndpPrefixInformationReserved2Offset:][:ndpPrefixInformationReserved2Length] - common.ClearArray(reserved2) + clear(reserved2) return used } @@ -686,7 +685,7 @@ func (o NDPRecursiveDNSServer) serializeInto(b []byte) int { used := copy(b, o) // Zero out the reserved bytes that are before the Lifetime field. - common.ClearArray(b[0:ndpRecursiveDNSServerLifetimeOffset]) + clear(b[0:ndpRecursiveDNSServerLifetimeOffset]) return used } @@ -779,7 +778,7 @@ func (o NDPDNSSearchList) serializeInto(b []byte) int { used := copy(b, o) // Zero out the reserved bytes that are before the Lifetime field. - common.ClearArray(b[0:ndpDNSSearchListLifetimeOffset]) + clear(b[0:ndpDNSSearchListLifetimeOffset]) return used } diff --git a/internal/fdbased_darwin/endpoint.go b/internal/fdbased_darwin/endpoint.go index 05371e7..b3292a2 100644 --- a/internal/fdbased_darwin/endpoint.go +++ b/internal/fdbased_darwin/endpoint.go @@ -50,7 +50,6 @@ import ( "github.com/sagernet/gvisor/pkg/tcpip/header" "github.com/sagernet/gvisor/pkg/tcpip/stack" rawfile "github.com/sagernet/sing-tun/internal/rawfile_darwin" - "github.com/sagernet/sing/common" "golang.org/x/sys/unix" ) @@ -200,10 +199,6 @@ type Options struct { // include CapabilitySaveRestore SaveRestore bool - // DisconnectOk if true, indicates that this NIC capability set should - // include CapabilityDisconnectOk. - DisconnectOk bool - // PacketDispatchMode specifies the type of inbound dispatcher to be // used for this endpoint. PacketDispatchMode PacketDispatchMode @@ -257,10 +252,6 @@ func New(opts *Options) (stack.LinkEndpoint, error) { caps |= stack.CapabilitySaveRestore } - if opts.DisconnectOk { - caps |= stack.CapabilityDisconnectOk - } - if len(opts.FDs) == 0 { return nil, fmt.Errorf("opts.FD is empty, at least one FD must be specified") } @@ -301,7 +292,7 @@ func New(opts *Options) (stack.LinkEndpoint, error) { e.fds = append(e.fds, fdInfo{fd: fd, isSocket: true}) if opts.ProcessorsPerChannel == 0 { - opts.ProcessorsPerChannel = common.Max(1, runtime.GOMAXPROCS(0)/len(opts.FDs)) + opts.ProcessorsPerChannel = max(1, runtime.GOMAXPROCS(0)/len(opts.FDs)) } inboundDispatcher, err := newRecvMMsgDispatcher(fd, e, opts) diff --git a/internal/fdbased_darwin/processors.go b/internal/fdbased_darwin/processors.go index 1ceca83..eb4cb2c 100644 --- a/internal/fdbased_darwin/processors.go +++ b/internal/fdbased_darwin/processors.go @@ -214,34 +214,47 @@ func tcpipConnectionID(pkt *stack.PacketBuffer) (connectionID, bool) { return cid, true } ipHdr := header.IPv6(h) + cid.srcAddr = ipHdr.SourceAddressSlice() + cid.dstAddr = ipHdr.DestinationAddressSlice() + cid.proto = header.IPv6ProtocolNumber - var tcpHdr header.TCP - if tcpip.TransportProtocolNumber(ipHdr.NextHeader()) == header.TCPProtocolNumber { - tcpHdr = header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen]) + if !header.IsExtensionHeader(ipHdr.NextHeader()) { + // Known transport protocols(not just TCP) store the src and dst ports + // in the first 4 bytes after the IPv6 fixed header. + tcpHdr := header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen]) + cid.srcPort = tcpHdr.SourcePort() + cid.dstPort = tcpHdr.DestinationPort() } else { // Slow path for IPv6 extension headers :(. dataBuf := pkt.Data().ToBuffer() dataBuf.TrimFront(header.IPv6MinimumSize) it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(ipHdr.NextHeader()), dataBuf) defer it.Release() + // All fragment packets need to be processed by the same goroutine, so + // only record the ports if this is not a fragment packet. + var isFragment bool for { hdr, done, err := it.Next() if done || err != nil { break } + if fh, ok := hdr.(header.IPv6FragmentExtHdr); ok && !fh.IsAtomic() { + isFragment = true + } hdr.Release() } - h, ok = pkt.Data().PullUp(int(it.HeaderOffset()) + tcpSrcDstPortLen) - if !ok { - return cid, true + 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() } - 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 new file mode 100644 index 0000000..e2b5706 --- /dev/null +++ b/netns_linux.go @@ -0,0 +1,73 @@ +package tun + +import ( + "context" + "net" + "runtime" + "strings" + + "github.com/sagernet/sing/common/control" + E "github.com/sagernet/sing/common/exceptions" + + "golang.org/x/sys/unix" +) + +func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) { + return execInNetworkNamespace(nameOrPath, func() (net.Listener, error) { + return config.Listen(ctx, network, address) + }) +} + +type networkNamespaceInterfaceFinder struct { + control.InterfaceFinder + options *Options +} + +func (f *networkNamespaceInterfaceFinder) Update() error { + return runInNetworkNamespace(f.options.NetNs, f.InterfaceFinder.Update) +} + +func execInNetworkNamespace[T any](nameOrPath string, block func() (T, error)) (T, error) { + if nameOrPath == "" { + return block() + } + type blockResult struct { + value T + err error + } + resultChannel := make(chan blockResult, 1) + go func() { + runtime.LockOSThread() + value, err := execInNetworkNamespaceThread(nameOrPath, block) + resultChannel <- blockResult{value, err} + }() + result := <-resultChannel + return result.value, result.err +} + +func execInNetworkNamespaceThread[T any](nameOrPath string, block func() (T, error)) (T, error) { + var defaultValue T + var path string + if strings.HasPrefix(nameOrPath, "/") { + path = nameOrPath + } else { + path = "/run/netns/" + nameOrPath + } + targetFd, err := unix.Open(path, unix.O_RDONLY|unix.O_CLOEXEC, 0) + if err != nil { + return defaultValue, E.Cause(err, "open netns ", nameOrPath) + } + defer unix.Close(targetFd) + err = unix.Setns(targetFd, unix.CLONE_NEWNET) + if err != nil { + return defaultValue, E.Cause(err, "set netns to ", nameOrPath) + } + return block() +} + +func runInNetworkNamespace(nameOrPath string, block func() error) error { + _, err := execInNetworkNamespace(nameOrPath, func() (struct{}, error) { + return struct{}{}, block() + }) + return err +} diff --git a/netns_other.go b/netns_other.go new file mode 100644 index 0000000..ab5a3a2 --- /dev/null +++ b/netns_other.go @@ -0,0 +1,12 @@ +//go:build !linux + +package tun + +import ( + "context" + "net" +) + +func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) { + return config.Listen(ctx, network, address) +} diff --git a/ping/cmsg_windows.go b/ping/cmsg_windows.go index 07c322c..be5be9b 100644 --- a/ping/cmsg_windows.go +++ b/ping/cmsg_windows.go @@ -1,11 +1,10 @@ package ping import ( + "encoding/binary" "fmt" "unsafe" - "github.com/sagernet/sing/common" - "golang.org/x/net/ipv6" "golang.org/x/sys/windows" ) @@ -37,9 +36,9 @@ func parseIPv6ControlMessage(cmsg []byte) (*ipv6.ControlMessage, error) { } switch cmsghdr.Type { case IPV6_TCLASS: - controlMessage.TrafficClass = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) + controlMessage.TrafficClass = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) case IPV6_HOPLIMIT: - controlMessage.HopLimit = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) + controlMessage.HopLimit = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) } cmsg = cmsg[msgSize:] } diff --git a/ping/socket_linux_unprivileged.go b/ping/socket_linux_unprivileged.go index f709684..1ad1548 100644 --- a/ping/socket_linux_unprivileged.go +++ b/ping/socket_linux_unprivileged.go @@ -9,7 +9,6 @@ import ( "time" "github.com/sagernet/sing-tun/gtcpip/header" - "github.com/sagernet/sing/common" "github.com/sagernet/sing/common/buf" "github.com/sagernet/sing/common/control" M "github.com/sagernet/sing/common/metadata" @@ -175,7 +174,7 @@ func (c *UnprivilegedConn) Close() error { for _, conn := range c.mapping { _ = conn.Close() } - common.ClearMap(c.mapping) + clear(c.mapping) return nil } diff --git a/redirect_linux.go b/redirect_linux.go index 04a1fee..f08d0f7 100644 --- a/redirect_linux.go +++ b/redirect_linux.go @@ -26,6 +26,7 @@ type autoRedirect struct { logger logger.Logger tableName string networkMonitor NetworkUpdateMonitor + ownedNetworkMonitor bool networkListener *list.Element[NetworkUpdateCallback] interfaceFinder control.InterfaceFinder localAddresses []netip.Prefix @@ -51,7 +52,7 @@ type autoRedirect struct { } func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) { - return &autoRedirect{ + r := &autoRedirect{ tunOptions: options.TunOptions, ctx: options.Context, handler: options.Handler, @@ -63,7 +64,11 @@ func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) { customRedirectPortFunc: options.CustomRedirectPort, routeAddressSet: options.RouteAddressSet, routeExcludeAddressSet: options.RouteExcludeAddressSet, - }, nil + } + if options.TunOptions.NetNs != "" { + r.interfaceFinder = &networkNamespaceInterfaceFinder{control.NewDefaultInterfaceFinder(), options.TunOptions} + } + return r, nil } func (r *autoRedirect) Start() error { @@ -89,8 +94,11 @@ func (r *autoRedirect) Start() error { } } } else { + if r.tunOptions.NetNs != "" && !r.useNFTables { + return E.New("auto_redirect in network namespace requires nftables") + } if r.useNFTables { - err = r.initializeNFTables() + err = runInNetworkNamespace(r.tunOptions.NetNs, r.initializeNFTables) if err != nil { return E.Cause(err, "missing nftables support") } @@ -132,7 +140,7 @@ func (r *autoRedirect) Start() error { listenAddr = netip.IPv4Unspecified() } server := newRedirectServer(r.ctx, r.handler, r.logger, listenAddr) - err = server.Start() + err = runInNetworkNamespace(r.tunOptions.NetNs, server.Start) if err != nil { return E.Cause(err, "start redirect server") } @@ -151,24 +159,43 @@ func (r *autoRedirect) Start() error { }) if err != nil { r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err) - } else if err = handler.Start(); err != nil { + } else if err = runInNetworkNamespace(r.tunOptions.NetNs, handler.Start); err != nil { r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err) } else { r.nfqueueHandler = handler r.nfqueueEnabled = true } } - r.cleanupNFTables() - err = r.setupNFTables() - if err != nil { - return E.Cause(err, "setup nftables") - } - if r.tunOptions.AutoRedirectMarkMode { - err = r.setupRedirectRoutes() + if r.tunOptions.NetNs != "" { + var monitor NetworkUpdateMonitor + monitor, err = NewNetworkUpdateMonitor(r.logger) if err != nil { - r.cleanupNFTables() - return E.Cause(err, "setup redirect routes") + 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 + }) + if err != nil { + return err } } else { r.cleanupIPTables() @@ -185,8 +212,14 @@ func (r *autoRedirect) Close() error { r.nfqueueHandler.Close() } if r.useNFTables { - r.cleanupNFTables() - r.cleanupRedirectRoutes() + _ = runInNetworkNamespace(r.tunOptions.NetNs, func() error { + r.cleanupNFTables() + r.cleanupRedirectRoutes() + return nil + }) + if r.ownedNetworkMonitor { + _ = r.networkMonitor.Close() + } } else { r.cleanupIPTables() } @@ -197,7 +230,7 @@ func (r *autoRedirect) Close() error { func (r *autoRedirect) UpdateRouteAddressSet() { if r.useNFTables { - err := r.nftablesUpdateRouteAddressSet() + err := runInNetworkNamespace(r.tunOptions.NetNs, r.nftablesUpdateRouteAddressSet) if err != nil { r.logger.Error("update route address set: ", err) } diff --git a/redirect_nftables.go b/redirect_nftables.go index 5944e4e..c71a770 100644 --- a/redirect_nftables.go +++ b/redirect_nftables.go @@ -299,27 +299,38 @@ func (r *autoRedirect) setupNFTables() error { if err != nil { return E.Cause(err, "flush nftables") } - r.startDockerFirewallMonitor() - err = r.configureDockerFirewall(false) - if err != nil && r.logger != nil { - r.logger.Warn("configure docker firewall: ", err) + if r.tunOptions.NetNs == "" { + r.startDockerFirewallMonitor() + err = r.configureDockerFirewall(false) + if err != nil && r.logger != nil { + r.logger.Warn("configure docker firewall: ", err) + } } r.networkListener = r.networkMonitor.RegisterCallback(func() { - err = r.nftablesUpdateLocalAddressSet() - if err != nil { - r.logger.Error("update local address set: ", err) - } - if r.tunOptions.AutoRedirectMarkMode { - err = r.updateRedirectRoutes() - if err != nil { - r.logger.Error("update redirect routes: ", err) - } + updateErr := runInNetworkNamespace(r.tunOptions.NetNs, r.updateNetworkAddresses) + if updateErr != nil { + r.logger.Error(updateErr) } }) return nil } +func (r *autoRedirect) updateNetworkAddresses() error { + err := r.nftablesUpdateLocalAddressSet() + if err != nil { + err = E.Cause(err, "update local address set") + } + if r.tunOptions.AutoRedirectMarkMode { + routeErr := r.updateRedirectRoutes() + if routeErr != nil { + routeErr = E.Cause(routeErr, "update redirect routes") + } + err = E.Errors(err, routeErr) + } + return err +} + // TODO: test if this works func (r *autoRedirect) nftablesUpdateLocalAddressSet() error { err := r.interfaceFinder.Update() @@ -376,6 +387,7 @@ func (r *autoRedirect) nftablesUpdateRouteAddressSet() error { func (r *autoRedirect) cleanupNFTables() { if r.networkListener != nil { r.networkMonitor.UnregisterCallback(r.networkListener) + r.networkListener = nil } r.stopDockerFirewallMonitor() nft, err := nftables.New() @@ -389,9 +401,11 @@ func (r *autoRedirect) cleanupNFTables() { _ = r.configureOpenWRTFirewall4(nft, true) _ = nft.Flush() _ = nft.CloseLasting() - err = r.configureDockerFirewall(true) - if err != nil && r.logger != nil { - r.logger.Warn("cleanup docker firewall: ", err) + if r.tunOptions.NetNs == "" { + 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 eaf2405..613e45d 100644 --- a/stack.go +++ b/stack.go @@ -14,6 +14,7 @@ import ( type Stack interface { Start() error + ResetNetwork() Close() error } @@ -23,6 +24,9 @@ type StackOptions struct { TunOptions Options UDPTimeout time.Duration ICMPTimeout time.Duration + UDPMapping NATMapping + UDPFiltering NATFiltering + UDPNATMax uint32 Handler Handler Logger logger.Logger ForwarderBindInterface bool diff --git a/stack_gvisor.go b/stack_gvisor.go index 03b2873..c226d05 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,6 +44,7 @@ type GVisor struct { endpoint stack.LinkEndpoint dispatcher *ForwardDispatcher icmpForwarder *ICMPForwarder + udpForwarder *UDPForwarder } type GVisorTun interface { @@ -78,11 +79,18 @@ func NewGVisor( inet6Address: inet6Address, inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress, inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress, - udpTimeout: options.UDPTimeout, icmpTimeout: options.ICMPTimeout, - broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address), - handler: options.Handler, - logger: options.Logger, + 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, } return gStack, nil } @@ -93,7 +101,7 @@ func (t *GVisor) Start() error { return err } if t.handler != nil { - t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout) + t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpNATOptions.Timeout, t.icmpTimeout) } linkEndpoint = &LinkEndpointFilter{ LinkEndpoint: linkEndpoint, @@ -110,7 +118,13 @@ func (t *GVisor) Start() error { return err } ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket) - ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket) + udpForwarder := NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpNATOptions) + err = udpForwarder.Start() + if err != nil { + return err + } + ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, udpForwarder.HandlePacket) + t.udpForwarder = udpForwarder icmpForwarder := NewICMPForwarder(ipStack, t.handler, t.logger) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket) @@ -120,11 +134,24 @@ func (t *GVisor) Start() error { return nil } +func (t *GVisor) ResetNetwork() { + if t.udpForwarder != nil { + t.udpForwarder.udpNat.Purge() + } + if t.icmpForwarder != nil { + t.icmpForwarder.Purge() + } + t.dispatcher.ResetNetwork() +} + func (t *GVisor) Close() error { t.dispatcher.Close() if t.icmpForwarder != nil { t.icmpForwarder.Close() } + if t.udpForwarder != nil { + t.udpForwarder.Close() + } if t.stack == nil { return nil } diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 11e82af..70e27ec 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -72,6 +72,15 @@ func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) return forwarder } +func (f *ICMPForwarder) Purge() { + f.flowAccess.Lock() + for key, flow := range f.flows { + flow.close(FlowCloseReset) + delete(f.flows, key) + } + f.flowAccess.Unlock() +} + func (f *ICMPForwarder) Close() error { f.returnPath.closed.Store(true) f.flowAccess.Lock() @@ -146,9 +155,15 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa } else { ipHdr := header.IPv6(pkt.NetworkHeader().Slice()) icmpHdr := header.ICMPv6(pkt.TransportHeader().Slice()) - if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 { + if icmpHdr.Type() != header.ICMPv6EchoRequest { return false } + if icmpHdr.Code() != 0 { + // The IPv6 built-in echo reply path lacks the LocalAddressTemporary + // check its IPv4 sibling has, so returning false would make the stack + // reply on behalf of arbitrary forwarded destinations. + return true + } identifier := icmpHdr.Ident() key := icmpFlowKey{ v6: true, diff --git a/stack_gvisor_udp.go b/stack_gvisor_udp.go index 2ae54cf..5cd0c93 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -8,7 +8,6 @@ import ( "net/netip" "os" "sync" - "time" _ "unsafe" "github.com/sagernet/gvisor/pkg/buffer" @@ -21,26 +20,35 @@ import ( E "github.com/sagernet/sing/common/exceptions" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" - "github.com/sagernet/sing/common/udpnat2" ) type UDPForwarder struct { ctx context.Context stack *stack.Stack handler Handler - udpNat *udpnat.Service + udpNat *UDPNat } -func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, timeout time.Duration) *UDPForwarder { +func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, options UDPNatOptions) *UDPForwarder { forwarder := &UDPForwarder{ ctx: ctx, stack: stack, handler: handler, } - forwarder.udpNat = udpnat.New(handler, forwarder.PreparePacketConnection, timeout, false) + options.Handler = handler + options.Prepare = forwarder.PreparePacketConnection + forwarder.udpNat = NewUDPNat(options) return forwarder } +func (f *UDPForwarder) Start() error { + return f.udpNat.Start() +} + +func (f *UDPForwarder) Close() error { + return f.udpNat.Close() +} + func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool { source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort) destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort) @@ -63,18 +71,26 @@ func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M firstPacket = append(firstPacket[:len(firstPacket):len(firstPacket)], view.AsSlice()...) } }) + var sourceNetwork tcpip.NetworkProtocolNumber + if source.Addr.Is4() { + sourceNetwork = header.IPv4ProtocolNumber + } else { + sourceNetwork = header.IPv6ProtocolNumber + } switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort(), firstPacket).Action { case ActionReject: gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer)) return false, nil, nil, nil case ActionDrop: return false, nil, nil, nil - } - var sourceNetwork tcpip.NetworkProtocolNumber - if source.Addr.Is4() { - sourceNetwork = header.IPv4ProtocolNumber - } else { - sourceNetwork = header.IPv6ProtocolNumber + case ActionHijackDNS: + f.handler.NewDNSPacket(firstPacket, source, destination, &UDPBackWriter{ + stack: f.stack, + source: AddressFromAddr(source.Addr), + sourcePort: source.Port, + sourceNetwork: sourceNetwork, + }) + return false, nil, nil, nil } writer := &UDPBackWriter{ stack: f.stack, diff --git a/stack_mixed.go b/stack_mixed.go index 4680380..69c8b27 100644 --- a/stack_mixed.go +++ b/stack_mixed.go @@ -19,9 +19,10 @@ import ( type Mixed struct { *System - tun GVisorTun - stack *stack.Stack - endpoint *channel.Endpoint + tun GVisorTun + stack *stack.Stack + endpoint *channel.Endpoint + udpForwarder *UDPForwarder } func NewMixed( @@ -47,7 +48,13 @@ func (m *Mixed) Start() error { if err != nil { return err } - ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpTimeout).HandlePacket) + udpForwarder := NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpNATOptions) + err = udpForwarder.Start() + if err != nil { + return err + } + ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, udpForwarder.HandlePacket) + m.udpForwarder = udpForwarder m.stack = ipStack m.endpoint = endpoint go m.tunLoop() @@ -55,10 +62,20 @@ func (m *Mixed) Start() error { return nil } +func (m *Mixed) ResetNetwork() { + m.System.ResetNetwork() + if m.udpForwarder != nil { + m.udpForwarder.udpNat.Purge() + } +} + func (m *Mixed) Close() error { if m.stack == nil { return nil } + if m.udpForwarder != nil { + m.udpForwarder.Close() + } m.endpoint.Attach(nil) m.stack.Close() for _, endpoint := range m.stack.CleanupEndpoints() { diff --git a/stack_system.go b/stack_system.go index 3cb0cb0..41644bd 100644 --- a/stack_system.go +++ b/stack_system.go @@ -5,7 +5,10 @@ import ( "errors" "net" "net/netip" + "os" "slices" + "sync" + "sync/atomic" "syscall" "time" @@ -19,7 +22,6 @@ import ( "github.com/sagernet/sing/common/logger" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" - "github.com/sagernet/sing/common/udpnat2" ) var ErrIncludeAllNetworks = E.New("`system` and `mixed` stack are not available when `includeAllNetworks` is enabled. See https://github.com/SagerNet/sing-tun/issues/25") @@ -28,6 +30,7 @@ type System struct { ctx context.Context tun Tun tunName string + netNs string mtu int handler Handler logger logger.Logger @@ -44,16 +47,25 @@ type System struct { icmpTimeout time.Duration tcpListener net.Listener tcpListener6 net.Listener - tcpPort uint16 - tcpPort6 uint16 - tcpNat *TCPNat - udpNat *udpnat.Service - dispatcher *ForwardDispatcher - bindInterface bool - interfaceFinder control.InterfaceFinder - frontHeadroom int - txChecksumOffload bool - multiPendingPackets bool + // 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 } type Session struct { @@ -68,6 +80,7 @@ func NewSystem(options StackOptions) (Stack, error) { ctx: options.Context, tun: options.Tun, tunName: options.TunOptions.Name, + netNs: options.TunOptions.NetNs, mtu: int(options.TunOptions.MTU), inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress, inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress, @@ -78,9 +91,17 @@ func NewSystem(options StackOptions) (Stack, error) { inet4Prefixes: options.TunOptions.Inet4Address, inet6Prefixes: options.TunOptions.Inet6Address, broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address), - bindInterface: options.ForwarderBindInterface, - interfaceFinder: options.InterfaceFinder, - multiPendingPackets: options.TunOptions.EXP_MultiPendingPackets, + 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, } if len(options.TunOptions.Inet4Address) > 0 { if !HasNextAddress(options.TunOptions.Inet4Address[0], 1) { @@ -102,8 +123,26 @@ func NewSystem(options StackOptions) (Stack, error) { return stack, nil } +func (s *System) ResetNetwork() { + if s.tcpNat != nil { + s.tcpNat.Purge() + } + if s.udpNat != nil { + s.udpNat.Purge() + } + s.dispatcher.ResetNetwork() +} + func (s *System) Close() error { + // lx/040: mark the deliberate shutdown BEFORE closing the listeners so + // acceptLoop exits quietly instead of treating it as a foreign kill. + s.closing.Store(true) s.dispatcher.Close() + if s.udpNat != nil { + s.udpNat.Close() + } + s.listenAccess.Lock() + defer s.listenAccess.Unlock() return common.Close( s.tcpListener, s.tcpListener6, @@ -119,8 +158,10 @@ func (s *System) Start() error { return nil } -func (s *System) start() error { - _ = fixWindowsFirewall() +// lx/040: TCP forwarder bind, shared by start() and the acceptLoop self-heal +// relisten path. isIPv6 selects the address family; the bind-to-interface +// Control and the EADDRNOTAVAIL retry loop match the original start() code. +func (s *System) listenTCP(isIPv6 bool) (net.Listener, error) { var listener net.ListenConfig if s.bindInterface { listener.Control = control.Append(listener.Control, func(network, address string, conn syscall.RawConn) error { @@ -131,40 +172,60 @@ func (s *System) start() error { return nil }) } + 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() { - for range 3 { - tcpListener, err = listener.Listen(s.ctx, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0")) - if !retryableListenError(err) { - break - } - time.Sleep(time.Second) - } + tcpListener, err = s.listenTCP(false) if err != nil { return err } s.tcpListener = tcpListener - s.tcpPort = M.SocksaddrFromNet(tcpListener.Addr()).Port - go s.acceptLoop(tcpListener) + s.tcpPort.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) + go s.acceptLoop(tcpListener, false) } if s.inet6NextAddress.IsValid() { - for range 3 { - tcpListener, err = listener.Listen(s.ctx, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0")) - if !retryableListenError(err) { - break - } - time.Sleep(time.Second) - } + tcpListener, err = s.listenTCP(true) if err != nil { return err } s.tcpListener6 = tcpListener - s.tcpPort6 = M.SocksaddrFromNet(tcpListener.Addr()).Port - go s.acceptLoop(tcpListener) + s.tcpPort6.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) + go s.acceptLoop(tcpListener, true) } s.tcpNat = NewNat(s.ctx, s.udpTimeout) - s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false) + udpNATOptions := s.udpNATOptions + udpNATOptions.Handler = s.handler + udpNATOptions.Prepare = s.preparePacketConnection + s.udpNat = NewUDPNat(udpNATOptions) + err = s.udpNat.Start() + if err != nil { + return err + } if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN { s.frontHeadroom = linuxTUN.FrontHeadroom() s.txChecksumOffload = linuxTUN.TXChecksumOffload() @@ -336,11 +397,28 @@ func (s *System) processPacket(packet []byte) bool { return writeBack } -func (s *System) acceptLoop(listener net.Listener) { +func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) { for { conn, err := listener.Accept() if err != nil { - return + // lx/040 (SPECS/TASKS/040): upstream silently returns on ANY Accept + // error, leaving the stack alive but every new TCP SYN NAT-rewritten + // onto a dead port (instant RST) until a VPN restart — the LxBox §047 + // "browser dead, QUIC alive" failure. A deliberate System.Close is the + // only quiet exit; anything else means the listener died out from + // under us (e.g. a foreign close on a reused fd number from the + // Java side of the shared Android process) — log it (the errno names + // the killer) and recreate the listener. + if s.closing.Load() { + return + } + newListener, healErr := s.healListener(listener, isIPv6, err) + if healErr != nil { + s.logger.Error("system stack: tcp", ipVersionSuffix(isIPv6), " accept loop died: ", err, "; relisten failed: ", healErr) + return + } + listener = newListener + continue } connPort := M.SocksaddrFromNet(conn.RemoteAddr()).Port session := s.tcpNat.LookupBack(connPort) @@ -352,6 +430,47 @@ func (s *System) acceptLoop(listener net.Listener) { } } +// lx/040: recreate a TCP forwarder listener that died out from under the +// stack. Returns the replacement listener after publishing it (listener field +// + atomic port) under listenAccess, or an error if the stack is closing or +// the bind failed. +func (s *System) healListener(dead net.Listener, isIPv6 bool, cause error) (net.Listener, error) { + port := &s.tcpPort + if isIPv6 { + port = &s.tcpPort6 + } + oldPort := port.Load() + s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener (port ", oldPort, ") accept failed: ", cause, " — recreating listener") + _ = dead.Close() // release netpoll state; harmless if already closed + newListener, err := s.listenTCP(isIPv6) + if err != nil { + return nil, err + } + s.listenAccess.Lock() + defer s.listenAccess.Unlock() + if s.closing.Load() { + _ = newListener.Close() + return nil, net.ErrClosed + } + if isIPv6 { + s.tcpListener6 = newListener + } else { + s.tcpListener = newListener + } + newPort := uint32(M.SocksaddrFromNet(newListener.Addr()).Port) + port.Store(newPort) + recoveries := s.acceptRecoveries.Add(1) + s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener recreated (port ", oldPort, " → ", newPort, ", recoveries: ", recoveries, ")") + return newListener, nil +} + +func ipVersionSuffix(isIPv6 bool) string { + if isIPv6 { + return "6" + } + return "4" +} + func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool { switch ipHdr.TransportProtocol() { case header.TCPProtocolNumber: @@ -361,7 +480,7 @@ func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool { if ipHdr.SourceAddr() == s.inet4Address && ipHdr.FragmentOffset() == 0 && len(ipHdr.Payload()) >= header.TCPMinimumSize && - header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort { + header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort.Load()) { return false } case header.ICMPv4ProtocolNumber: @@ -380,7 +499,7 @@ func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool { } if ipHdr.SourceAddr() == s.inet6Address && len(ipHdr.Payload()) >= header.TCPMinimumSize && - header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 { + header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort6.Load()) { return false } case header.ICMPv6ProtocolNumber: @@ -444,7 +563,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort()) if !destination.Addr().IsGlobalUnicast() { return false, nil - } else if source.Addr() == s.inet4Address && source.Port() == s.tcpPort { + } else if source.Addr() == s.inet4Address && source.Port() == uint16(s.tcpPort.Load()) { session := s.tcpNat.LookupBack(destination.Port()) if session == nil { return false, E.New("ipv4: tcp: session not found: ", destination.Port()) @@ -470,7 +589,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err } rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, s.inet4NextAddress, natPort, true, - s.inet4Address, s.tcpPort, true) + s.inet4Address, uint16(s.tcpPort.Load()), true) } } return true, nil @@ -481,7 +600,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort()) if !destination.Addr().IsGlobalUnicast() { return false, nil - } else if source.Addr() == s.inet6Address && source.Port() == s.tcpPort6 { + } else if source.Addr() == s.inet6Address && source.Port() == uint16(s.tcpPort6.Load()) { session := s.tcpNat.LookupBack(destination.Port()) if session == nil { return false, E.New("ipv6: tcp: session not found: ", destination.Port()) @@ -507,7 +626,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err } rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, s.inet6NextAddress, natPort, true, - s.inet6Address, s.tcpPort6, true) + s.inet6Address, uint16(s.tcpPort6.Load()), true) } } return true, nil @@ -682,20 +801,22 @@ type systemUDPPacketWriter4 struct { txChecksumOffload bool } -func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { - newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len()) - defer newPacket.Release() - newPacket.Resize(w.frontHeadroom, 0) - newPacket.Write(w.header) - newPacket.Write(buffer.Bytes()) - ipHdr := header.IPv4(newPacket.Bytes()) - ipHdr.SetTotalLength(uint16(newPacket.Len())) +func (w *systemUDPPacketWriter4) FrontHeadroom() int { + return w.frontHeadroom + len(w.header) +} + +func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { + payloadLen := buffer.Len() + buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer) + copy(buffer.ExtendHeader(len(w.header)), w.header) + ipHdr := header.IPv4(buffer.Bytes()) + ipHdr.SetTotalLength(uint16(buffer.Len())) ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetSourceAddr(destination.Addr) udpHdr := header.UDP(ipHdr.Payload()) udpHdr.SetDestinationPort(udpHdr.SourcePort()) udpHdr.SetSourcePort(destination.Port) - udpHdr.SetLength(uint16(buffer.Len() + header.UDPMinimumSize)) + udpHdr.SetLength(uint16(payloadLen + header.UDPMinimumSize)) if !w.txChecksumOffload { udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum( header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()), @@ -704,12 +825,61 @@ func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.S udpHdr.SetChecksum(0) } ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) + return buffer +} + +func (w *systemUDPPacketWriter4) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { + buffer = w.preparePacket(buffer, destination) if PacketOffset > 0 { - PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version) - } else { - newPacket.Advance(-w.frontHeadroom) + PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv4Version) + } + if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 { + buffer.Advance(-remainingHeadroom) + } + return buffer +} + +func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + buffer = w.prepareWritePacket(buffer, destination) + defer buffer.Release() + return common.Error(w.tun.Write(buffer.Bytes())) +} + +func (w *systemUDPPacketWriter4) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) { + switch w.tun.(type) { + case LinuxTUN, DarwinTUN: + return w, true + default: + return nil, false + } +} + +func (w *systemUDPPacketWriter4) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error { + if len(buffers) == 0 || len(buffers) != len(destinations) { + buf.ReleaseMulti(buffers) + return os.ErrInvalid + } + defer func() { + buf.ReleaseMulti(buffers) + }() + switch tunInterface := w.tun.(type) { + case LinuxTUN: + packets := make([][]byte, len(buffers)) + for index, buffer := range buffers { + buffer = w.preparePacket(buffer, destinations[index]) + buffer.Advance(-w.frontHeadroom) + buffers[index] = buffer + packets[index] = buffer.Bytes() + } + return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom)) + case DarwinTUN: + for index, buffer := range buffers { + buffers[index] = w.preparePacket(buffer, destinations[index]) + } + return tunInterface.BatchWrite(buffers) + default: + return os.ErrInvalid } - return common.Error(w.tun.Write(newPacket.Bytes())) } type systemUDPPacketWriter6 struct { @@ -720,14 +890,16 @@ type systemUDPPacketWriter6 struct { txChecksumOffload bool } -func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { - newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len()) - defer newPacket.Release() - newPacket.Resize(w.frontHeadroom, 0) - newPacket.Write(w.header) - newPacket.Write(buffer.Bytes()) - ipHdr := header.IPv6(newPacket.Bytes()) - udpLen := uint16(header.UDPMinimumSize + buffer.Len()) +func (w *systemUDPPacketWriter6) FrontHeadroom() int { + return w.frontHeadroom + len(w.header) +} + +func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { + payloadLen := buffer.Len() + buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer) + copy(buffer.ExtendHeader(len(w.header)), w.header) + ipHdr := header.IPv6(buffer.Bytes()) + udpLen := uint16(header.UDPMinimumSize + payloadLen) ipHdr.SetPayloadLength(udpLen) ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetSourceAddr(destination.Addr) @@ -742,12 +914,61 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S } else { udpHdr.SetChecksum(0) } + return buffer +} + +func (w *systemUDPPacketWriter6) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { + buffer = w.preparePacket(buffer, destination) if PacketOffset > 0 { - PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version) - } else { - newPacket.Advance(-w.frontHeadroom) + PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv6Version) + } + if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 { + buffer.Advance(-remainingHeadroom) + } + return buffer +} + +func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + buffer = w.prepareWritePacket(buffer, destination) + defer buffer.Release() + return common.Error(w.tun.Write(buffer.Bytes())) +} + +func (w *systemUDPPacketWriter6) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) { + switch w.tun.(type) { + case LinuxTUN, DarwinTUN: + return w, true + default: + return nil, false + } +} + +func (w *systemUDPPacketWriter6) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error { + if len(buffers) == 0 || len(buffers) != len(destinations) { + buf.ReleaseMulti(buffers) + return os.ErrInvalid + } + defer func() { + buf.ReleaseMulti(buffers) + }() + switch tunInterface := w.tun.(type) { + case LinuxTUN: + packets := make([][]byte, len(buffers)) + for index, buffer := range buffers { + buffer = w.preparePacket(buffer, destinations[index]) + buffer.Advance(-w.frontHeadroom) + buffers[index] = buffer + packets[index] = buffer.Bytes() + } + return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom)) + case DarwinTUN: + for index, buffer := range buffers { + buffers[index] = w.preparePacket(buffer, destinations[index]) + } + return tunInterface.BatchWrite(buffers) + default: + return os.ErrInvalid } - return common.Error(w.tun.Write(newPacket.Bytes())) } func newSystemWriteback(tunInterface Tun, frontHeadroom int) ForwardWriteback { diff --git a/stack_system_nat.go b/stack_system_nat.go index 2fec29c..1dd5377 100644 --- a/stack_system_nat.go +++ b/stack_system_nat.go @@ -86,6 +86,15 @@ func (n *TCPNat) checkTimeout() { n.addrAccess.Unlock() } +func (n *TCPNat) Purge() { + n.addrAccess.Lock() + n.portAccess.Lock() + clear(n.addrMap) + clear(n.portMap) + n.portAccess.Unlock() + n.addrAccess.Unlock() +} + func (n *TCPNat) LookupBack(port uint16) *TCPSession { n.portAccess.RLock() session := n.portMap[port] diff --git a/stack_system_packet.go b/stack_system_packet.go index a8f8076..d00b95d 100644 --- a/stack_system_packet.go +++ b/stack_system_packet.go @@ -5,7 +5,6 @@ import ( "syscall" "github.com/sagernet/sing-tun/gtcpip/header" - "github.com/sagernet/sing/common" ) func PacketIPVersion(packet []byte) int { @@ -14,7 +13,7 @@ func PacketIPVersion(packet []byte) int { func PacketFillHeader(packet []byte, ipVersion int) { if PacketOffset > 0 { - common.ClearArray(packet[:3]) + clear(packet[:3]) switch ipVersion { case header.IPv4Version: packet[3] = syscall.AF_INET diff --git a/stack_system_selfheal_test.go b/stack_system_selfheal_test.go new file mode 100644 index 0000000..ab063c2 --- /dev/null +++ b/stack_system_selfheal_test.go @@ -0,0 +1,117 @@ +package tun + +// lx/040 (SPECS/TASKS/040-SINGTUN_ACCEPTLOOP_SELFHEAL): acceptLoop self-heal. +// +// Red/green против апстрима 2d9b8aed5fe2: там acceptLoop(listener) при любой +// ошибке Accept молча выходит навсегда — восстановления нет, порт не меняется, +// новый connect вечно бьётся в мёртвый сокет. Для red-прогона на чистом +// апстрим-чекауте достаточно адаптировать хелперы ниже (currentTCPPort → +// s.tcpPort, spawnAcceptLoop → go s.acceptLoop(ln)): тест упадёт по таймауту +// ожидания восстановления. + +import ( + "context" + "fmt" + "net" + "net/netip" + "testing" + "time" + + "github.com/sagernet/sing/common/logger" +) + +func newSelfHealTestSystem(t *testing.T) *System { + t.Helper() + s := &System{ + ctx: context.Background(), + logger: logger.NOP(), + inet4Address: netip.MustParseAddr("127.0.0.1"), + udpTimeout: time.Minute, + } + s.tcpNat = NewNat(s.ctx, s.udpTimeout) + ln, err := s.listenTCP(false) + if err != nil { + t.Fatalf("listenTCP: %v", err) + } + s.tcpListener = ln + s.tcpPort.Store(uint32(ln.Addr().(*net.TCPAddr).Port)) + spawnAcceptLoop(s, ln) + return s +} + +func currentTCPPort(s *System) uint32 { + return s.tcpPort.Load() +} + +func spawnAcceptLoop(s *System, ln net.Listener) { + go s.acceptLoop(ln, false) +} + +func dialForwarder(t *testing.T, port uint32) error { + t.Helper() + conn, err := net.DialTimeout("tcp4", fmt.Sprintf("127.0.0.1:%d", port), time.Second) + if err == nil { + _ = conn.Close() + } + return err +} + +// Убийство listener'а мимо System.Close (эмуляция чужого close по +// переиспользованному fd-номеру) должно приводить к пересозданию listener'а +// и продолжению приёма TCP, а не к вечной смерти петли. +func TestSystemAcceptLoopSelfHeal(t *testing.T) { + s := newSelfHealTestSystem(t) + oldPort := currentTCPPort(s) + + if err := dialForwarder(t, oldPort); err != nil { + t.Fatalf("healthy listener refused connect: %v", err) + } + + // Убить listener из-под стека: closing НЕ выставлен. + _ = s.tcpListener.Close() + + deadline := time.Now().Add(5 * time.Second) + healed := false + for time.Now().Before(deadline) { + if s.acceptRecoveries.Load() > 0 { + healed = true + break + } + time.Sleep(10 * time.Millisecond) + } + if !healed { + t.Fatalf("acceptLoop did not recover within 5s (upstream behavior: silent permanent death)") + } + + newPort := currentTCPPort(s) + if newPort == oldPort { + t.Fatalf("recovered port equals dead port %d — relisten did not publish a new port", oldPort) + } + if err := dialForwarder(t, newPort); err != nil { + t.Fatalf("connect to recreated listener (port %d) failed: %v", newPort, err) + } + if got := s.acceptRecoveries.Load(); got != 1 { + t.Fatalf("acceptRecoveries = %d, want 1", got) + } + + s.closing.Store(true) + _ = s.tcpListener.Close() +} + +// Штатное закрытие (closing выставлен, как это делает System.Close) обязано +// оставаться тихим: без пересозданий и без роста счётчика. +func TestSystemAcceptLoopQuietOnClose(t *testing.T) { + s := newSelfHealTestSystem(t) + oldPort := currentTCPPort(s) + + s.closing.Store(true) + _ = s.tcpListener.Close() + + time.Sleep(300 * time.Millisecond) + if got := s.acceptRecoveries.Load(); got != 0 { + t.Fatalf("deliberate close triggered %d recoveries, want 0", got) + } + if port := currentTCPPort(s); port != oldPort { + t.Fatalf("deliberate close changed port %d → %d", oldPort, port) + } +} diff --git a/tun.go b/tun.go index c6518f4..e770122 100644 --- a/tun.go +++ b/tun.go @@ -14,12 +14,14 @@ import ( E "github.com/sagernet/sing/common/exceptions" F "github.com/sagernet/sing/common/format" "github.com/sagernet/sing/common/logger" + M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" "github.com/sagernet/sing/common/ranges" ) type Handler interface { JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort, firstPacket []byte) FlowVerdict + NewDNSPacket(payload []byte, source M.Socksaddr, destination M.Socksaddr, writer N.PacketWriter) N.TCPConnectionHandlerEx N.UDPConnectionHandlerEx } @@ -66,6 +68,7 @@ const ( type Options struct { Name string + NetNs string Inet4Address []netip.Prefix Inet6Address []netip.Prefix MTU uint32 diff --git a/tun_linux.go b/tun_linux.go index 487051d..41d4dc4 100644 --- a/tun_linux.go +++ b/tun_linux.go @@ -51,37 +51,38 @@ type NativeTun struct { } func New(options Options) (Tun, error) { - var nativeTun *NativeTun if options.FileDescriptor == 0 { - 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)) - } - } else { - nativeTun = &NativeTun{ - tunFd: options.FileDescriptor, - tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"), - options: options, - } - if options.GSO { - err := nativeTun.enableGSO() + return execInNetworkNamespace(options.NetNs, func() (Tun, error) { + 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)) + } + return nativeTun, nil + }) + } + 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) } } } @@ -290,10 +291,10 @@ func (t *NativeTun) Name() (string, error) { func (t *NativeTun) Start() error { if t.options.FileDescriptor == 0 { - if !t.options.EXP_ExternalConfiguration { + if !t.options.EXP_ExternalConfiguration && t.options.NetNs == "" { t.options.InterfaceMonitor.RegisterMyInterface(t.options.Name) } - err := t.start() + err := runInNetworkNamespace(t.options.NetNs, t.start) if err != nil { return err } @@ -354,7 +355,7 @@ func (t *NativeTun) start() error { return E.Cause(err, "set rules") } - if t.options.DNSMode != DNSModeDisabled { + if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { err = t.setSearchDomainForSystemdResolved() if err != nil { return E.Cause(err, "set search domain") @@ -374,11 +375,13 @@ func (t *NativeTun) Close() error { if t.options.EXP_ExternalConfiguration { return common.Close(common.PtrOrNil(t.tunFile)) } - if t.options.DNSMode != DNSModeDisabled { + if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { t.unsetSearchDomainForSystemdResolved() } - t.unsetAddresses() - return E.Errors(t.unsetRoute(), t.unsetRules(), common.Close(common.PtrOrNil(t.tunFile))) + return E.Errors(runInNetworkNamespace(t.options.NetNs, func() error { + 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) { @@ -625,16 +628,18 @@ func (t *NativeTun) UpdateRouteOptions(tunOptions Options) error { t.options = tunOptions return nil } - 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) + 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) + }) } func (t *NativeTun) routes(tunLink netlink.Link) ([]netlink.Route, error) { diff --git a/udp_egress.go b/udp_egress.go new file mode 100644 index 0000000..0ace53e --- /dev/null +++ b/udp_egress.go @@ -0,0 +1,287 @@ +package tun + +import ( + "context" + "net" + "net/netip" + "runtime" + "slices" + "sync" + "sync/atomic" + + "github.com/sagernet/sing/common/buf" + "github.com/sagernet/sing/common/control" + E "github.com/sagernet/sing/common/exceptions" + "github.com/sagernet/sing/common/logger" + "github.com/sagernet/sing/common/x/list" +) + +const udpEgressBufferSize = 65535 + +type UDPEgressPoolOptions struct { + Logger logger.Logger + Network string + Control control.Func + InterfaceFinder control.InterfaceFinder + InterfaceMonitor DefaultInterfaceMonitor + ExcludeInterface string + IsExempt func() bool +} + +type UDPEgressPool struct { + logger logger.Logger + network string + control control.Func + interfaceFinder control.InterfaceFinder + interfaceMonitor DefaultInterfaceMonitor + excludeInterface string + isExempt func() bool + access sync.Mutex + port uint16 + anchorInterfaceIndex int + receiveDone chan struct{} + members map[udpEgressSpec]*udpEgressMember + state atomic.Pointer[[]*udpEgressMember] + packetChan chan udpEgressPacket + memberReaders sync.WaitGroup + finderElement *list.Element[control.InterfaceUpdateCallback] +} + +type udpEgressSpec struct { + interfaceIndex int + interfaceName string + prefix netip.Prefix +} + +type udpEgressMember struct { + prefix netip.Prefix + conn *net.UDPConn +} + +type udpEgressPacket struct { + buffer *buf.Buffer + source netip.AddrPort +} + +func NewUDPEgressPool(options UDPEgressPoolOptions) *UDPEgressPool { + return &UDPEgressPool{ + logger: options.Logger, + network: options.Network, + control: options.Control, + interfaceFinder: options.InterfaceFinder, + interfaceMonitor: options.InterfaceMonitor, + excludeInterface: options.ExcludeInterface, + isExempt: options.IsExempt, + anchorInterfaceIndex: -1, + members: make(map[udpEgressSpec]*udpEgressMember), + packetChan: make(chan udpEgressPacket, 128), + } +} + +func (p *UDPEgressPool) Close() { + p.SetEgressPort(0) + p.access.Lock() + defer p.access.Unlock() + if p.finderElement != nil { + p.interfaceFinder.UnregisterInterfaceUpdateCallback(p.finderElement) + p.finderElement = nil + } +} + +func (p *UDPEgressPool) SetEgressPort(port uint16) bool { + p.access.Lock() + defer p.access.Unlock() + if p.port == port { + return p.state.Load() != nil + } + if p.receiveDone != nil { + close(p.receiveDone) + p.receiveDone = nil + } + p.port = 0 + p.state.Store(nil) + for spec, member := range p.members { + delete(p.members, spec) + member.conn.Close() + } + p.memberReaders.Wait() + for { + select { + case packet := <-p.packetChan: + packet.buffer.Release() + default: + goto drained + } + } +drained: + p.anchorInterfaceIndex = -1 + if port == 0 { + return false + } + p.port = port + defaultInterface := p.interfaceMonitor.DefaultInterface() + if defaultInterface != nil { + p.anchorInterfaceIndex = defaultInterface.Index + } + p.receiveDone = make(chan struct{}) + if p.finderElement == nil { + p.finderElement = p.interfaceFinder.RegisterInterfaceUpdateCallback(func(interfaces []control.Interface) { + p.access.Lock() + defer p.access.Unlock() + p.rebuildLocked() + }) + } + p.rebuildLocked() + return p.state.Load() != nil +} + +func (p *UDPEgressPool) LookupEgress(destination netip.AddrPort) *net.UDPConn { + members := p.state.Load() + if members == nil { + return nil + } + address := destination.Addr().Unmap() + for _, member := range *members { + if member.prefix.Contains(address) { + return member.conn + } + } + return nil +} + +func (p *UDPEgressPool) ReceiveEgress(buffer []byte) (int, netip.AddrPort, error) { + p.access.Lock() + receiveDone := p.receiveDone + p.access.Unlock() + if receiveDone == nil { + return 0, netip.AddrPort{}, net.ErrClosed + } + select { + case <-receiveDone: + return 0, netip.AddrPort{}, net.ErrClosed + default: + } + select { + case packet := <-p.packetChan: + copied := copy(buffer, packet.buffer.Bytes()) + packet.buffer.Release() + return copied, packet.source, nil + case <-receiveDone: + return 0, netip.AddrPort{}, net.ErrClosed + } +} + +func (p *UDPEgressPool) rebuildLocked() { + if p.port == 0 { + return + } + specs := make(map[udpEgressSpec]struct{}) + if !p.isExempt() { + for _, networkInterface := range p.interfaceFinder.Interfaces() { + if networkInterface.Flags&net.FlagUp == 0 || + networkInterface.Flags&net.FlagLoopback != 0 || + networkInterface.Flags&net.FlagPointToPoint != 0 || + networkInterface.Flags&net.FlagBroadcast == 0 || + networkInterface.Index == p.anchorInterfaceIndex || + networkInterface.Name == p.excludeInterface { + continue + } + for _, prefix := range networkInterface.Addresses { + if !prefix.Addr().IsGlobalUnicast() { + continue + } + if p.network == "udp4" && !prefix.Addr().Is4() { + continue + } + if p.network == "udp6" && prefix.Addr().Is4() { + continue + } + specs[udpEgressSpec{ + interfaceIndex: networkInterface.Index, + interfaceName: networkInterface.Name, + prefix: prefix, + }] = struct{}{} + } + } + } + for spec, member := range p.members { + _, loaded := specs[spec] + if loaded { + continue + } + delete(p.members, spec) + member.conn.Close() + } + for spec := range specs { + _, loaded := p.members[spec] + if loaded { + continue + } + memberConn, err := p.listenMember(spec) + if err != nil { + p.logger.Warn(E.Cause(err, "listen egress member on ", spec.interfaceName, " (", spec.prefix.Addr(), ")")) + continue + } + member := &udpEgressMember{ + prefix: spec.prefix.Masked(), + conn: memberConn, + } + p.members[spec] = member + p.memberReaders.Add(1) + go p.readMember(member, p.receiveDone) + } + members := make([]*udpEgressMember, 0, len(p.members)) + for _, member := range p.members { + members = append(members, member) + } + slices.SortFunc(members, func(firstMember, secondMember *udpEgressMember) int { + return secondMember.prefix.Bits() - firstMember.prefix.Bits() + }) + if len(members) == 0 { + p.state.Store(nil) + } else { + p.state.Store(&members) + } +} + +func (p *UDPEgressPool) listenMember(spec udpEgressSpec) (*net.UDPConn, error) { + var listenConfig net.ListenConfig + if runtime.GOOS == "darwin" || runtime.GOOS == "ios" { + listenConfig.Control = control.ReuseAddrOnly() + } + listenConfig.Control = control.Append(listenConfig.Control, control.DisableUDPNetReset()) + listenConfig.Control = control.Append(listenConfig.Control, control.BindToInterface(p.interfaceFinder, spec.interfaceName, spec.interfaceIndex)) + listenConfig.Control = control.Append(listenConfig.Control, p.control) + var network string + if spec.prefix.Addr().Is4() { + network = "udp4" + } else { + network = "udp6" + } + packetConn, err := listenConfig.ListenPacket(context.Background(), network, netip.AddrPortFrom(spec.prefix.Addr(), p.port).String()) + if err != nil { + return nil, err + } + return packetConn.(*net.UDPConn), nil +} + +func (p *UDPEgressPool) readMember(member *udpEgressMember, doneChan <-chan struct{}) { + defer p.memberReaders.Done() + for { + buffer := buf.NewSize(udpEgressBufferSize) + dataLength, source, err := member.conn.ReadFromUDPAddrPort(buffer.FreeBytes()) + if err != nil { + buffer.Release() + return + } + buffer.Extend(dataLength) + select { + case p.packetChan <- udpEgressPacket{buffer: buffer, source: source}: + case <-doneChan: + buffer.Release() + return + default: + buffer.Release() + } + } +} diff --git a/udp_egress_conn.go b/udp_egress_conn.go new file mode 100644 index 0000000..93cd5d4 --- /dev/null +++ b/udp_egress_conn.go @@ -0,0 +1,124 @@ +package tun + +import ( + "net" + "net/netip" + "sync" + "time" + + "github.com/sagernet/sing/common/buf" + E "github.com/sagernet/sing/common/exceptions" +) + +type UDPEgressConn struct { + anchor *net.UDPConn + pool *UDPEgressPool + packetChan chan udpEgressConnPacket + doneChan chan struct{} + closeOnce sync.Once + readWait sync.WaitGroup +} + +type udpEgressConnPacket struct { + buffer *buf.Buffer + source netip.AddrPort + err error +} + +func NewUDPEgressConn(anchor *net.UDPConn, pool *UDPEgressPool) *UDPEgressConn { + conn := &UDPEgressConn{ + anchor: anchor, + pool: pool, + packetChan: make(chan udpEgressConnPacket, 64), + doneChan: make(chan struct{}), + } + conn.readWait.Add(2) + go conn.read(anchor.ReadFromUDPAddrPort) + go conn.read(pool.ReceiveEgress) + return conn +} + +func (c *UDPEgressConn) read(readPacket func([]byte) (int, netip.AddrPort, error)) { + defer c.readWait.Done() + for { + buffer := buf.NewSize(udpEgressBufferSize) + dataLength, source, err := readPacket(buffer.FreeBytes()) + if err != nil { + buffer.Release() + if E.IsClosed(err) { + return + } + select { + case c.packetChan <- udpEgressConnPacket{err: err}: + case <-c.doneChan: + return + } + continue + } + buffer.Extend(dataLength) + select { + case c.packetChan <- udpEgressConnPacket{buffer: buffer, source: source}: + case <-c.doneChan: + buffer.Release() + return + } + } +} + +func (c *UDPEgressConn) ReadFromUDPAddrPort(buffer []byte) (int, netip.AddrPort, error) { + select { + case packet := <-c.packetChan: + if packet.err != nil { + return 0, netip.AddrPort{}, packet.err + } + copied := copy(buffer, packet.buffer.Bytes()) + packet.buffer.Release() + return copied, packet.source, nil + case <-c.doneChan: + return 0, netip.AddrPort{}, net.ErrClosed + } +} + +func (c *UDPEgressConn) WriteToUDPAddrPort(buffer []byte, destination netip.AddrPort) (int, error) { + memberConn := c.pool.LookupEgress(destination) + if memberConn != nil { + return memberConn.WriteToUDPAddrPort(buffer, destination) + } + return c.anchor.WriteToUDPAddrPort(buffer, destination) +} + +func (c *UDPEgressConn) LocalAddr() net.Addr { + return c.anchor.LocalAddr() +} + +func (c *UDPEgressConn) SetDeadline(t time.Time) error { + return c.anchor.SetDeadline(t) +} + +func (c *UDPEgressConn) SetReadDeadline(t time.Time) error { + return c.anchor.SetReadDeadline(t) +} + +func (c *UDPEgressConn) SetWriteDeadline(t time.Time) error { + return c.anchor.SetWriteDeadline(t) +} + +func (c *UDPEgressConn) Close() error { + c.closeOnce.Do(func() { + close(c.doneChan) + c.anchor.Close() + c.pool.Close() + c.readWait.Wait() + for { + select { + case packet := <-c.packetChan: + if packet.buffer != nil { + packet.buffer.Release() + } + default: + return + } + } + }) + return nil +} diff --git a/udp_nat.go b/udp_nat.go new file mode 100644 index 0000000..6d5dd94 --- /dev/null +++ b/udp_nat.go @@ -0,0 +1,789 @@ +package tun + +import ( + "context" + "io" + "net" + "net/netip" + "os" + "runtime" + "slices" + "sync" + "sync/atomic" + "time" + + "github.com/sagernet/sing/common" + "github.com/sagernet/sing/common/buf" + "github.com/sagernet/sing/common/canceler" + "github.com/sagernet/sing/common/control" + "github.com/sagernet/sing/common/memory" + M "github.com/sagernet/sing/common/metadata" + N "github.com/sagernet/sing/common/network" + "github.com/sagernet/sing/common/pipe" + "github.com/sagernet/sing/common/x/list" + "github.com/sagernet/sing/contrab/freelru" + "github.com/sagernet/sing/contrab/maphash" +) + +type NATMapping uint8 + +const ( + NATMappingEndpointIndependent NATMapping = iota + NATMappingAddressDependent + NATMappingAddressAndPortDependent +) + +type NATFiltering uint8 + +const ( + NATFilteringEndpointIndependent NATFiltering = iota + NATFilteringAddressDependent + NATFilteringAddressAndPortDependent +) + +type UDPNatPrepareFunc func(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) + +type UDPNatOptions struct { + Handler N.UDPConnectionHandlerEx + Prepare UDPNatPrepareFunc + Timeout time.Duration + Shared bool + Mapping NATMapping + Filtering NATFiltering + MaxSize uint32 + + InterfaceFinder control.InterfaceFinder + ExcludeInterface []string +} + +type udpNatSessionKey struct { + sourceAddr netip.Addr + destinationAddr netip.Addr + sourcePort uint16 + destinationPort uint16 + interfaceIndex uint32 +} + +type udpNatFilterKey struct { + sessionID uint64 + peer netip.AddrPort +} + +type udpNatEgressEntry struct { + prefix netip.Prefix + interfaceIndex uint32 +} + +const udpNatEgressLinearThreshold = 8 + +type udpNatEgressBuckets struct { + inet4 [256][]udpNatEgressEntry + inet6 [256][]udpNatEgressEntry +} + +type udpNatEgressTable struct { + entries []udpNatEgressEntry + buckets *udpNatEgressBuckets +} + +type UDPNat struct { + handler N.UDPConnectionHandlerEx + prepare UDPNatPrepareFunc + timeout time.Duration + mapping NATMapping + filtering NATFiltering + cache *freelru.Cache[udpNatSessionKey, *udpNatConn] + filterCache *freelru.Cache[udpNatFilterKey, *udpNatConn] + nextFilterSessionID atomic.Uint64 + interfaceFinder control.InterfaceFinder + excludeInterface []string + interfaceElement *list.Element[control.InterfaceUpdateCallback] + egress atomic.Pointer[udpNatEgressTable] + classAccess sync.Mutex + classConns map[uint32]map[*udpNatConn]struct{} + cleanup *udpNatCleanupQueue + state atomic.Uint32 + lifecycleAccess sync.Mutex + closeOnce sync.Once + cleanupDone chan struct{} + cleanupWait sync.WaitGroup +} + +func NewUDPNat(options UDPNatOptions) *UDPNat { + if options.Timeout == 0 { + panic("invalid timeout") + } + maxSize := options.MaxSize + if maxSize == 0 { + if runtime.GOOS == "ios" { + maxSize = 4096 + } else if totalMemory := memory.Total(); totalMemory == 0 { + maxSize = 16384 + } else { + maxSize = uint32(min(max(totalMemory/16384, 4096), 16384)) + } + } + hasher := maphash.NewHasher[udpNatSessionKey]() + cache := common.Must1(freelru.New[udpNatSessionKey, *udpNatConn](maxSize, hasher.Hash32, options.Shared)) + var filterCache *freelru.Cache[udpNatFilterKey, *udpNatConn] + if NATMapping(options.Filtering) > options.Mapping { + filterHasher := maphash.NewHasher[udpNatFilterKey]() + filterCache = common.Must1(freelru.New[udpNatFilterKey, *udpNatConn](maxSize, filterHasher.Hash32, options.Shared)) + } + service := &UDPNat{ + handler: options.Handler, + prepare: options.Prepare, + timeout: options.Timeout, + mapping: options.Mapping, + filtering: options.Filtering, + cache: cache, + filterCache: filterCache, + interfaceFinder: options.InterfaceFinder, + excludeInterface: options.ExcludeInterface, + classConns: make(map[uint32]map[*udpNatConn]struct{}), + cleanupDone: make(chan struct{}), + } + service.cleanup = newUDPNatCleanupQueue(service) + cache.SetLifetime(options.Timeout) + cache.SetHealthCheck(func(_ udpNatSessionKey, conn *udpNatConn) bool { + select { + case <-conn.doneChan: + return false + default: + return true + } + }) + cache.SetOnEvict(func(_ udpNatSessionKey, conn *udpNatConn) { + conn.closeFromCache() + }) + if filterCache != nil { + filterCache.SetOnEvict(func(key udpNatFilterKey, conn *udpNatConn) { + conn.removeFilterPeer(key.peer) + }) + } + return service +} + +func (s *UDPNat) Close() error { + s.closeOnce.Do(func() { + s.lifecycleAccess.Lock() + previousState := s.state.Swap(udpNatStateClosed) + if previousState == udpNatStateStarted { + close(s.cleanupDone) + } + s.lifecycleAccess.Unlock() + if previousState == udpNatStateStarted { + s.cleanupWait.Wait() + } + if s.interfaceElement != nil { + s.interfaceFinder.UnregisterInterfaceUpdateCallback(s.interfaceElement) + s.interfaceElement = nil + } + s.cache.Purge() + if s.filterCache != nil { + s.filterCache.Purge() + } + s.cleanup.clear() + }) + return nil +} + +func (s *UDPNat) reloadInterfaces() { + s.updateInterfaces(s.interfaceFinder.Interfaces()) +} + +func (s *UDPNat) updateInterfaces(interfaces []control.Interface) { + var entries []udpNatEgressEntry + for _, networkInterface := range interfaces { + if networkInterface.Flags&net.FlagUp == 0 || + networkInterface.Flags&net.FlagLoopback != 0 || + networkInterface.Flags&net.FlagPointToPoint != 0 || + networkInterface.Flags&net.FlagBroadcast == 0 { + continue + } + if slices.Contains(s.excludeInterface, networkInterface.Name) { + continue + } + for _, prefix := range networkInterface.Addresses { + if !prefix.Addr().IsGlobalUnicast() { + continue + } + entries = append(entries, udpNatEgressEntry{prefix.Masked(), uint32(networkInterface.Index)}) + } + } + s.egress.Store(newUDPNatEgressTable(entries)) + var closeConns []*udpNatConn + s.classAccess.Lock() + for interfaceIndex, conns := range s.classConns { + if !slices.ContainsFunc(entries, func(entry udpNatEgressEntry) bool { + return entry.interfaceIndex == interfaceIndex + }) { + for conn := range conns { + closeConns = append(closeConns, conn) + } + delete(s.classConns, interfaceIndex) + } + } + s.classAccess.Unlock() + for _, conn := range closeConns { + conn.Close() + } +} + +func (s *UDPNat) classify(destination M.Socksaddr) uint32 { + table := s.egress.Load() + if table == nil || !destination.IsIP() { + return 0 + } + return table.lookup(destination.Addr.Unmap()) +} + +func newUDPNatEgressTable(entries []udpNatEgressEntry) *udpNatEgressTable { + entries = slices.Clone(entries) + slices.SortStableFunc(entries, func(a, b udpNatEgressEntry) int { + return b.prefix.Bits() - a.prefix.Bits() + }) + table := &udpNatEgressTable{entries: entries} + if len(entries) <= udpNatEgressLinearThreshold { + return table + } + buckets := new(udpNatEgressBuckets) + for _, entry := range entries { + address := entry.prefix.Addr().Unmap() + bits := entry.prefix.Bits() + var target *[256][]udpNatEgressEntry + var firstByte byte + if address.Is4() { + target = &buckets.inet4 + firstByte = address.As4()[0] + } else { + target = &buckets.inet6 + firstByte = address.As16()[0] + } + if bits >= 8 { + target[firstByte] = append(target[firstByte], entry) + continue + } + var mask byte + if bits > 0 { + mask = ^byte(0) << (8 - bits) + } + firstByte &= mask + for index := 0; index < 1<<(8-bits); index++ { + bucketIndex := firstByte + byte(index) + target[bucketIndex] = append(target[bucketIndex], entry) + } + } + table.buckets = buckets + return table +} + +func (t *udpNatEgressTable) lookup(address netip.Addr) uint32 { + entries := t.entries + if t.buckets != nil { + if address.Is4() { + entries = t.buckets.inet4[address.As4()[0]] + } else { + entries = t.buckets.inet6[address.As16()[0]] + } + } + for _, entry := range entries { + if entry.prefix.Contains(address) { + return entry.interfaceIndex + } + } + return 0 +} + +func (s *UDPNat) registerClass(conn *udpNatConn) { + s.classAccess.Lock() + conns := s.classConns[conn.interfaceIndex] + if conns == nil { + conns = make(map[*udpNatConn]struct{}) + s.classConns[conn.interfaceIndex] = conns + } + conns[conn] = struct{}{} + s.classAccess.Unlock() +} + +func (s *UDPNat) unregisterClass(conn *udpNatConn) { + s.classAccess.Lock() + conns := s.classConns[conn.interfaceIndex] + if conns != nil { + delete(conns, conn) + if len(conns) == 0 { + delete(s.classConns, conn.interfaceIndex) + } + } + s.classAccess.Unlock() +} + +func (s *UDPNat) NewPacket(bufferSlices [][]byte, source M.Socksaddr, destination M.Socksaddr, userData any) { + conn, ok := s.getOrCreateConn(source, destination, userData) + if !ok { + return + } + readWaitOptions := conn.loadReadWaitOptions() + var dataLen int + for _, bufferSlice := range bufferSlices { + dataLen += len(bufferSlice) + } + buffer := readWaitOptions.NewBufferSize(dataLen) + for _, bufferSlice := range bufferSlices { + buffer.Write(bufferSlice) + } + readWaitOptions.PostReturn(buffer) + conn.enqueue(buffer, destination) +} + +func (s *UDPNat) getOrCreateConn(source M.Socksaddr, destination M.Socksaddr, userData any) (*udpNatConn, bool) { + if s.state.Load() != udpNatStateStarted { + return nil, false + } + key := udpNatSessionKey{ + sourceAddr: source.Addr.Unmap(), + sourcePort: source.Port, + } + switch s.mapping { + case NATMappingEndpointIndependent: + key.interfaceIndex = s.classify(destination) + case NATMappingAddressDependent: + key.destinationAddr = destination.Addr.Unmap() + case NATMappingAddressAndPortDependent: + key.destinationAddr = destination.Addr.Unmap() + key.destinationPort = destination.Port + } + var ( + newContext context.Context + newOnClose N.CloseHandlerFunc + ) + conn, loaded, ok := s.cache.GetAndRefreshOrAdd(key, func() (*udpNatConn, bool) { + ok, ctx, writer, onClose := s.prepare(source, destination, userData) + if !ok { + return nil, false + } + newConn := &udpNatConn{ + service: s, + key: key, + writer: writer, + localAddr: source, + packetChan: make(chan *N.PacketBuffer, 64), + doneChan: make(chan struct{}), + readDeadline: pipe.MakeDeadline(), + } + newConn.cleanupEntry = &udpNatCleanupEntry{ + conn: newConn, + index: -1, + } + if s.filtering != NATFilteringEndpointIndependent { + if destination.IsIP() { + newConn.filterPeer = s.filterPeer(destination) + newConn.filterPeerValid = true + } + if s.filterCache != nil { + filterSessionID := s.nextFilterSessionID.Add(1) + if filterSessionID == 0 { + filterSessionID = s.nextFilterSessionID.Add(1) + } + newConn.filterSessionID = filterSessionID + } + } + interfaceIndex := key.interfaceIndex + if s.mapping != NATMappingEndpointIndependent { + interfaceIndex = s.classify(destination) + } + if interfaceIndex != 0 { + newConn.interfaceIndex = interfaceIndex + s.registerClass(newConn) + } + newContext = ctx + newOnClose = onClose + return newConn, true + }) + if !ok { + return nil, false + } + if s.state.Load() != udpNatStateStarted { + conn.Close() + s.cache.Peek(key) + return nil, false + } + if !loaded { + s.cleanup.addOrUpdate(conn.cleanupEntry, time.Now().Add(s.timeout)) + if conn.isClosed() { + return nil, false + } + go s.handler.NewPacketConnectionEx(newContext, conn, source, destination, newOnClose) + } + conn.addFilterPeer(destination) + return conn, true +} + +func (c *udpNatConn) enqueue(buffer *buf.Buffer, destination M.Socksaddr) { + c.packetAccess.RLock() + select { + case <-c.doneChan: + buffer.Release() + c.packetAccess.RUnlock() + return + default: + } + packet := N.NewPacketBuffer() + *packet = N.PacketBuffer{ + Buffer: buffer, + Destination: destination, + } + select { + case c.packetChan <- packet: + default: + packet.Buffer.Release() + N.PutPacketBuffer(packet) + } + c.packetAccess.RUnlock() +} + +func (s *UDPNat) NewPacketBatch(buffers []*buf.Buffer, sources []M.Socksaddr, destination M.Socksaddr, userData any) { + if len(buffers) != len(sources) { + buf.ReleaseMulti(buffers) + return + } + for index, buffer := range buffers { + conn, ok := s.getOrCreateConn(sources[index], destination, userData) + if !ok { + buffer.Release() + continue + } + readWaitOptions := conn.loadReadWaitOptions() + conn.enqueue(readWaitOptions.Copy(buffer), destination) + } +} + +func (s *UDPNat) filterPeer(destination M.Socksaddr) netip.AddrPort { + if s.filtering == NATFilteringAddressDependent { + return netip.AddrPortFrom(destination.Addr.Unmap(), 0) + } + return netip.AddrPortFrom(destination.Addr.Unmap(), destination.Port) +} + +func (s *UDPNat) Purge() { + if s.filterCache != nil { + s.filterCache.Purge() + } + s.cache.Purge() +} + +func (s *UDPNat) PurgeExpired() { + s.cache.PurgeExpired() +} + +var ( + _ N.PacketConn = (*udpNatConn)(nil) + _ canceler.PacketConn = (*udpNatConn)(nil) + _ N.PacketBatchReadWaitCreator = (*udpNatConn)(nil) + _ N.PacketBatchWriteCreator = (*udpNatConn)(nil) +) + +type udpNatConn struct { + service *UDPNat + key udpNatSessionKey + interfaceIndex uint32 + writer N.PacketWriter + localAddr M.Socksaddr + packetChan chan *N.PacketBuffer + packetAccess sync.RWMutex + closeOnce sync.Once + doneChan chan struct{} + readDeadline pipe.Deadline + readWaitOptions atomic.Pointer[N.ReadWaitOptions] + readBatch *udpNatReadBatch + cleanupEntry *udpNatCleanupEntry + filterSessionID uint64 + filterPeer netip.AddrPort + filterPeerValid bool + filterAccess sync.Mutex + filterPeers map[netip.AddrPort]struct{} +} + +type udpNatReadBatch struct { + buffers []*buf.Buffer + destinations []M.Socksaddr +} + +func (c *udpNatConn) loadReadWaitOptions() N.ReadWaitOptions { + options := c.readWaitOptions.Load() + if options == nil { + return N.ReadWaitOptions{} + } + return *options +} + +func (c *udpNatConn) addFilterPeer(destination M.Socksaddr) { + if c.filterSessionID == 0 || !destination.IsIP() { + return + } + key := udpNatFilterKey{ + sessionID: c.filterSessionID, + peer: c.service.filterPeer(destination), + } + if c.filterPeerValid && c.filterPeer == key.peer { + return + } + if c.isClosed() || c.service.state.Load() != udpNatStateStarted { + return + } + c.service.filterCache.Add(key, c) + c.filterAccess.Lock() + if c.isClosed() || c.service.state.Load() != udpNatStateStarted { + c.filterAccess.Unlock() + c.service.filterCache.Remove(key) + return + } + if c.filterPeers == nil { + c.filterPeers = make(map[netip.AddrPort]struct{}) + } + c.filterPeers[key.peer] = struct{}{} + c.filterAccess.Unlock() + filterConn, loaded := c.service.filterCache.Peek(key) + if !loaded || filterConn != c { + c.removeFilterPeer(key.peer) + return + } + if c.isClosed() || c.service.state.Load() != udpNatStateStarted { + c.service.filterCache.Remove(key) + } +} + +func (c *udpNatConn) removeFilterPeer(peer netip.AddrPort) { + c.filterAccess.Lock() + delete(c.filterPeers, peer) + c.filterAccess.Unlock() +} + +func (c *udpNatConn) clearFilterPeers() { + if c.filterSessionID == 0 { + return + } + c.filterAccess.Lock() + filterPeers := c.filterPeers + c.filterPeers = nil + c.filterAccess.Unlock() + for peer := range filterPeers { + c.service.filterCache.Remove(udpNatFilterKey{ + sessionID: c.filterSessionID, + peer: peer, + }) + } +} + +func (c *udpNatConn) allowPeer(destination M.Socksaddr) bool { + if c.service.filtering == NATFilteringEndpointIndependent || !destination.IsIP() { + return true + } + peer := c.service.filterPeer(destination) + if c.filterPeerValid && c.filterPeer == peer { + return true + } + if c.filterSessionID == 0 { + return false + } + filterConn, loaded := c.service.filterCache.Get(udpNatFilterKey{ + sessionID: c.filterSessionID, + peer: peer, + }) + return loaded && filterConn == c +} + +func (c *udpNatConn) ReadPacket(buffer *buf.Buffer) (addr M.Socksaddr, err error) { + select { + case p := <-c.packetChan: + _, err = buffer.ReadOnceFrom(p.Buffer) + destination := p.Destination + p.Buffer.Release() + N.PutPacketBuffer(p) + return destination, err + case <-c.doneChan: + return M.Socksaddr{}, io.ErrClosedPipe + case <-c.readDeadline.Wait(): + return M.Socksaddr{}, os.ErrDeadlineExceeded + } +} + +func (c *udpNatConn) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + if !c.allowPeer(destination) { + buffer.Release() + return nil + } + return c.writer.WritePacket(buffer, destination) +} + +func (c *udpNatConn) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) { + if c.service.filtering != NATFilteringEndpointIndependent { + return nil, false + } + if creator, isCreator := c.writer.(N.PacketBatchWriteCreator); isCreator { + return creator.CreatePacketBatchWriter() + } + if writer, isWriter := c.writer.(N.PacketBatchWriter); isWriter { + return writer, true + } + return nil, false +} + +func (c *udpNatConn) InitializeReadWaiter(options N.ReadWaitOptions) (needCopy bool) { + c.readWaitOptions.Store(&options) + return false +} + +func (c *udpNatConn) WaitReadPacket() (buffer *buf.Buffer, destination M.Socksaddr, err error) { + return c.waitReadPacket(c.loadReadWaitOptions()) +} + +func (c *udpNatConn) waitReadPacket(options N.ReadWaitOptions) (buffer *buf.Buffer, destination M.Socksaddr, err error) { + select { + case packet := <-c.packetChan: + buffer = options.Copy(packet.Buffer) + destination = packet.Destination + N.PutPacketBuffer(packet) + return + case <-c.doneChan: + return nil, M.Socksaddr{}, io.ErrClosedPipe + case <-c.readDeadline.Wait(): + return nil, M.Socksaddr{}, os.ErrDeadlineExceeded + } +} + +func (c *udpNatConn) CreatePacketBatchReadWaiter() (N.PacketBatchReadWaiter, bool) { + return c, true +} + +func (c *udpNatConn) WaitReadPackets() (buffers []*buf.Buffer, destinations []M.Socksaddr, err error) { + options := c.loadReadWaitOptions() + buffer, destination, err := c.waitReadPacket(options) + if err != nil { + return nil, nil, err + } + batchSize := options.BatchSize + if batchSize <= 0 { + batchSize = 1 + } + batch := c.readBatch + if batch == nil { + batch = new(udpNatReadBatch) + c.readBatch = batch + } else { + clear(batch.buffers) + clear(batch.destinations) + } + buffers = batch.buffers[:0] + destinations = batch.destinations[:0] + defer func() { + batch.buffers = buffers + batch.destinations = destinations + }() + buffers = append(buffers, buffer) + destinations = append(destinations, destination) + for len(buffers) < batchSize { + select { + case packet := <-c.packetChan: + buffers = append(buffers, options.Copy(packet.Buffer)) + destinations = append(destinations, packet.Destination) + N.PutPacketBuffer(packet) + default: + return + } + } + return +} + +func (c *udpNatConn) Timeout() time.Duration { + rawConn, lifetime, loaded := c.service.cache.PeekWithLifetime(c.key) + if !loaded || rawConn != c { + return 0 + } + if lifetime.UnixMilli() == 0 { + return 0 + } + return time.Until(lifetime) +} + +func (c *udpNatConn) SetTimeout(timeout time.Duration) bool { + updated := c.service.cache.UpdateLifetime(c.key, c, timeout) + if !updated { + return false + } + if timeout == 0 { + c.service.cleanup.remove(c.cleanupEntry) + } else { + c.service.cleanup.addOrUpdate(c.cleanupEntry, time.Now().Add(timeout)) + } + return true +} + +func (c *udpNatConn) Close() error { + c.close() + if c.service.state.Load() == udpNatStateStarted { + c.service.cleanup.addOrUpdate(c.cleanupEntry, time.Now()) + } + return nil +} + +func (c *udpNatConn) close() { + c.closeOnce.Do(func() { + c.packetAccess.Lock() + close(c.doneChan) + drained := false + for !drained { + select { + case packet := <-c.packetChan: + packet.Buffer.Release() + N.PutPacketBuffer(packet) + default: + drained = true + } + } + c.packetAccess.Unlock() + c.clearFilterPeers() + if c.interfaceIndex != 0 { + c.service.unregisterClass(c) + } + }) +} + +func (c *udpNatConn) closeFromCache() { + c.close() + c.service.cleanup.remove(c.cleanupEntry) +} + +func (c *udpNatConn) isClosed() bool { + select { + case <-c.doneChan: + return true + default: + return false + } +} + +func (c *udpNatConn) LocalAddr() net.Addr { + return c.localAddr +} + +func (c *udpNatConn) RemoteAddr() net.Addr { + return M.Socksaddr{} +} + +func (c *udpNatConn) SetDeadline(t time.Time) error { + return os.ErrInvalid +} + +func (c *udpNatConn) SetReadDeadline(t time.Time) error { + c.readDeadline.Set(t) + return nil +} + +func (c *udpNatConn) SetWriteDeadline(t time.Time) error { + return os.ErrInvalid +} + +func (c *udpNatConn) Upstream() any { + return c.writer +} diff --git a/udp_nat_cleanup.go b/udp_nat_cleanup.go new file mode 100644 index 0000000..c63b8af --- /dev/null +++ b/udp_nat_cleanup.go @@ -0,0 +1,219 @@ +package tun + +import ( + "container/heap" + "os" + "sync" + "time" +) + +const ( + udpNatStateCreated uint32 = iota + udpNatStateStarted + udpNatStateClosed +) + +func (s *UDPNat) Start() error { + s.lifecycleAccess.Lock() + defer s.lifecycleAccess.Unlock() + switch s.state.Load() { + case udpNatStateCreated: + if s.interfaceFinder != nil { + s.interfaceElement = s.interfaceFinder.RegisterInterfaceUpdateCallback(s.updateInterfaces) + s.reloadInterfaces() + } + s.state.Store(udpNatStateStarted) + s.cleanupWait.Add(1) + go s.cleanupLoop() + return nil + case udpNatStateStarted: + return nil + default: + return os.ErrClosed + } +} + +type udpNatCleanupEntry struct { + conn *udpNatConn + deadline time.Time + index int +} + +type udpNatCleanupQueue struct { + service *UDPNat + access sync.Mutex + wake chan struct{} + entries udpNatCleanupHeap +} + +func newUDPNatCleanupQueue(service *UDPNat) *udpNatCleanupQueue { + queue := &udpNatCleanupQueue{ + service: service, + wake: make(chan struct{}, 1), + } + return queue +} + +func (q *udpNatCleanupQueue) notify() { + select { + case q.wake <- struct{}{}: + default: + } +} + +func (q *udpNatCleanupQueue) addOrUpdate(entry *udpNatCleanupEntry, deadline time.Time) { + if entry == nil || q.service.state.Load() == udpNatStateClosed { + return + } + q.access.Lock() + now := time.Now() + if entry.conn.isClosed() && deadline.After(now) { + deadline = now + } + entry.deadline = deadline + if entry.index == -1 { + heap.Push(&q.entries, entry) + } else { + heap.Fix(&q.entries, entry.index) + } + q.access.Unlock() + q.notify() +} + +func (q *udpNatCleanupQueue) remove(entry *udpNatCleanupEntry) { + if entry == nil { + return + } + q.access.Lock() + if entry.index != -1 { + heap.Remove(&q.entries, entry.index) + } + q.access.Unlock() + q.notify() +} + +func (q *udpNatCleanupQueue) next() (time.Time, bool) { + q.access.Lock() + defer q.access.Unlock() + if len(q.entries) == 0 { + return time.Time{}, false + } + return q.entries[0].deadline, true +} + +func (q *udpNatCleanupQueue) popDue(now time.Time) *udpNatCleanupEntry { + q.access.Lock() + defer q.access.Unlock() + if len(q.entries) == 0 || q.entries[0].deadline.After(now) { + return nil + } + return heap.Pop(&q.entries).(*udpNatCleanupEntry) +} + +func (q *udpNatCleanupQueue) clear() { + q.access.Lock() + for _, entry := range q.entries { + entry.index = -1 + } + clear(q.entries) + q.entries = nil + q.access.Unlock() + q.notify() +} + +type udpNatCleanupHeap []*udpNatCleanupEntry + +func (h udpNatCleanupHeap) Len() int { + return len(h) +} + +func (h udpNatCleanupHeap) Less(i int, j int) bool { + return h[i].deadline.Before(h[j].deadline) +} + +func (h udpNatCleanupHeap) Swap(i int, j int) { + h[i], h[j] = h[j], h[i] + h[i].index = i + h[j].index = j +} + +func (h *udpNatCleanupHeap) Push(value any) { + entry := value.(*udpNatCleanupEntry) + entry.index = len(*h) + *h = append(*h, entry) +} + +func (h *udpNatCleanupHeap) Pop() any { + oldItems := *h + lastIndex := len(oldItems) - 1 + entry := oldItems[lastIndex] + oldItems[lastIndex] = nil + entry.index = -1 + *h = oldItems[:lastIndex] + return entry +} + +func (s *UDPNat) cleanupLoop() { + defer s.cleanupWait.Done() + timer := time.NewTimer(time.Hour) + stopUDPNatCleanupTimer(timer) + defer timer.Stop() + for { + select { + case <-s.cleanup.wake: + default: + } + deadline, loaded := s.cleanup.next() + if !loaded { + select { + case <-s.cleanupDone: + return + case <-s.cleanup.wake: + continue + } + } + waitDuration := time.Until(deadline) + if waitDuration > 0 { + timer.Reset(waitDuration) + select { + case <-s.cleanupDone: + stopUDPNatCleanupTimer(timer) + return + case <-s.cleanup.wake: + stopUDPNatCleanupTimer(timer) + continue + case <-timer.C: + } + } + for { + entry := s.cleanup.popDue(time.Now()) + if entry == nil { + break + } + s.cleanupEntry(entry) + } + } +} + +func (s *UDPNat) cleanupEntry(entry *udpNatCleanupEntry) { + conn, lifetime, loaded := s.cache.PeekWithLifetime(entry.conn.key) + if !loaded || conn != entry.conn { + return + } + if lifetime.UnixMilli() == 0 { + return + } + if conn.isClosed() { + lifetime = time.Now() + } + s.cleanup.addOrUpdate(entry, lifetime) +} + +func stopUDPNatCleanupTimer(timer *time.Timer) { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } +}