refactor: New udpnat

This commit is contained in:
世界 2026-07-16 20:57:59 +08:00
parent d1af8aaf7e
commit 95bc107a1c
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
10 changed files with 1218 additions and 52 deletions

2
go.mod
View file

@ -11,7 +11,7 @@ require (
github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1
github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a
github.com/sagernet/nftables v0.3.0-mod.2 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 github.com/stretchr/testify v1.11.1
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8

4
go.sum
View file

@ -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/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 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.2/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ=
github.com/sagernet/sing v0.8.0 h1:OwLEwbcYfZHvu4olZVljxxC1XRicBqJ1HfiFr6F2WEE= github.com/sagernet/sing v0.8.12-0.20260716111929-074fc9988b34 h1:rgSs2ttiz8EaubsOt0SkzsqciY0m0PRp3w/fOisPoNo=
github.com/sagernet/sing v0.8.0/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak= 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 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8= github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=

View file

@ -23,6 +23,9 @@ type StackOptions struct {
TunOptions Options TunOptions Options
UDPTimeout time.Duration UDPTimeout time.Duration
ICMPTimeout time.Duration ICMPTimeout time.Duration
UDPMapping NATMapping
UDPFiltering NATFiltering
UDPNATMax uint32
Handler Handler Handler Handler
Logger logger.Logger Logger logger.Logger
ForwarderBindInterface bool ForwarderBindInterface bool

View file

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

View file

@ -8,7 +8,6 @@ import (
"net/netip" "net/netip"
"os" "os"
"sync" "sync"
"time"
_ "unsafe" _ "unsafe"
"github.com/sagernet/gvisor/pkg/buffer" "github.com/sagernet/gvisor/pkg/buffer"
@ -21,26 +20,35 @@ import (
E "github.com/sagernet/sing/common/exceptions" E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata" M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network" N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/udpnat2"
) )
type UDPForwarder struct { type UDPForwarder struct {
ctx context.Context ctx context.Context
stack *stack.Stack stack *stack.Stack
handler Handler 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{ forwarder := &UDPForwarder{
ctx: ctx, ctx: ctx,
stack: stack, stack: stack,
handler: handler, handler: handler,
} }
forwarder.udpNat = udpnat.New(handler, forwarder.PreparePacketConnection, timeout, false) options.Handler = handler
options.Prepare = forwarder.PreparePacketConnection
forwarder.udpNat = NewUDPNat(options)
return forwarder 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 { func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort) source := M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort)
destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort) destination := M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort)

View file

@ -22,6 +22,7 @@ type Mixed struct {
tun GVisorTun tun GVisorTun
stack *stack.Stack stack *stack.Stack
endpoint *channel.Endpoint endpoint *channel.Endpoint
udpForwarder *UDPForwarder
} }
func NewMixed( func NewMixed(
@ -47,7 +48,13 @@ func (m *Mixed) Start() error {
if err != nil { if err != nil {
return err 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.stack = ipStack
m.endpoint = endpoint m.endpoint = endpoint
go m.tunLoop() go m.tunLoop()
@ -59,6 +66,9 @@ func (m *Mixed) Close() error {
if m.stack == nil { if m.stack == nil {
return nil return nil
} }
if m.udpForwarder != nil {
m.udpForwarder.Close()
}
m.endpoint.Attach(nil) m.endpoint.Attach(nil)
m.stack.Close() m.stack.Close()
for _, endpoint := range m.stack.CleanupEndpoints() { for _, endpoint := range m.stack.CleanupEndpoints() {

View file

@ -5,6 +5,7 @@ import (
"errors" "errors"
"net" "net"
"net/netip" "net/netip"
"os"
"slices" "slices"
"syscall" "syscall"
"time" "time"
@ -19,7 +20,6 @@ import (
"github.com/sagernet/sing/common/logger" "github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata" M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network" 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") 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 tcpPort uint16
tcpPort6 uint16 tcpPort6 uint16
tcpNat *TCPNat tcpNat *TCPNat
udpNat *udpnat.Service udpNat *UDPNat
udpNATOptions UDPNatOptions
dispatcher *ForwardDispatcher dispatcher *ForwardDispatcher
bindInterface bool bindInterface bool
interfaceFinder control.InterfaceFinder interfaceFinder control.InterfaceFinder
@ -80,6 +81,14 @@ func NewSystem(options StackOptions) (Stack, error) {
inet4Prefixes: options.TunOptions.Inet4Address, inet4Prefixes: options.TunOptions.Inet4Address,
inet6Prefixes: options.TunOptions.Inet6Address, inet6Prefixes: options.TunOptions.Inet6Address,
broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address), broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address),
udpNATOptions: UDPNatOptions{
Timeout: options.UDPTimeout,
Mapping: options.UDPMapping,
Filtering: options.UDPFiltering,
MaxSize: options.UDPNATMax,
InterfaceFinder: options.InterfaceFinder,
ExcludeInterface: []string{options.TunOptions.Name},
},
bindInterface: options.ForwarderBindInterface, bindInterface: options.ForwarderBindInterface,
interfaceFinder: options.InterfaceFinder, interfaceFinder: options.InterfaceFinder,
multiPendingPackets: options.TunOptions.EXP_MultiPendingPackets, multiPendingPackets: options.TunOptions.EXP_MultiPendingPackets,
@ -106,6 +115,9 @@ func NewSystem(options StackOptions) (Stack, error) {
func (s *System) Close() error { func (s *System) Close() error {
s.dispatcher.Close() s.dispatcher.Close()
if s.udpNat != nil {
s.udpNat.Close()
}
return common.Close( return common.Close(
s.tcpListener, s.tcpListener,
s.tcpListener6, s.tcpListener6,
@ -166,7 +178,14 @@ func (s *System) start() error {
go s.acceptLoop(tcpListener) go s.acceptLoop(tcpListener)
} }
s.tcpNat = NewNat(s.ctx, s.udpTimeout) 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 { if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom() s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload() s.txChecksumOffload = linuxTUN.TXChecksumOffload()
@ -684,20 +703,22 @@ type systemUDPPacketWriter4 struct {
txChecksumOffload bool txChecksumOffload bool
} }
func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { func (w *systemUDPPacketWriter4) FrontHeadroom() int {
newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len()) return w.frontHeadroom + len(w.header)
defer newPacket.Release() }
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(w.header) func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
newPacket.Write(buffer.Bytes()) payloadLen := buffer.Len()
ipHdr := header.IPv4(newPacket.Bytes()) buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer)
ipHdr.SetTotalLength(uint16(newPacket.Len())) copy(buffer.ExtendHeader(len(w.header)), w.header)
ipHdr := header.IPv4(buffer.Bytes())
ipHdr.SetTotalLength(uint16(buffer.Len()))
ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetDestinationAddress(ipHdr.SourceAddress())
ipHdr.SetSourceAddr(destination.Addr) ipHdr.SetSourceAddr(destination.Addr)
udpHdr := header.UDP(ipHdr.Payload()) udpHdr := header.UDP(ipHdr.Payload())
udpHdr.SetDestinationPort(udpHdr.SourcePort()) udpHdr.SetDestinationPort(udpHdr.SourcePort())
udpHdr.SetSourcePort(destination.Port) udpHdr.SetSourcePort(destination.Port)
udpHdr.SetLength(uint16(buffer.Len() + header.UDPMinimumSize)) udpHdr.SetLength(uint16(payloadLen + header.UDPMinimumSize))
if !w.txChecksumOffload { if !w.txChecksumOffload {
udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum( udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum(
header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), ipHdr.PayloadLength()), 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) udpHdr.SetChecksum(0)
} }
ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
if PacketOffset > 0 { return buffer
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version) }
} else {
newPacket.Advance(-w.frontHeadroom) func (w *systemUDPPacketWriter4) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
buffer = w.preparePacket(buffer, destination)
if PacketOffset > 0 {
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv4Version)
}
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 {
buffer.Advance(-remainingHeadroom)
}
return buffer
}
func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
buffer = w.prepareWritePacket(buffer, destination)
defer buffer.Release()
return common.Error(w.tun.Write(buffer.Bytes()))
}
func (w *systemUDPPacketWriter4) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
switch w.tun.(type) {
case LinuxTUN, DarwinTUN:
return w, true
default:
return nil, false
}
}
func (w *systemUDPPacketWriter4) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error {
if len(buffers) == 0 || len(buffers) != len(destinations) {
buf.ReleaseMulti(buffers)
return os.ErrInvalid
}
defer func() {
buf.ReleaseMulti(buffers)
}()
switch tunInterface := w.tun.(type) {
case LinuxTUN:
packets := make([][]byte, len(buffers))
for index, buffer := range buffers {
buffer = w.preparePacket(buffer, destinations[index])
buffer.Advance(-w.frontHeadroom)
buffers[index] = buffer
packets[index] = buffer.Bytes()
}
return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom))
case DarwinTUN:
for index, buffer := range buffers {
buffers[index] = w.preparePacket(buffer, destinations[index])
}
return tunInterface.BatchWrite(buffers)
default:
return os.ErrInvalid
} }
return common.Error(w.tun.Write(newPacket.Bytes()))
} }
type systemUDPPacketWriter6 struct { type systemUDPPacketWriter6 struct {
@ -722,14 +792,16 @@ type systemUDPPacketWriter6 struct {
txChecksumOffload bool txChecksumOffload bool
} }
func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { func (w *systemUDPPacketWriter6) FrontHeadroom() int {
newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len()) return w.frontHeadroom + len(w.header)
defer newPacket.Release() }
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(w.header) func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
newPacket.Write(buffer.Bytes()) payloadLen := buffer.Len()
ipHdr := header.IPv6(newPacket.Bytes()) buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer)
udpLen := uint16(header.UDPMinimumSize + buffer.Len()) copy(buffer.ExtendHeader(len(w.header)), w.header)
ipHdr := header.IPv6(buffer.Bytes())
udpLen := uint16(header.UDPMinimumSize + payloadLen)
ipHdr.SetPayloadLength(udpLen) ipHdr.SetPayloadLength(udpLen)
ipHdr.SetDestinationAddress(ipHdr.SourceAddress()) ipHdr.SetDestinationAddress(ipHdr.SourceAddress())
ipHdr.SetSourceAddr(destination.Addr) ipHdr.SetSourceAddr(destination.Addr)
@ -744,12 +816,61 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
} else { } else {
udpHdr.SetChecksum(0) udpHdr.SetChecksum(0)
} }
if PacketOffset > 0 { return buffer
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version) }
} else {
newPacket.Advance(-w.frontHeadroom) func (w *systemUDPPacketWriter6) prepareWritePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer {
buffer = w.preparePacket(buffer, destination)
if PacketOffset > 0 {
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv6Version)
}
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 {
buffer.Advance(-remainingHeadroom)
}
return buffer
}
func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
buffer = w.prepareWritePacket(buffer, destination)
defer buffer.Release()
return common.Error(w.tun.Write(buffer.Bytes()))
}
func (w *systemUDPPacketWriter6) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
switch w.tun.(type) {
case LinuxTUN, DarwinTUN:
return w, true
default:
return nil, false
}
}
func (w *systemUDPPacketWriter6) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error {
if len(buffers) == 0 || len(buffers) != len(destinations) {
buf.ReleaseMulti(buffers)
return os.ErrInvalid
}
defer func() {
buf.ReleaseMulti(buffers)
}()
switch tunInterface := w.tun.(type) {
case LinuxTUN:
packets := make([][]byte, len(buffers))
for index, buffer := range buffers {
buffer = w.preparePacket(buffer, destinations[index])
buffer.Advance(-w.frontHeadroom)
buffers[index] = buffer
packets[index] = buffer.Bytes()
}
return common.Error(tunInterface.BatchWrite(packets, w.frontHeadroom))
case DarwinTUN:
for index, buffer := range buffers {
buffers[index] = w.preparePacket(buffer, destinations[index])
}
return tunInterface.BatchWrite(buffers)
default:
return os.ErrInvalid
} }
return common.Error(w.tun.Write(newPacket.Bytes()))
} }
func newSystemWriteback(tunInterface Tun, frontHeadroom int) ForwardWriteback { func newSystemWriteback(tunInterface Tun, frontHeadroom int) ForwardWriteback {

View file

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

789
udp_nat.go Normal file
View file

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

219
udp_nat_cleanup.go Normal file
View file

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