From 994d6ccdbf3fd9c0fcbad4661b54050f077ca9d4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sat, 11 Jul 2026 00:57:57 +0800 Subject: [PATCH 01/10] Add netns support --- netns_linux.go | 73 ++++++++++++++++++++++++++++++++++ netns_other.go | 12 ++++++ redirect_linux.go | 67 +++++++++++++++++++++++-------- redirect_nftables.go | 46 +++++++++++++-------- stack_system.go | 6 ++- tun.go | 1 + tun_linux.go | 95 +++++++++++++++++++++++--------------------- 7 files changed, 220 insertions(+), 80 deletions(-) create mode 100644 netns_linux.go create mode 100644 netns_other.go 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/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_system.go b/stack_system.go index 3cb0cb0..f2e2edc 100644 --- a/stack_system.go +++ b/stack_system.go @@ -28,6 +28,7 @@ type System struct { ctx context.Context tun Tun tunName string + netNs string mtu int handler Handler logger logger.Logger @@ -68,6 +69,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, @@ -135,7 +137,7 @@ func (s *System) start() error { var err error if s.inet4NextAddress.IsValid() { for range 3 { - tcpListener, err = listener.Listen(s.ctx, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0")) + tcpListener, err = listenNetworkNamespace(s.ctx, s.netNs, listener, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0")) if !retryableListenError(err) { break } @@ -150,7 +152,7 @@ func (s *System) start() error { } if s.inet6NextAddress.IsValid() { for range 3 { - tcpListener, err = listener.Listen(s.ctx, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0")) + tcpListener, err = listenNetworkNamespace(s.ctx, s.netNs, listener, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0")) if !retryableListenError(err) { break } diff --git a/tun.go b/tun.go index c6518f4..14344b6 100644 --- a/tun.go +++ b/tun.go @@ -66,6 +66,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) { From d1af8aaf7eaaf52a148712eddaacc590b9619b92 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Thu, 16 Jul 2026 18:45:49 +0800 Subject: [PATCH 02/10] Fix lint errors --- gtcpip/header/ipv4.go | 11 +++++------ gtcpip/header/ipv6_extension_headers.go | 3 +-- gtcpip/header/ndp_options.go | 9 ++++----- internal/fdbased_darwin/endpoint.go | 3 +-- ping/cmsg_windows.go | 7 +++---- ping/socket_linux_unprivileged.go | 3 +-- 6 files changed, 15 insertions(+), 21 deletions(-) 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..54766d7 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" ) @@ -301,7 +300,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/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 } From 95bc107a1c771c1b3be2625eb23c7ac4d3356775 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Thu, 16 Jul 2026 20:57:59 +0800 Subject: [PATCH 03/10] refactor: New udpnat --- go.mod | 2 +- go.sum | 4 +- stack.go | 3 + stack_gvisor.go | 31 +- stack_gvisor_udp.go | 18 +- stack_mixed.go | 18 +- stack_system.go | 183 ++++++++-- stack_system_packet.go | 3 +- udp_nat.go | 789 +++++++++++++++++++++++++++++++++++++++++ udp_nat_cleanup.go | 219 ++++++++++++ 10 files changed, 1218 insertions(+), 52 deletions(-) create mode 100644 udp_nat.go create mode 100644 udp_nat_cleanup.go diff --git a/go.mod b/go.mod index 8bab948..441c024 100644 --- a/go.mod +++ b/go.mod @@ -11,7 +11,7 @@ require ( github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a github.com/sagernet/nftables v0.3.0-mod.2 - github.com/sagernet/sing v0.8.0 + github.com/sagernet/sing v0.8.12-0.20260716111929-074fc9988b34 github.com/stretchr/testify v1.11.1 go4.org/netipx v0.0.0-20231129151722-fdeea329fbba golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 diff --git a/go.sum b/go.sum index 54ed5fd..63dbfd4 100644 --- a/go.sum +++ b/go.sum @@ -24,8 +24,8 @@ github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZN 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/sing v0.8.12-0.20260716111929-074fc9988b34 h1:rgSs2ttiz8EaubsOt0SkzsqciY0m0PRp3w/fOisPoNo= +github.com/sagernet/sing v0.8.12-0.20260716111929-074fc9988b34/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= diff --git a/stack.go b/stack.go index eaf2405..b2d9568 100644 --- a/stack.go +++ b/stack.go @@ -23,6 +23,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..8a02601 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) @@ -125,6 +139,9 @@ func (t *GVisor) Close() error { 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_udp.go b/stack_gvisor_udp.go index 2ae54cf..3dce60a 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) diff --git a/stack_mixed.go b/stack_mixed.go index 4680380..a238622 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() @@ -59,6 +66,9 @@ 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 f2e2edc..148515a 100644 --- a/stack_system.go +++ b/stack_system.go @@ -5,6 +5,7 @@ import ( "errors" "net" "net/netip" + "os" "slices" "syscall" "time" @@ -19,7 +20,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") @@ -48,7 +48,8 @@ type System struct { tcpPort uint16 tcpPort6 uint16 tcpNat *TCPNat - udpNat *udpnat.Service + udpNat *UDPNat + udpNATOptions UDPNatOptions dispatcher *ForwardDispatcher bindInterface bool interfaceFinder control.InterfaceFinder @@ -80,9 +81,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) { @@ -106,6 +115,9 @@ func NewSystem(options StackOptions) (Stack, error) { func (s *System) Close() error { s.dispatcher.Close() + if s.udpNat != nil { + s.udpNat.Close() + } return common.Close( s.tcpListener, s.tcpListener6, @@ -166,7 +178,14 @@ func (s *System) start() error { go s.acceptLoop(tcpListener) } 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() @@ -684,20 +703,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()), @@ -706,12 +727,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 { @@ -722,14 +792,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) @@ -744,12 +816,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_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/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: + } + } +} From 1ba7d79118e60c6964f1b8c268eb1a7dffcec2d8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 17 Jul 2026 10:40:08 +0800 Subject: [PATCH 04/10] Add UDPEgressPool --- go.mod | 2 +- go.sum | 4 +- udp_egress.go | 287 +++++++++++++++++++++++++++++++++++++++++++++ udp_egress_conn.go | 124 ++++++++++++++++++++ 4 files changed, 414 insertions(+), 3 deletions(-) create mode 100644 udp_egress.go create mode 100644 udp_egress_conn.go diff --git a/go.mod b/go.mod index 441c024..10ad126 100644 --- a/go.mod +++ b/go.mod @@ -11,7 +11,7 @@ require ( github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a github.com/sagernet/nftables v0.3.0-mod.2 - github.com/sagernet/sing v0.8.12-0.20260716111929-074fc9988b34 + 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 diff --git a/go.sum b/go.sum index 63dbfd4..bf39ee5 100644 --- a/go.sum +++ b/go.sum @@ -24,8 +24,8 @@ github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZN 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.12-0.20260716111929-074fc9988b34 h1:rgSs2ttiz8EaubsOt0SkzsqciY0m0PRp3w/fOisPoNo= -github.com/sagernet/sing v0.8.12-0.20260716111929-074fc9988b34/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA= +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= 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 +} From 79084fa79883527d5f70d009a46a9ecd715ada3c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 19 Jul 2026 13:02:55 +0800 Subject: [PATCH 05/10] Add stateless DNS hijack --- flow.go | 1 + flow_dispatch.go | 7 ++++ flow_dns.go | 89 +++++++++++++++++++++++++++++++++++++++++++++ stack_gvisor_udp.go | 20 +++++++--- tun.go | 2 + 5 files changed, 113 insertions(+), 6 deletions(-) create mode 100644 flow_dns.go 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..c5e6ea6 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -291,6 +291,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 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/stack_gvisor_udp.go b/stack_gvisor_udp.go index 3dce60a..5cd0c93 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -71,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/tun.go b/tun.go index 14344b6..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 } From b59636919cbf65d526e95cf184f69b653d3ed92c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 19 Jul 2026 17:41:50 +0800 Subject: [PATCH 06/10] Add Stack.ResetNetwork --- flow_dispatch.go | 23 ++++++++++++++++++----- stack.go | 1 + stack_gvisor.go | 10 ++++++++++ stack_gvisor_icmp.go | 9 +++++++++ stack_mixed.go | 7 +++++++ stack_system.go | 10 ++++++++++ stack_system_nat.go | 9 +++++++++ 7 files changed, 64 insertions(+), 5 deletions(-) diff --git a/flow_dispatch.go b/flow_dispatch.go index c5e6ea6..87f2c02 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -114,11 +114,12 @@ type ForwardDispatcher struct { udpTimeout time.Duration icmpTimeout time.Duration - 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 @@ -544,10 +545,22 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) { } } +func (d *ForwardDispatcher) ResetNetwork() { + if d == nil { + return + } + d.resetPending.Store(true) +} + func (d *ForwardDispatcher) Flush() { if d == nil { 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/stack.go b/stack.go index b2d9568..613e45d 100644 --- a/stack.go +++ b/stack.go @@ -14,6 +14,7 @@ import ( type Stack interface { Start() error + ResetNetwork() Close() error } diff --git a/stack_gvisor.go b/stack_gvisor.go index 8a02601..c226d05 100644 --- a/stack_gvisor.go +++ b/stack_gvisor.go @@ -134,6 +134,16 @@ 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 { diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 11e82af..55cbbd5 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() diff --git a/stack_mixed.go b/stack_mixed.go index a238622..69c8b27 100644 --- a/stack_mixed.go +++ b/stack_mixed.go @@ -62,6 +62,13 @@ 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 diff --git a/stack_system.go b/stack_system.go index 148515a..dd561da 100644 --- a/stack_system.go +++ b/stack_system.go @@ -113,6 +113,16 @@ 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 { s.dispatcher.Close() if s.udpNat != nil { 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] From e5c21070ae46dfd24c0921a593a1ac1ff859e78a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Tue, 21 Jul 2026 14:48:16 +0800 Subject: [PATCH 07/10] Fix flow close race --- flow_dispatch.go | 30 ++++++++++++++++++++++++++---- 1 file changed, 26 insertions(+), 4 deletions(-) diff --git a/flow_dispatch.go b/flow_dispatch.go index 87f2c02..232895a 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -3,6 +3,7 @@ package tun import ( "maps" "net/netip" + "sync" "sync/atomic" "time" @@ -113,6 +114,7 @@ type ForwardDispatcher struct { logger logger.Logger udpTimeout time.Duration icmpTimeout time.Duration + access sync.RWMutex table map[flowKey]*flowEntry lastSweep int64 @@ -168,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() + defer d.access.RUnlock() + if d.returnPath.closed.Load() { + return false + } key := parsed.flowKey() now := d.now() entry, loaded := d.table[key] @@ -553,7 +570,12 @@ func (d *ForwardDispatcher) ResetNetwork() { } func (d *ForwardDispatcher) Flush() { - if d == nil { + 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) { From 2d9b8aed5fe22e8194030ad268bbbd1db3c62630 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 29 Jul 2026 13:45:28 +0800 Subject: [PATCH 08/10] Fix deadlock between flow judgement and close --- flow_dispatch.go | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/flow_dispatch.go b/flow_dispatch.go index 232895a..e872177 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -201,8 +201,8 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool { return false } d.access.RLock() - defer d.access.RUnlock() if d.returnPath.closed.Load() { + d.access.RUnlock() return false } key := parsed.flowKey() @@ -213,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 { @@ -273,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 { From da24acaf4de3896e28021f89bf77be83e3a33678 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 5 Aug 2026 08:10:16 +0800 Subject: [PATCH 09/10] Update gvisor to 20260727.0 --- go.mod | 26 +++++++------- go.sum | 50 ++++++++++++++------------- internal/fdbased_darwin/endpoint.go | 8 ----- internal/fdbased_darwin/processors.go | 37 +++++++++++++------- stack_gvisor_icmp.go | 8 ++++- 5 files changed, 71 insertions(+), 58 deletions(-) diff --git a/go.mod b/go.mod index 10ad126..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/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 bf39ee5..0c72db9 100644 --- a/go.sum +++ b/go.sum @@ -1,29 +1,31 @@ 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/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= @@ -32,17 +34,17 @@ github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1Y 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/internal/fdbased_darwin/endpoint.go b/internal/fdbased_darwin/endpoint.go index 54766d7..b3292a2 100644 --- a/internal/fdbased_darwin/endpoint.go +++ b/internal/fdbased_darwin/endpoint.go @@ -199,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 @@ -256,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") } 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/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 55cbbd5..70e27ec 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -155,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, From d31d20ba5811fa8106c214e994078c9ae8ca154f Mon Sep 17 00:00:00 2001 From: Leadaxe <247031499+Leadaxe@users.noreply.github.com> Date: Fri, 31 Jul 2026 23:54:20 +0300 Subject: [PATCH 10/10] system stack: self-heal the TCP forwarder accept loop (sing-box-lx SPEC 040) Upstream acceptLoop treats any Accept error as terminal and silently returns, leaving the stack alive but every new TCP SYN NAT-rewritten onto a dead port (instant RST) until a full restart. When the listener fd is closed out from under the stack (a stray close on a reused fd number from another runtime in the same process), all new TCP dies forever while UDP/QUIC/DNS keep working. - System.Close() now marks a deliberate shutdown first; acceptLoop still exits quietly on it. - Any other Accept error is logged (the errno names the killer path), the listener is recreated on the same address, the forwarder port is republished atomically, and the loop keeps serving. - If the rebind fails, the loop logs an error and gives up - no worse than upstream. - acceptRecoveries counter is kept as telemetry. tcpPort/tcpPort6 become atomic (written by the heal path, read from the tunLoop dispatch/NAT path); listener replacement is serialized against Close() with a mutex. --- stack_system.go | 166 ++++++++++++++++++++++++++-------- stack_system_selfheal_test.go | 117 ++++++++++++++++++++++++ 2 files changed, 244 insertions(+), 39 deletions(-) create mode 100644 stack_system_selfheal_test.go diff --git a/stack_system.go b/stack_system.go index dd561da..41644bd 100644 --- a/stack_system.go +++ b/stack_system.go @@ -7,6 +7,8 @@ import ( "net/netip" "os" "slices" + "sync" + "sync/atomic" "syscall" "time" @@ -45,17 +47,25 @@ type System struct { icmpTimeout time.Duration tcpListener net.Listener tcpListener6 net.Listener - tcpPort uint16 - tcpPort6 uint16 - tcpNat *TCPNat - udpNat *UDPNat - udpNATOptions UDPNatOptions - 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 { @@ -124,10 +134,15 @@ func (s *System) 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, @@ -143,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 { @@ -155,37 +172,50 @@ 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 = listenNetworkNamespace(s.ctx, s.netNs, listener, "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 = listenNetworkNamespace(s.ctx, s.netNs, listener, "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) udpNATOptions := s.udpNATOptions @@ -367,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) @@ -383,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: @@ -392,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: @@ -411,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: @@ -475,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()) @@ -501,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 @@ -512,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()) @@ -538,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 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) + } +}