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 +}