Compare commits

..

No commits in common. "d31d20ba5811fa8106c214e994078c9ae8ca154f" and "d0d4ebd8dbae7c5c6035d90360de8ea04b0cfe42" have entirely different histories.

31 changed files with 265 additions and 2391 deletions

View file

@ -21,7 +21,6 @@ const (
ActionReject ActionReject
ActionDrop ActionDrop
ActionBypass ActionBypass
ActionHijackDNS
) )
type FlowTracker interface { type FlowTracker interface {

View file

@ -3,7 +3,6 @@ package tun
import ( import (
"maps" "maps"
"net/netip" "net/netip"
"sync"
"sync/atomic" "sync/atomic"
"time" "time"
@ -114,11 +113,9 @@ type ForwardDispatcher struct {
logger logger.Logger logger logger.Logger
udpTimeout time.Duration udpTimeout time.Duration
icmpTimeout time.Duration icmpTimeout time.Duration
access sync.RWMutex
table map[flowKey]*flowEntry table map[flowKey]*flowEntry
lastSweep int64 lastSweep int64
resetPending atomic.Bool
ports map[Port]*portNAT ports map[Port]*portNAT
natList atomic.Pointer[[]*portNAT] natList atomic.Pointer[[]*portNAT]
revNAT atomic.Pointer[map[netip.Addr]*portNAT] revNAT atomic.Pointer[map[netip.Addr]*portNAT]
@ -170,41 +167,26 @@ func (d *ForwardDispatcher) Close() {
return return
} }
d.returnPath.closed.Store(true) d.returnPath.closed.Store(true)
d.access.Lock()
flows := make([]*forwardFlow, 0, len(d.table))
for _, entry := range d.table { for _, entry := range d.table {
if entry.flow != nil { if entry.flow != nil {
flows = append(flows, entry.flow) entry.flow.close(FlowCloseReset)
} }
} }
ports := make([]Port, 0, len(d.ports))
for port, nat := range d.ports { for port, nat := range d.ports {
if nat != nil { if nat != nil {
ports = append(ports, port)
}
}
d.access.Unlock()
for _, flow := range flows {
flow.close(FlowCloseReset)
}
for _, port := range ports {
port.DetachReturn(&d.returnPath) port.DetachReturn(&d.returnPath)
} }
}
} }
func (d *ForwardDispatcher) Dispatch(packet []byte) bool { func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
if d == nil || d.returnPath.closed.Load() { if d == nil {
return false return false
} }
parsed, ok := parseForwardPacket(packet) parsed, ok := parseForwardPacket(packet)
if !ok || parsed.fragment || !parsed.hasFlow { if !ok || parsed.fragment || !parsed.hasFlow {
return false return false
} }
d.access.RLock()
if d.returnPath.closed.Load() {
d.access.RUnlock()
return false
}
key := parsed.flowKey() key := parsed.flowKey()
now := d.now() now := d.now()
entry, loaded := d.table[key] entry, loaded := d.table[key]
@ -213,16 +195,13 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
loaded = false loaded = false
} }
if loaded { if loaded {
handled := d.handleHit(key, entry, &parsed, packet, now) return d.handleHit(key, entry, &parsed, packet, now)
d.access.RUnlock()
return handled
} }
d.access.RUnlock()
if parsed.protocol == uint8(header.TCPProtocolNumber) && if parsed.protocol == uint8(header.TCPProtocolNumber) &&
(parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) { (parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) {
return false return false
} }
return d.judgeAndInstall(key, &parsed, packet) return d.judgeAndInstall(key, &parsed, packet, now)
} }
func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool { func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool {
@ -276,18 +255,12 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
} }
} }
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte) bool { func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
var firstPacket []byte var firstPacket []byte
if packet.protocol == uint8(header.UDPProtocolNumber) { if packet.protocol == uint8(header.UDPProtocolNumber) {
firstPacket = header.UDP(packet.transport).Payload() firstPacket = header.UDP(packet.transport).Payload()
} }
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination, firstPacket) 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 { switch verdict.Action {
case ActionFlow: case ActionFlow:
if verdict.Port != nil { if verdict.Port != nil {
@ -318,13 +291,6 @@ func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket,
case ActionDrop: case ActionDrop:
d.installSimple(key, ActionDrop, packet.protocol, now) d.installSimple(key, ActionDrop, packet.protocol, now)
return true 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: default:
d.installSimple(key, ActionAccept, packet.protocol, now) d.installSimple(key, ActionAccept, packet.protocol, now)
return false return false
@ -571,27 +537,10 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) {
} }
} }
func (d *ForwardDispatcher) ResetNetwork() { func (d *ForwardDispatcher) Flush() {
if d == nil { if d == nil {
return return
} }
d.resetPending.Store(true)
}
func (d *ForwardDispatcher) Flush() {
if d == nil || d.returnPath.closed.Load() {
return
}
d.access.RLock()
defer d.access.RUnlock()
if d.returnPath.closed.Load() {
return
}
if d.resetPending.Swap(false) {
for key, entry := range d.table {
d.removeEntry(key, entry, FlowCloseReset)
}
}
for _, nat := range d.activeNATs { for _, nat := range d.activeNATs {
d.flushPort(nat) d.flushPort(nat)
} }

View file

@ -1,89 +0,0 @@
package tun
import (
"net/netip"
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
func (d *ForwardDispatcher) hijackDNSPacket(packet *forwardPacket) {
writer := &dnsResponseWriter{
writeback: d.writeback,
source: packet.source,
}
d.handler.NewDNSPacket(header.UDP(packet.transport).Payload(), M.SocksaddrFromNetIP(packet.source), M.SocksaddrFromNetIP(packet.destination), writer)
}
var _ N.PacketWriter = (*dnsResponseWriter)(nil)
type dnsResponseWriter struct {
writeback ForwardWriteback
source netip.AddrPort
}
func (w *dnsResponseWriter) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
defer buffer.Release()
if !destination.IsIP() {
return E.New("invalid destination: ", destination)
}
sourceAddr := w.source.Addr().Unmap()
destinationAddr := destination.Addr.Unmap()
headroom := w.writeback.ReturnHeadroom()
udpLen := header.UDPMinimumSize + buffer.Len()
var (
packet []byte
udpHdr header.UDP
ipHdr header.Network
)
if sourceAddr.Is4() {
if !destinationAddr.Is4() {
return E.New("send IPv6 packet to IPv4 connection")
}
size := header.IPv4MinimumSize + udpLen
packet = make([]byte, headroom+size)
inet4Hdr := header.IPv4(packet[headroom:])
inet4Hdr.Encode(&header.IPv4Fields{
TotalLength: uint16(size),
TTL: synthesizedTTL,
Protocol: uint8(header.UDPProtocolNumber),
SrcAddr: destinationAddr,
DstAddr: sourceAddr,
})
udpHdr = header.UDP(inet4Hdr.Payload())
ipHdr = inet4Hdr
} else {
if destinationAddr.Is4() {
destinationAddr = netip.AddrFrom16(destinationAddr.As16())
}
size := header.IPv6MinimumSize + udpLen
packet = make([]byte, headroom+size)
inet6Hdr := header.IPv6(packet[headroom:])
inet6Hdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(udpLen),
TransportProtocol: header.UDPProtocolNumber,
HopLimit: synthesizedTTL,
SrcAddr: destinationAddr,
DstAddr: sourceAddr,
})
udpHdr = header.UDP(inet6Hdr.Payload())
ipHdr = inet6Hdr
}
udpHdr.Encode(&header.UDPFields{
SrcPort: destination.Port,
DstPort: w.source.Port(),
Length: uint16(udpLen),
})
copy(udpHdr.Payload(), buffer.Bytes())
udpHdr.SetChecksum(^checksum.Checksum(udpHdr.Payload(), udpHdr.CalculateChecksum(
header.PseudoHeaderChecksum(header.UDPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), uint16(udpLen)),
)))
if inet4Hdr, isInet4 := ipHdr.(header.IPv4); isInet4 {
inet4Hdr.SetChecksum(^inet4Hdr.CalculateChecksum())
}
return w.writeback.WriteReturnPackets([][]byte{packet})
}

28
go.mod
View file

@ -1,32 +1,32 @@
module github.com/sagernet/sing-tun module github.com/sagernet/sing-tun
go 1.25.0 go 1.24.7
require ( require (
github.com/florianl/go-nfqueue/v2 v2.1.0 github.com/florianl/go-nfqueue/v2 v2.0.2
github.com/go-ole/go-ole v1.3.0 github.com/go-ole/go-ole v1.3.0
github.com/google/btree v1.1.3 github.com/google/btree v1.1.3
github.com/mdlayher/netlink v1.11.2 github.com/mdlayher/netlink v1.9.0
github.com/sagernet/fswatch v0.1.2 github.com/sagernet/fswatch v0.1.1
github.com/sagernet/gvisor v0.0.0-20260727.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.4 github.com/sagernet/nftables v0.3.0-mod.2
github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8 github.com/sagernet/sing v0.8.0
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-20250711185948-6ae5c78190dc golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8
golang.org/x/net v0.57.0 golang.org/x/net v0.50.0
golang.org/x/sys v0.47.0 golang.org/x/sys v0.41.0
) )
require ( require (
github.com/davecgh/go-spew v1.1.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect
github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/fsnotify/fsnotify v1.7.0 // indirect
github.com/google/go-cmp v0.7.0 // indirect github.com/google/go-cmp v0.7.0 // indirect
github.com/mdlayher/socket v0.6.0 // indirect github.com/mdlayher/socket v0.5.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/vishvananda/netns v0.0.4 // indirect github.com/vishvananda/netns v0.0.4 // indirect
golang.org/x/sync v0.20.0 // indirect golang.org/x/sync v0.7.0 // indirect
golang.org/x/time v0.15.0 // indirect golang.org/x/time v0.7.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect
) )

54
go.sum
View file

@ -1,50 +1,48 @@
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= 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/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/florianl/go-nfqueue/v2 v2.1.0 h1:Fywt30TY/evxyDySpXjxQ1jsRW7nQbLpOhELqpr4068= github.com/florianl/go-nfqueue/v2 v2.0.2 h1:FL5lQTeetgpCvac1TRwSfgaXUn0YSO7WzGvWNIp3JPE=
github.com/florianl/go-nfqueue/v2 v2.1.0/go.mod h1:8PKUM5rYoVFO5IZV1bifx4/b0jHAglKkHXr9PRwzi4Y= github.com/florianl/go-nfqueue/v2 v2.0.2/go.mod h1:VA09+iPOT43OMoCKNfXHyzujQUty2xmzyCRkBOlmabc=
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= github.com/fsnotify/fsnotify v1.7.0 h1:8JEhPFa5W2WU7YfeZzPNqzMP6Lwt7L2715Ggo0nosvA=
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/fsnotify/fsnotify v1.7.0/go.mod h1:40Bi/Hjc2AVfZrqy+aj+yEI+/bRxZnMJyTJwOpGvigM=
github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE=
github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= 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 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4= 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 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/jsimonetti/rtnetlink/v2 v2.2.0 h1:/KfZ310gOAFrXXol5VwnFEt+ucldD/0dsSRZwpHCP9w= github.com/mdlayher/netlink v1.9.0 h1:G8+GLq2x3v4D4MVIqDdNUhTUC7TKiCy/6MDkmItfKco=
github.com/jsimonetti/rtnetlink/v2 v2.2.0/go.mod h1:lbjDHxC+5RJ08lzPeA90Ls2pEoId3F08MoEMlhfHxeI= github.com/mdlayher/netlink v1.9.0/go.mod h1:YBnl5BXsCoRuwBjKKlZ+aYmEoq0r12FDA/3JC+94KDg=
github.com/mdlayher/netlink v1.11.2 h1:HKh2jqe+omdSWcQ88nrT7INE61B0NXfiSPFdgL4YbNI= github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos=
github.com/mdlayher/netlink v1.11.2/go.mod h1:uT2Yc/QLaZubzDpZIBi9d4GoeLwtp3x1AMeqSRrK2sA= github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ=
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 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/sagernet/fswatch v0.1.2 h1:/TT7k4mkce1qFPxamLO842WjqBgbTBiXP2mlUjp9PFk= github.com/sagernet/fswatch v0.1.1 h1:YqID+93B7VRfqIH3PArW/XpJv5H4OLEVWDfProGoRQs=
github.com/sagernet/fswatch v0.1.2/go.mod h1:5BpGmpUQVd3Mc5r313HRpvADHRg3/rKn5QbwFteB880= github.com/sagernet/fswatch v0.1.1/go.mod h1:nz85laH0mkQqJfaOrqPpkwtU1znMFNVTpT/5oRsVz/o=
github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1 h1:IdQ7yTKkB2wv8txwshxUroPlO4npOYAV71xb7xQ7Lys= github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1 h1:AzCE2RhBjLJ4WIWc/GejpNh+z30d5H1hwaB0nD9eY3o=
github.com/sagernet/gvisor v0.0.0-20260727.0-sing-box-mod.1/go.mod h1:9O3SQskYuCfdHNvHEsWuEAgoyKEF74PiWp4NsNUia8g= github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1/go.mod h1:NJKBtm9nVEK3iyOYWsUlrDQuoGh4zJ4KOPhSYVidvQ4=
github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZNjr6sGeT00J8uU7JF4cNUdb44/Duis= github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a h1:ObwtHN2VpqE0ZNjr6sGeT00J8uU7JF4cNUdb44/Duis=
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.4 h1:vnOtcDYeSXv2e5RoRuGH0lrpttQFJ8iC4ICS2nhlDSo= github.com/sagernet/nftables v0.3.0-mod.2 h1:ck2KMU02OxL1eDFgGaWYglMDpoOZ7OHzxje+vW5Q0OQ=
github.com/sagernet/nftables v0.3.0-mod.4/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.12-0.20260717023913-84ab32b56cb8 h1:dyRIj+MZ2rc9JVzJoG04jxu+MpvHrLIZLJr0QjNAMGg= github.com/sagernet/sing v0.8.0 h1:OwLEwbcYfZHvu4olZVljxxC1XRicBqJ1HfiFr6F2WEE=
github.com/sagernet/sing v0.8.12-0.20260717023913-84ab32b56cb8/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA= github.com/sagernet/sing v0.8.0/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1 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=
github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM= 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 h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y= go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc h1:TS73t7x3KarrNd5qAipmspBDS1rkMcgVG/fS1aRb4Rc= golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8 h1:yixxcjnhBmY0nkL253HFVIm0JsFHwrHdT3Yh6szTnfY=
golang.org/x/exp v0.0.0-20250711185948-6ae5c78190dc/go.mod h1:A+z0yzpGtvnG90cToK5n2tu8UJVP2XUATh+r+sfOOOc= golang.org/x/exp v0.0.0-20240613232115-7f521ea00fb8/go.mod h1:jj3sYF3dwk5D+ghuXyeI3r5MFf+NT2An6/9dOA95KSI=
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.7.0 h1:YsImfSBoP9QPYL0xyKJPq0gcaJdG3rInoqxTWbfQu9M=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=

View file

@ -22,6 +22,7 @@ import (
"github.com/sagernet/sing-tun/gtcpip" "github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/checksum" "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 // RFC 971 defines the fields of the IPv4 header on page 11 using the following
@ -334,7 +335,7 @@ func (b IPv4) FragmentOffset() uint16 {
} }
func (b IPv4) FragmentOffsetDarwinRaw() uint16 { func (b IPv4) FragmentOffsetDarwinRaw() uint16 {
return binary.NativeEndian.Uint16(b[flagsFO:]) << 3 return common.NativeEndian.Uint16(b[flagsFO:]) << 3
} }
// TotalLength returns the "total length" field of the IPv4 header. // TotalLength returns the "total length" field of the IPv4 header.
@ -343,7 +344,7 @@ func (b IPv4) TotalLength() uint16 {
} }
func (b IPv4) TotalLengthDarwinRaw() uint16 { func (b IPv4) TotalLengthDarwinRaw() uint16 {
return binary.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength()) return common.NativeEndian.Uint16(b[IPv4TotalLenOffset:]) + uint16(b.HeaderLength())
} }
// Checksum returns the checksum field of the IPv4 header. // Checksum returns the checksum field of the IPv4 header.
@ -440,7 +441,7 @@ func (b IPv4) SetTotalLength(totalLength uint16) {
} }
func (b IPv4) SetTotalLengthDarwinRaw(totalLength uint16) { func (b IPv4) SetTotalLengthDarwinRaw(totalLength uint16) {
binary.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength) common.NativeEndian.PutUint16(b[IPv4TotalLenOffset:], totalLength)
} }
// SetChecksum sets the checksum field of the IPv4 header. // SetChecksum sets the checksum field of the IPv4 header.
@ -457,7 +458,7 @@ func (b IPv4) SetFlagsFragmentOffset(flags uint8, offset uint16) {
func (b IPv4) SetFlagsFragmentOffsetDarwinRaw(flags uint8, offset uint16) { func (b IPv4) SetFlagsFragmentOffsetDarwinRaw(flags uint8, offset uint16) {
v := (uint16(flags) << 13) | (offset >> 3) v := (uint16(flags) << 13) | (offset >> 3)
binary.NativeEndian.PutUint16(b[flagsFO:], v) common.NativeEndian.PutUint16(b[flagsFO:], v)
} }
// SetID sets the identification field. // SetID sets the identification field.
@ -1178,7 +1179,7 @@ func (s IPv4OptionsSerializer) Serialize(b []byte) uint8 {
// header ends on a 32 bit boundary. The padding is zero. // header ends on a 32 bit boundary. The padding is zero.
padded := padIPv4OptionsLength(total) padded := padIPv4OptionsLength(total)
b = b[:padded-total] b = b[:padded-total]
clear(b) common.ClearArray(b)
return padded return padded
} }

View file

@ -21,6 +21,7 @@ import (
"math" "math"
"github.com/sagernet/sing-tun/gtcpip" "github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing/common"
) )
// IPv6ExtensionHeaderIdentifier is an IPv6 extension header identifier. // IPv6ExtensionHeaderIdentifier is an IPv6 extension header identifier.
@ -128,7 +129,7 @@ func padIPv6Option(b []byte) {
b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6Pad1ExtHdrOptionIdentifier) b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6Pad1ExtHdrOptionIdentifier)
default: // Pad with PadN. default: // Pad with PadN.
s := b[ipv6ExtHdrOptionPayloadOffset:] s := b[ipv6ExtHdrOptionPayloadOffset:]
clear(s) common.ClearArray(s)
b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6PadNExtHdrOptionIdentifier) b[ipv6ExtHdrOptionTypeOffset] = uint8(ipv6PadNExtHdrOptionIdentifier)
b[ipv6ExtHdrOptionLengthOffset] = uint8(len(s)) b[ipv6ExtHdrOptionLengthOffset] = uint8(len(s))
} }

View file

@ -24,6 +24,7 @@ import (
"time" "time"
"github.com/sagernet/sing-tun/gtcpip" "github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing/common"
) )
// ndpOptionIdentifier is an NDP option type identifier. // ndpOptionIdentifier is an NDP option type identifier.
@ -340,7 +341,7 @@ func (b NDPOptions) Serialize(s NDPOptionsSerializer) int {
// Zero out remaining (padding) bytes, if any exists. // Zero out remaining (padding) bytes, if any exists.
if used+2 < l { if used+2 < l {
clear(b[used+2 : l]) common.ClearArray(b[used+2 : l])
} }
b = b[l:] b = b[l:]
@ -566,7 +567,7 @@ func (o NDPPrefixInformation) serializeInto(b []byte) int {
// Zero out the Reserved2 field. // Zero out the Reserved2 field.
reserved2 := b[ndpPrefixInformationReserved2Offset:][:ndpPrefixInformationReserved2Length] reserved2 := b[ndpPrefixInformationReserved2Offset:][:ndpPrefixInformationReserved2Length]
clear(reserved2) common.ClearArray(reserved2)
return used return used
} }
@ -685,7 +686,7 @@ func (o NDPRecursiveDNSServer) serializeInto(b []byte) int {
used := copy(b, o) used := copy(b, o)
// Zero out the reserved bytes that are before the Lifetime field. // Zero out the reserved bytes that are before the Lifetime field.
clear(b[0:ndpRecursiveDNSServerLifetimeOffset]) common.ClearArray(b[0:ndpRecursiveDNSServerLifetimeOffset])
return used return used
} }
@ -778,7 +779,7 @@ func (o NDPDNSSearchList) serializeInto(b []byte) int {
used := copy(b, o) used := copy(b, o)
// Zero out the reserved bytes that are before the Lifetime field. // Zero out the reserved bytes that are before the Lifetime field.
clear(b[0:ndpDNSSearchListLifetimeOffset]) common.ClearArray(b[0:ndpDNSSearchListLifetimeOffset])
return used return used
} }

View file

@ -50,6 +50,7 @@ import (
"github.com/sagernet/gvisor/pkg/tcpip/header" "github.com/sagernet/gvisor/pkg/tcpip/header"
"github.com/sagernet/gvisor/pkg/tcpip/stack" "github.com/sagernet/gvisor/pkg/tcpip/stack"
rawfile "github.com/sagernet/sing-tun/internal/rawfile_darwin" rawfile "github.com/sagernet/sing-tun/internal/rawfile_darwin"
"github.com/sagernet/sing/common"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
@ -199,6 +200,10 @@ type Options struct {
// include CapabilitySaveRestore // include CapabilitySaveRestore
SaveRestore bool 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 // PacketDispatchMode specifies the type of inbound dispatcher to be
// used for this endpoint. // used for this endpoint.
PacketDispatchMode PacketDispatchMode PacketDispatchMode PacketDispatchMode
@ -252,6 +257,10 @@ func New(opts *Options) (stack.LinkEndpoint, error) {
caps |= stack.CapabilitySaveRestore caps |= stack.CapabilitySaveRestore
} }
if opts.DisconnectOk {
caps |= stack.CapabilityDisconnectOk
}
if len(opts.FDs) == 0 { if len(opts.FDs) == 0 {
return nil, fmt.Errorf("opts.FD is empty, at least one FD must be specified") return nil, fmt.Errorf("opts.FD is empty, at least one FD must be specified")
} }
@ -292,7 +301,7 @@ func New(opts *Options) (stack.LinkEndpoint, error) {
e.fds = append(e.fds, fdInfo{fd: fd, isSocket: true}) e.fds = append(e.fds, fdInfo{fd: fd, isSocket: true})
if opts.ProcessorsPerChannel == 0 { if opts.ProcessorsPerChannel == 0 {
opts.ProcessorsPerChannel = max(1, runtime.GOMAXPROCS(0)/len(opts.FDs)) opts.ProcessorsPerChannel = common.Max(1, runtime.GOMAXPROCS(0)/len(opts.FDs))
} }
inboundDispatcher, err := newRecvMMsgDispatcher(fd, e, opts) inboundDispatcher, err := newRecvMMsgDispatcher(fd, e, opts)

View file

@ -214,47 +214,34 @@ func tcpipConnectionID(pkt *stack.PacketBuffer) (connectionID, bool) {
return cid, true return cid, true
} }
ipHdr := header.IPv6(h) ipHdr := header.IPv6(h)
cid.srcAddr = ipHdr.SourceAddressSlice()
cid.dstAddr = ipHdr.DestinationAddressSlice()
cid.proto = header.IPv6ProtocolNumber
if !header.IsExtensionHeader(ipHdr.NextHeader()) { var tcpHdr header.TCP
// Known transport protocols(not just TCP) store the src and dst ports if tcpip.TransportProtocolNumber(ipHdr.NextHeader()) == header.TCPProtocolNumber {
// in the first 4 bytes after the IPv6 fixed header. tcpHdr = header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen])
tcpHdr := header.TCP(h[header.IPv6FixedHeaderSize:][:tcpSrcDstPortLen])
cid.srcPort = tcpHdr.SourcePort()
cid.dstPort = tcpHdr.DestinationPort()
} else { } else {
// Slow path for IPv6 extension headers :(. // Slow path for IPv6 extension headers :(.
dataBuf := pkt.Data().ToBuffer() dataBuf := pkt.Data().ToBuffer()
dataBuf.TrimFront(header.IPv6MinimumSize) dataBuf.TrimFront(header.IPv6MinimumSize)
it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(ipHdr.NextHeader()), dataBuf) it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(ipHdr.NextHeader()), dataBuf)
defer it.Release() 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 { for {
hdr, done, err := it.Next() hdr, done, err := it.Next()
if done || err != nil { if done || err != nil {
break break
} }
if fh, ok := hdr.(header.IPv6FragmentExtHdr); ok && !fh.IsAtomic() {
isFragment = true
}
hdr.Release() hdr.Release()
} }
if !isFragment {
h, ok = pkt.Data().PullUp(int(it.HeaderOffset()) + tcpSrcDstPortLen) h, ok = pkt.Data().PullUp(int(it.HeaderOffset()) + tcpSrcDstPortLen)
if !ok { if !ok {
return cid, true return cid, true
} }
// Known transport protocols store the src and dst ports tcpHdr = header.TCP(h[it.HeaderOffset():][:tcpSrcDstPortLen])
// in the first 4 bytes after the IPv6 fixed header. }
tcpHdr := header.TCP(h[it.HeaderOffset():][:tcpSrcDstPortLen]) cid.srcAddr = ipHdr.SourceAddressSlice()
cid.dstAddr = ipHdr.DestinationAddressSlice()
cid.srcPort = tcpHdr.SourcePort() cid.srcPort = tcpHdr.SourcePort()
cid.dstPort = tcpHdr.DestinationPort() cid.dstPort = tcpHdr.DestinationPort()
} cid.proto = header.IPv6ProtocolNumber
}
default: default:
return cid, true return cid, true
} }

View file

@ -1,73 +0,0 @@
package tun
import (
"context"
"net"
"runtime"
"strings"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"golang.org/x/sys/unix"
)
func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) {
return execInNetworkNamespace(nameOrPath, func() (net.Listener, error) {
return config.Listen(ctx, network, address)
})
}
type networkNamespaceInterfaceFinder struct {
control.InterfaceFinder
options *Options
}
func (f *networkNamespaceInterfaceFinder) Update() error {
return runInNetworkNamespace(f.options.NetNs, f.InterfaceFinder.Update)
}
func execInNetworkNamespace[T any](nameOrPath string, block func() (T, error)) (T, error) {
if nameOrPath == "" {
return block()
}
type blockResult struct {
value T
err error
}
resultChannel := make(chan blockResult, 1)
go func() {
runtime.LockOSThread()
value, err := execInNetworkNamespaceThread(nameOrPath, block)
resultChannel <- blockResult{value, err}
}()
result := <-resultChannel
return result.value, result.err
}
func execInNetworkNamespaceThread[T any](nameOrPath string, block func() (T, error)) (T, error) {
var defaultValue T
var path string
if strings.HasPrefix(nameOrPath, "/") {
path = nameOrPath
} else {
path = "/run/netns/" + nameOrPath
}
targetFd, err := unix.Open(path, unix.O_RDONLY|unix.O_CLOEXEC, 0)
if err != nil {
return defaultValue, E.Cause(err, "open netns ", nameOrPath)
}
defer unix.Close(targetFd)
err = unix.Setns(targetFd, unix.CLONE_NEWNET)
if err != nil {
return defaultValue, E.Cause(err, "set netns to ", nameOrPath)
}
return block()
}
func runInNetworkNamespace(nameOrPath string, block func() error) error {
_, err := execInNetworkNamespace(nameOrPath, func() (struct{}, error) {
return struct{}{}, block()
})
return err
}

View file

@ -1,12 +0,0 @@
//go:build !linux
package tun
import (
"context"
"net"
)
func listenNetworkNamespace(ctx context.Context, nameOrPath string, config net.ListenConfig, network, address string) (net.Listener, error) {
return config.Listen(ctx, network, address)
}

View file

@ -1,10 +1,11 @@
package ping package ping
import ( import (
"encoding/binary"
"fmt" "fmt"
"unsafe" "unsafe"
"github.com/sagernet/sing/common"
"golang.org/x/net/ipv6" "golang.org/x/net/ipv6"
"golang.org/x/sys/windows" "golang.org/x/sys/windows"
) )
@ -36,9 +37,9 @@ func parseIPv6ControlMessage(cmsg []byte) (*ipv6.ControlMessage, error) {
} }
switch cmsghdr.Type { switch cmsghdr.Type {
case IPV6_TCLASS: case IPV6_TCLASS:
controlMessage.TrafficClass = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) controlMessage.TrafficClass = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
case IPV6_HOPLIMIT: case IPV6_HOPLIMIT:
controlMessage.HopLimit = int(binary.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4])) controlMessage.HopLimit = int(common.NativeEndian.Uint32(cmsg[alignedSizeofCmsghdr : alignedSizeofCmsghdr+4]))
} }
cmsg = cmsg[msgSize:] cmsg = cmsg[msgSize:]
} }

View file

@ -9,6 +9,7 @@ import (
"time" "time"
"github.com/sagernet/sing-tun/gtcpip/header" "github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf" "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control" "github.com/sagernet/sing/common/control"
M "github.com/sagernet/sing/common/metadata" M "github.com/sagernet/sing/common/metadata"
@ -174,7 +175,7 @@ func (c *UnprivilegedConn) Close() error {
for _, conn := range c.mapping { for _, conn := range c.mapping {
_ = conn.Close() _ = conn.Close()
} }
clear(c.mapping) common.ClearMap(c.mapping)
return nil return nil
} }

View file

@ -26,7 +26,6 @@ type autoRedirect struct {
logger logger.Logger logger logger.Logger
tableName string tableName string
networkMonitor NetworkUpdateMonitor networkMonitor NetworkUpdateMonitor
ownedNetworkMonitor bool
networkListener *list.Element[NetworkUpdateCallback] networkListener *list.Element[NetworkUpdateCallback]
interfaceFinder control.InterfaceFinder interfaceFinder control.InterfaceFinder
localAddresses []netip.Prefix localAddresses []netip.Prefix
@ -52,7 +51,7 @@ type autoRedirect struct {
} }
func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) { func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) {
r := &autoRedirect{ return &autoRedirect{
tunOptions: options.TunOptions, tunOptions: options.TunOptions,
ctx: options.Context, ctx: options.Context,
handler: options.Handler, handler: options.Handler,
@ -64,11 +63,7 @@ func NewAutoRedirect(options AutoRedirectOptions) (AutoRedirect, error) {
customRedirectPortFunc: options.CustomRedirectPort, customRedirectPortFunc: options.CustomRedirectPort,
routeAddressSet: options.RouteAddressSet, routeAddressSet: options.RouteAddressSet,
routeExcludeAddressSet: options.RouteExcludeAddressSet, routeExcludeAddressSet: options.RouteExcludeAddressSet,
} }, nil
if options.TunOptions.NetNs != "" {
r.interfaceFinder = &networkNamespaceInterfaceFinder{control.NewDefaultInterfaceFinder(), options.TunOptions}
}
return r, nil
} }
func (r *autoRedirect) Start() error { func (r *autoRedirect) Start() error {
@ -94,11 +89,8 @@ func (r *autoRedirect) Start() error {
} }
} }
} else { } else {
if r.tunOptions.NetNs != "" && !r.useNFTables {
return E.New("auto_redirect in network namespace requires nftables")
}
if r.useNFTables { if r.useNFTables {
err = runInNetworkNamespace(r.tunOptions.NetNs, r.initializeNFTables) err = r.initializeNFTables()
if err != nil { if err != nil {
return E.Cause(err, "missing nftables support") return E.Cause(err, "missing nftables support")
} }
@ -140,7 +132,7 @@ func (r *autoRedirect) Start() error {
listenAddr = netip.IPv4Unspecified() listenAddr = netip.IPv4Unspecified()
} }
server := newRedirectServer(r.ctx, r.handler, r.logger, listenAddr) server := newRedirectServer(r.ctx, r.handler, r.logger, listenAddr)
err = runInNetworkNamespace(r.tunOptions.NetNs, server.Start) err = server.Start()
if err != nil { if err != nil {
return E.Cause(err, "start redirect server") return E.Cause(err, "start redirect server")
} }
@ -159,43 +151,24 @@ func (r *autoRedirect) Start() error {
}) })
if err != nil { if err != nil {
r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err) r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
} else if err = runInNetworkNamespace(r.tunOptions.NetNs, handler.Start); err != nil { } else if err = handler.Start(); err != nil {
r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err) r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
} else { } else {
r.nfqueueHandler = handler r.nfqueueHandler = handler
r.nfqueueEnabled = true r.nfqueueEnabled = true
} }
} }
if r.tunOptions.NetNs != "" {
var monitor NetworkUpdateMonitor
monitor, err = NewNetworkUpdateMonitor(r.logger)
if err != nil {
return E.Cause(err, "create netns network monitor")
}
err = runInNetworkNamespace(r.tunOptions.NetNs, monitor.Start)
if err != nil {
return E.Cause(err, "start netns network monitor")
}
r.networkMonitor = monitor
r.ownedNetworkMonitor = true
}
err = runInNetworkNamespace(r.tunOptions.NetNs, func() error {
r.cleanupNFTables() r.cleanupNFTables()
setupErr := r.setupNFTables() err = r.setupNFTables()
if setupErr != nil { if err != nil {
return E.Cause(setupErr, "setup nftables") return E.Cause(err, "setup nftables")
} }
if r.tunOptions.AutoRedirectMarkMode { if r.tunOptions.AutoRedirectMarkMode {
setupErr = r.setupRedirectRoutes() err = r.setupRedirectRoutes()
if setupErr != nil {
r.cleanupNFTables()
return E.Cause(setupErr, "setup redirect routes")
}
}
return nil
})
if err != nil { if err != nil {
return err r.cleanupNFTables()
return E.Cause(err, "setup redirect routes")
}
} }
} else { } else {
r.cleanupIPTables() r.cleanupIPTables()
@ -212,14 +185,8 @@ func (r *autoRedirect) Close() error {
r.nfqueueHandler.Close() r.nfqueueHandler.Close()
} }
if r.useNFTables { if r.useNFTables {
_ = runInNetworkNamespace(r.tunOptions.NetNs, func() error {
r.cleanupNFTables() r.cleanupNFTables()
r.cleanupRedirectRoutes() r.cleanupRedirectRoutes()
return nil
})
if r.ownedNetworkMonitor {
_ = r.networkMonitor.Close()
}
} else { } else {
r.cleanupIPTables() r.cleanupIPTables()
} }
@ -230,7 +197,7 @@ func (r *autoRedirect) Close() error {
func (r *autoRedirect) UpdateRouteAddressSet() { func (r *autoRedirect) UpdateRouteAddressSet() {
if r.useNFTables { if r.useNFTables {
err := runInNetworkNamespace(r.tunOptions.NetNs, r.nftablesUpdateRouteAddressSet) err := r.nftablesUpdateRouteAddressSet()
if err != nil { if err != nil {
r.logger.Error("update route address set: ", err) r.logger.Error("update route address set: ", err)
} }

View file

@ -299,38 +299,27 @@ func (r *autoRedirect) setupNFTables() error {
if err != nil { if err != nil {
return E.Cause(err, "flush nftables") return E.Cause(err, "flush nftables")
} }
if r.tunOptions.NetNs == "" {
r.startDockerFirewallMonitor() r.startDockerFirewallMonitor()
err = r.configureDockerFirewall(false) err = r.configureDockerFirewall(false)
if err != nil && r.logger != nil { if err != nil && r.logger != nil {
r.logger.Warn("configure docker firewall: ", err) r.logger.Warn("configure docker firewall: ", err)
} }
}
r.networkListener = r.networkMonitor.RegisterCallback(func() { r.networkListener = r.networkMonitor.RegisterCallback(func() {
updateErr := runInNetworkNamespace(r.tunOptions.NetNs, r.updateNetworkAddresses) err = r.nftablesUpdateLocalAddressSet()
if updateErr != nil { if err != nil {
r.logger.Error(updateErr) r.logger.Error("update local address set: ", err)
}
if r.tunOptions.AutoRedirectMarkMode {
err = r.updateRedirectRoutes()
if err != nil {
r.logger.Error("update redirect routes: ", err)
}
} }
}) })
return nil 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 // TODO: test if this works
func (r *autoRedirect) nftablesUpdateLocalAddressSet() error { func (r *autoRedirect) nftablesUpdateLocalAddressSet() error {
err := r.interfaceFinder.Update() err := r.interfaceFinder.Update()
@ -387,7 +376,6 @@ func (r *autoRedirect) nftablesUpdateRouteAddressSet() error {
func (r *autoRedirect) cleanupNFTables() { func (r *autoRedirect) cleanupNFTables() {
if r.networkListener != nil { if r.networkListener != nil {
r.networkMonitor.UnregisterCallback(r.networkListener) r.networkMonitor.UnregisterCallback(r.networkListener)
r.networkListener = nil
} }
r.stopDockerFirewallMonitor() r.stopDockerFirewallMonitor()
nft, err := nftables.New() nft, err := nftables.New()
@ -401,12 +389,10 @@ func (r *autoRedirect) cleanupNFTables() {
_ = r.configureOpenWRTFirewall4(nft, true) _ = r.configureOpenWRTFirewall4(nft, true)
_ = nft.Flush() _ = nft.Flush()
_ = nft.CloseLasting() _ = nft.CloseLasting()
if r.tunOptions.NetNs == "" {
err = r.configureDockerFirewall(true) err = r.configureDockerFirewall(true)
if err != nil && r.logger != nil { if err != nil && r.logger != nil {
r.logger.Warn("cleanup docker firewall: ", err) r.logger.Warn("cleanup docker firewall: ", err)
} }
}
} }
func (r *autoRedirect) nftablesCreatePreMatchChains(nft *nftables.Conn, table *nftables.Table) error { func (r *autoRedirect) nftablesCreatePreMatchChains(nft *nftables.Conn, table *nftables.Table) error {

View file

@ -14,7 +14,6 @@ import (
type Stack interface { type Stack interface {
Start() error Start() error
ResetNetwork()
Close() error Close() error
} }
@ -24,9 +23,6 @@ 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,7 +44,6 @@ 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 {
@ -79,15 +78,8 @@ 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,
@ -101,7 +93,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.udpNATOptions.Timeout, t.icmpTimeout) t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout)
} }
linkEndpoint = &LinkEndpointFilter{ linkEndpoint = &LinkEndpointFilter{
LinkEndpoint: linkEndpoint, LinkEndpoint: linkEndpoint,
@ -118,13 +110,7 @@ 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)
udpForwarder := NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpNATOptions) ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket)
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)
@ -134,24 +120,11 @@ func (t *GVisor) Start() error {
return nil 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 { func (t *GVisor) Close() error {
t.dispatcher.Close() t.dispatcher.Close()
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

@ -72,15 +72,6 @@ func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger)
return forwarder 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 { func (f *ICMPForwarder) Close() error {
f.returnPath.closed.Store(true) f.returnPath.closed.Store(true)
f.flowAccess.Lock() f.flowAccess.Lock()
@ -155,15 +146,9 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
} else { } else {
ipHdr := header.IPv6(pkt.NetworkHeader().Slice()) ipHdr := header.IPv6(pkt.NetworkHeader().Slice())
icmpHdr := header.ICMPv6(pkt.TransportHeader().Slice()) icmpHdr := header.ICMPv6(pkt.TransportHeader().Slice())
if icmpHdr.Type() != header.ICMPv6EchoRequest { if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false 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() identifier := icmpHdr.Ident()
key := icmpFlowKey{ key := icmpFlowKey{
v6: true, v6: true,

View file

@ -8,6 +8,7 @@ import (
"net/netip" "net/netip"
"os" "os"
"sync" "sync"
"time"
_ "unsafe" _ "unsafe"
"github.com/sagernet/gvisor/pkg/buffer" "github.com/sagernet/gvisor/pkg/buffer"
@ -20,35 +21,26 @@ 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 udpNat *udpnat.Service
} }
func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, options UDPNatOptions) *UDPForwarder { func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, timeout time.Duration) *UDPForwarder {
forwarder := &UDPForwarder{ forwarder := &UDPForwarder{
ctx: ctx, ctx: ctx,
stack: stack, stack: stack,
handler: handler, handler: handler,
} }
options.Handler = handler forwarder.udpNat = udpnat.New(handler, forwarder.PreparePacketConnection, timeout, false)
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)
@ -71,26 +63,18 @@ func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M
firstPacket = append(firstPacket[:len(firstPacket):len(firstPacket)], view.AsSlice()...) 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 { switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort(), firstPacket).Action {
case ActionReject: case ActionReject:
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer)) gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
return false, nil, nil, nil return false, nil, nil, nil
case ActionDrop: case ActionDrop:
return false, nil, nil, nil return false, nil, nil, nil
case ActionHijackDNS: }
f.handler.NewDNSPacket(firstPacket, source, destination, &UDPBackWriter{ var sourceNetwork tcpip.NetworkProtocolNumber
stack: f.stack, if source.Addr.Is4() {
source: AddressFromAddr(source.Addr), sourceNetwork = header.IPv4ProtocolNumber
sourcePort: source.Port, } else {
sourceNetwork: sourceNetwork, sourceNetwork = header.IPv6ProtocolNumber
})
return false, nil, nil, nil
} }
writer := &UDPBackWriter{ writer := &UDPBackWriter{
stack: f.stack, stack: f.stack,

View file

@ -22,7 +22,6 @@ 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(
@ -48,13 +47,7 @@ func (m *Mixed) Start() error {
if err != nil { if err != nil {
return err return err
} }
udpForwarder := NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpNATOptions) ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(m.ctx, ipStack, m.handler, m.udpTimeout).HandlePacket)
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()
@ -62,20 +55,10 @@ func (m *Mixed) Start() error {
return nil return nil
} }
func (m *Mixed) ResetNetwork() {
m.System.ResetNetwork()
if m.udpForwarder != nil {
m.udpForwarder.udpNat.Purge()
}
}
func (m *Mixed) Close() error { 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,10 +5,7 @@ import (
"errors" "errors"
"net" "net"
"net/netip" "net/netip"
"os"
"slices" "slices"
"sync"
"sync/atomic"
"syscall" "syscall"
"time" "time"
@ -22,6 +19,7 @@ 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")
@ -30,7 +28,6 @@ type System struct {
ctx context.Context ctx context.Context
tun Tun tun Tun
tunName string tunName string
netNs string
mtu int mtu int
handler Handler handler Handler
logger logger.Logger logger logger.Logger
@ -47,19 +44,10 @@ type System struct {
icmpTimeout time.Duration icmpTimeout time.Duration
tcpListener net.Listener tcpListener net.Listener
tcpListener6 net.Listener tcpListener6 net.Listener
// lx/040: ports are written by acceptLoop on self-heal relisten and read tcpPort uint16
// concurrently from the tunLoop path (dispatch filter + NAT rewrite) — tcpPort6 uint16
// 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 tcpNat *TCPNat
udpNat *UDPNat udpNat *udpnat.Service
udpNATOptions UDPNatOptions
dispatcher *ForwardDispatcher dispatcher *ForwardDispatcher
bindInterface bool bindInterface bool
interfaceFinder control.InterfaceFinder interfaceFinder control.InterfaceFinder
@ -80,7 +68,6 @@ func NewSystem(options StackOptions) (Stack, error) {
ctx: options.Context, ctx: options.Context,
tun: options.Tun, tun: options.Tun,
tunName: options.TunOptions.Name, tunName: options.TunOptions.Name,
netNs: options.TunOptions.NetNs,
mtu: int(options.TunOptions.MTU), mtu: int(options.TunOptions.MTU),
inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress, inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress,
inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress, inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress,
@ -91,14 +78,6 @@ 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,
@ -123,26 +102,8 @@ func NewSystem(options StackOptions) (Stack, error) {
return stack, nil 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 { 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() s.dispatcher.Close()
if s.udpNat != nil {
s.udpNat.Close()
}
s.listenAccess.Lock()
defer s.listenAccess.Unlock()
return common.Close( return common.Close(
s.tcpListener, s.tcpListener,
s.tcpListener6, s.tcpListener6,
@ -158,10 +119,8 @@ func (s *System) Start() error {
return nil return nil
} }
// lx/040: TCP forwarder bind, shared by start() and the acceptLoop self-heal func (s *System) start() error {
// relisten path. isIPv6 selects the address family; the bind-to-interface _ = fixWindowsFirewall()
// Control and the EADDRNOTAVAIL retry loop match the original start() code.
func (s *System) listenTCP(isIPv6 bool) (net.Listener, error) {
var listener net.ListenConfig var listener net.ListenConfig
if s.bindInterface { if s.bindInterface {
listener.Control = control.Append(listener.Control, func(network, address string, conn syscall.RawConn) error { listener.Control = control.Append(listener.Control, func(network, address string, conn syscall.RawConn) error {
@ -172,60 +131,40 @@ func (s *System) listenTCP(isIPv6 bool) (net.Listener, error) {
return nil return nil
}) })
} }
network := "tcp4" var tcpListener net.Listener
address := s.inet4Address var err error
if isIPv6 { if s.inet4NextAddress.IsValid() {
network = "tcp6"
address = s.inet6Address
}
var (
tcpListener net.Listener
err error
)
for range 3 { for range 3 {
tcpListener, err = listenNetworkNamespace(s.ctx, s.netNs, listener, network, net.JoinHostPort(address.String(), "0")) tcpListener, err = listener.Listen(s.ctx, "tcp4", net.JoinHostPort(s.inet4Address.String(), "0"))
if !retryableListenError(err) { if !retryableListenError(err) {
break break
} }
time.Sleep(time.Second) time.Sleep(time.Second)
} }
if err != nil {
return nil, err
}
return tcpListener, nil
}
func (s *System) start() error {
_ = fixWindowsFirewall()
var tcpListener net.Listener
var err error
if s.inet4NextAddress.IsValid() {
tcpListener, err = s.listenTCP(false)
if err != nil { if err != nil {
return err return err
} }
s.tcpListener = tcpListener s.tcpListener = tcpListener
s.tcpPort.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) s.tcpPort = M.SocksaddrFromNet(tcpListener.Addr()).Port
go s.acceptLoop(tcpListener, false) go s.acceptLoop(tcpListener)
} }
if s.inet6NextAddress.IsValid() { if s.inet6NextAddress.IsValid() {
tcpListener, err = s.listenTCP(true) for range 3 {
tcpListener, err = listener.Listen(s.ctx, "tcp6", net.JoinHostPort(s.inet6Address.String(), "0"))
if !retryableListenError(err) {
break
}
time.Sleep(time.Second)
}
if err != nil { if err != nil {
return err return err
} }
s.tcpListener6 = tcpListener s.tcpListener6 = tcpListener
s.tcpPort6.Store(uint32(M.SocksaddrFromNet(tcpListener.Addr()).Port)) s.tcpPort6 = M.SocksaddrFromNet(tcpListener.Addr()).Port
go s.acceptLoop(tcpListener, true) go s.acceptLoop(tcpListener)
} }
s.tcpNat = NewNat(s.ctx, s.udpTimeout) s.tcpNat = NewNat(s.ctx, s.udpTimeout)
udpNATOptions := s.udpNATOptions s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false)
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()
@ -397,29 +336,12 @@ func (s *System) processPacket(packet []byte) bool {
return writeBack return writeBack
} }
func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) { func (s *System) acceptLoop(listener net.Listener) {
for { for {
conn, err := listener.Accept() conn, err := listener.Accept()
if err != nil { if err != nil {
// lx/040 (SPECS/TASKS/040): upstream silently returns on ANY Accept
// error, leaving the stack alive but every new TCP SYN NAT-rewritten
// onto a dead port (instant RST) until a VPN restart — the LxBox §047
// "browser dead, QUIC alive" failure. A deliberate System.Close is the
// only quiet exit; anything else means the listener died out from
// under us (e.g. a foreign close on a reused fd number from the
// Java side of the shared Android process) — log it (the errno names
// the killer) and recreate the listener.
if s.closing.Load() {
return 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 connPort := M.SocksaddrFromNet(conn.RemoteAddr()).Port
session := s.tcpNat.LookupBack(connPort) session := s.tcpNat.LookupBack(connPort)
if session == nil { if session == nil {
@ -430,47 +352,6 @@ func (s *System) acceptLoop(listener net.Listener, isIPv6 bool) {
} }
} }
// lx/040: recreate a TCP forwarder listener that died out from under the
// stack. Returns the replacement listener after publishing it (listener field
// + atomic port) under listenAccess, or an error if the stack is closing or
// the bind failed.
func (s *System) healListener(dead net.Listener, isIPv6 bool, cause error) (net.Listener, error) {
port := &s.tcpPort
if isIPv6 {
port = &s.tcpPort6
}
oldPort := port.Load()
s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener (port ", oldPort, ") accept failed: ", cause, " — recreating listener")
_ = dead.Close() // release netpoll state; harmless if already closed
newListener, err := s.listenTCP(isIPv6)
if err != nil {
return nil, err
}
s.listenAccess.Lock()
defer s.listenAccess.Unlock()
if s.closing.Load() {
_ = newListener.Close()
return nil, net.ErrClosed
}
if isIPv6 {
s.tcpListener6 = newListener
} else {
s.tcpListener = newListener
}
newPort := uint32(M.SocksaddrFromNet(newListener.Addr()).Port)
port.Store(newPort)
recoveries := s.acceptRecoveries.Add(1)
s.logger.Warn("system stack: tcp", ipVersionSuffix(isIPv6), " listener recreated (port ", oldPort, " → ", newPort, ", recoveries: ", recoveries, ")")
return newListener, nil
}
func ipVersionSuffix(isIPv6 bool) string {
if isIPv6 {
return "6"
}
return "4"
}
func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool { func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() { switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber: case header.TCPProtocolNumber:
@ -480,7 +361,7 @@ func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
if ipHdr.SourceAddr() == s.inet4Address && if ipHdr.SourceAddr() == s.inet4Address &&
ipHdr.FragmentOffset() == 0 && ipHdr.FragmentOffset() == 0 &&
len(ipHdr.Payload()) >= header.TCPMinimumSize && len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort.Load()) { header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort {
return false return false
} }
case header.ICMPv4ProtocolNumber: case header.ICMPv4ProtocolNumber:
@ -499,7 +380,7 @@ func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool {
} }
if ipHdr.SourceAddr() == s.inet6Address && if ipHdr.SourceAddr() == s.inet6Address &&
len(ipHdr.Payload()) >= header.TCPMinimumSize && len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == uint16(s.tcpPort6.Load()) { header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 {
return false return false
} }
case header.ICMPv6ProtocolNumber: case header.ICMPv6ProtocolNumber:
@ -563,7 +444,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort()) destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
if !destination.Addr().IsGlobalUnicast() { if !destination.Addr().IsGlobalUnicast() {
return false, nil return false, nil
} else if source.Addr() == s.inet4Address && source.Port() == uint16(s.tcpPort.Load()) { } else if source.Addr() == s.inet4Address && source.Port() == s.tcpPort {
session := s.tcpNat.LookupBack(destination.Port()) session := s.tcpNat.LookupBack(destination.Port())
if session == nil { if session == nil {
return false, E.New("ipv4: tcp: session not found: ", destination.Port()) return false, E.New("ipv4: tcp: session not found: ", destination.Port())
@ -589,7 +470,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
} }
rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload, rewriteIPv4TCP(ipHdr, tcpHdr, s.txChecksumOffload,
s.inet4NextAddress, natPort, true, s.inet4NextAddress, natPort, true,
s.inet4Address, uint16(s.tcpPort.Load()), true) s.inet4Address, s.tcpPort, true)
} }
} }
return true, nil return true, nil
@ -600,7 +481,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort()) destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
if !destination.Addr().IsGlobalUnicast() { if !destination.Addr().IsGlobalUnicast() {
return false, nil return false, nil
} else if source.Addr() == s.inet6Address && source.Port() == uint16(s.tcpPort6.Load()) { } else if source.Addr() == s.inet6Address && source.Port() == s.tcpPort6 {
session := s.tcpNat.LookupBack(destination.Port()) session := s.tcpNat.LookupBack(destination.Port())
if session == nil { if session == nil {
return false, E.New("ipv6: tcp: session not found: ", destination.Port()) return false, E.New("ipv6: tcp: session not found: ", destination.Port())
@ -626,7 +507,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
} }
rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload, rewriteIPv6TCP(ipHdr, tcpHdr, s.txChecksumOffload,
s.inet6NextAddress, natPort, true, s.inet6NextAddress, natPort, true,
s.inet6Address, uint16(s.tcpPort6.Load()), true) s.inet6Address, s.tcpPort6, true)
} }
} }
return true, nil return true, nil
@ -801,22 +682,20 @@ type systemUDPPacketWriter4 struct {
txChecksumOffload bool txChecksumOffload bool
} }
func (w *systemUDPPacketWriter4) FrontHeadroom() int { func (w *systemUDPPacketWriter4) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
return w.frontHeadroom + len(w.header) newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len())
} defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { newPacket.Write(w.header)
payloadLen := buffer.Len() newPacket.Write(buffer.Bytes())
buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer) ipHdr := header.IPv4(newPacket.Bytes())
copy(buffer.ExtendHeader(len(w.header)), w.header) ipHdr.SetTotalLength(uint16(newPacket.Len()))
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(payloadLen + header.UDPMinimumSize)) udpHdr.SetLength(uint16(buffer.Len() + 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()),
@ -825,61 +704,12 @@ func (w *systemUDPPacketWriter4) preparePacket(buffer *buf.Buffer, destination M
udpHdr.SetChecksum(0) udpHdr.SetChecksum(0)
} }
ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) 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 { if PacketOffset > 0 {
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv4Version) PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} } else {
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 { newPacket.Advance(-w.frontHeadroom)
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 {
@ -890,16 +720,14 @@ type systemUDPPacketWriter6 struct {
txChecksumOffload bool txChecksumOffload bool
} }
func (w *systemUDPPacketWriter6) FrontHeadroom() int { func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
return w.frontHeadroom + len(w.header) newPacket := buf.NewSize(w.frontHeadroom + len(w.header) + buffer.Len())
} defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M.Socksaddr) *buf.Buffer { newPacket.Write(w.header)
payloadLen := buffer.Len() newPacket.Write(buffer.Bytes())
buffer = (N.ReadWaitOptions{FrontHeadroom: w.FrontHeadroom()}).Copy(buffer) ipHdr := header.IPv6(newPacket.Bytes())
copy(buffer.ExtendHeader(len(w.header)), w.header) udpLen := uint16(header.UDPMinimumSize + buffer.Len())
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)
@ -914,61 +742,12 @@ func (w *systemUDPPacketWriter6) preparePacket(buffer *buf.Buffer, destination M
} else { } else {
udpHdr.SetChecksum(0) 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 { if PacketOffset > 0 {
PacketFillHeader(buffer.ExtendHeader(PacketOffset), header.IPv6Version) PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} } else {
if remainingHeadroom := w.frontHeadroom - PacketOffset; remainingHeadroom > 0 { newPacket.Advance(-w.frontHeadroom)
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

@ -86,15 +86,6 @@ func (n *TCPNat) checkTimeout() {
n.addrAccess.Unlock() 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 { func (n *TCPNat) LookupBack(port uint16) *TCPSession {
n.portAccess.RLock() n.portAccess.RLock()
session := n.portMap[port] session := n.portMap[port]

View file

@ -5,6 +5,7 @@ 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 {
@ -13,7 +14,7 @@ func PacketIPVersion(packet []byte) int {
func PacketFillHeader(packet []byte, ipVersion int) { func PacketFillHeader(packet []byte, ipVersion int) {
if PacketOffset > 0 { if PacketOffset > 0 {
clear(packet[:3]) common.ClearArray(packet[:3])
switch ipVersion { switch ipVersion {
case header.IPv4Version: case header.IPv4Version:
packet[3] = syscall.AF_INET packet[3] = syscall.AF_INET

View file

@ -1,117 +0,0 @@
package tun
// lx/040 (SPECS/TASKS/040-SINGTUN_ACCEPTLOOP_SELFHEAL): acceptLoop self-heal.
//
// Red/green против апстрима 2d9b8aed5fe2: там acceptLoop(listener) при любой
// ошибке Accept молча выходит навсегда — восстановления нет, порт не меняется,
// новый connect вечно бьётся в мёртвый сокет. Для red-прогона на чистом
// апстрим-чекауте достаточно адаптировать хелперы ниже (currentTCPPort →
// s.tcpPort, spawnAcceptLoop → go s.acceptLoop(ln)): тест упадёт по таймауту
// ожидания восстановления.
import (
"context"
"fmt"
"net"
"net/netip"
"testing"
"time"
"github.com/sagernet/sing/common/logger"
)
func newSelfHealTestSystem(t *testing.T) *System {
t.Helper()
s := &System{
ctx: context.Background(),
logger: logger.NOP(),
inet4Address: netip.MustParseAddr("127.0.0.1"),
udpTimeout: time.Minute,
}
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
ln, err := s.listenTCP(false)
if err != nil {
t.Fatalf("listenTCP: %v", err)
}
s.tcpListener = ln
s.tcpPort.Store(uint32(ln.Addr().(*net.TCPAddr).Port))
spawnAcceptLoop(s, ln)
return s
}
func currentTCPPort(s *System) uint32 {
return s.tcpPort.Load()
}
func spawnAcceptLoop(s *System, ln net.Listener) {
go s.acceptLoop(ln, false)
}
func dialForwarder(t *testing.T, port uint32) error {
t.Helper()
conn, err := net.DialTimeout("tcp4", fmt.Sprintf("127.0.0.1:%d", port), time.Second)
if err == nil {
_ = conn.Close()
}
return err
}
// Убийство listener'а мимо System.Close (эмуляция чужого close по
// переиспользованному fd-номеру) должно приводить к пересозданию listener'а
// и продолжению приёма TCP, а не к вечной смерти петли.
func TestSystemAcceptLoopSelfHeal(t *testing.T) {
s := newSelfHealTestSystem(t)
oldPort := currentTCPPort(s)
if err := dialForwarder(t, oldPort); err != nil {
t.Fatalf("healthy listener refused connect: %v", err)
}
// Убить listener из-под стека: closing НЕ выставлен.
_ = s.tcpListener.Close()
deadline := time.Now().Add(5 * time.Second)
healed := false
for time.Now().Before(deadline) {
if s.acceptRecoveries.Load() > 0 {
healed = true
break
}
time.Sleep(10 * time.Millisecond)
}
if !healed {
t.Fatalf("acceptLoop did not recover within 5s (upstream behavior: silent permanent death)")
}
newPort := currentTCPPort(s)
if newPort == oldPort {
t.Fatalf("recovered port equals dead port %d — relisten did not publish a new port", oldPort)
}
if err := dialForwarder(t, newPort); err != nil {
t.Fatalf("connect to recreated listener (port %d) failed: %v", newPort, err)
}
if got := s.acceptRecoveries.Load(); got != 1 {
t.Fatalf("acceptRecoveries = %d, want 1", got)
}
s.closing.Store(true)
_ = s.tcpListener.Close()
}
// Штатное закрытие (closing выставлен, как это делает System.Close) обязано
// оставаться тихим: без пересозданий и без роста счётчика.
func TestSystemAcceptLoopQuietOnClose(t *testing.T) {
s := newSelfHealTestSystem(t)
oldPort := currentTCPPort(s)
s.closing.Store(true)
_ = s.tcpListener.Close()
time.Sleep(300 * time.Millisecond)
if got := s.acceptRecoveries.Load(); got != 0 {
t.Fatalf("deliberate close triggered %d recoveries, want 0", got)
}
if port := currentTCPPort(s); port != oldPort {
t.Fatalf("deliberate close changed port %d → %d", oldPort, port)
}
}

3
tun.go
View file

@ -14,14 +14,12 @@ import (
E "github.com/sagernet/sing/common/exceptions" E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format" F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/common/logger" "github.com/sagernet/sing/common/logger"
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/ranges" "github.com/sagernet/sing/common/ranges"
) )
type Handler interface { type Handler interface {
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort, firstPacket []byte) FlowVerdict 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.TCPConnectionHandlerEx
N.UDPConnectionHandlerEx N.UDPConnectionHandlerEx
} }
@ -68,7 +66,6 @@ const (
type Options struct { type Options struct {
Name string Name string
NetNs string
Inet4Address []netip.Prefix Inet4Address []netip.Prefix
Inet6Address []netip.Prefix Inet6Address []netip.Prefix
MTU uint32 MTU uint32

View file

@ -51,8 +51,8 @@ type NativeTun struct {
} }
func New(options Options) (Tun, error) { func New(options Options) (Tun, error) {
var nativeTun *NativeTun
if options.FileDescriptor == 0 { if options.FileDescriptor == 0 {
return execInNetworkNamespace(options.NetNs, func() (Tun, error) {
tunFd, err := open(options.Name, options.GSO) tunFd, err := open(options.Name, options.GSO)
if err != nil { if err != nil {
return nil, E.Cause(err, "open tun") return nil, E.Cause(err, "open tun")
@ -61,7 +61,7 @@ func New(options Options) (Tun, error) {
if err != nil { if err != nil {
return nil, E.Errors(err, unix.Close(tunFd)) return nil, E.Errors(err, unix.Close(tunFd))
} }
nativeTun := &NativeTun{ nativeTun = &NativeTun{
tunFd: tunFd, tunFd: tunFd,
tunFile: os.NewFile(uintptr(tunFd), "tun"), tunFile: os.NewFile(uintptr(tunFd), "tun"),
options: options, options: options,
@ -70,10 +70,8 @@ func New(options Options) (Tun, error) {
if err != nil { if err != nil {
return nil, E.Errors(err, unix.Close(tunFd)) return nil, E.Errors(err, unix.Close(tunFd))
} }
return nativeTun, nil } else {
}) nativeTun = &NativeTun{
}
nativeTun := &NativeTun{
tunFd: options.FileDescriptor, tunFd: options.FileDescriptor,
tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"), tunFile: os.NewFile(uintptr(options.FileDescriptor), "tun"),
options: options, options: options,
@ -86,6 +84,7 @@ func New(options Options) (Tun, error) {
} }
} }
} }
}
return nativeTun, nil return nativeTun, nil
} }
@ -291,10 +290,10 @@ func (t *NativeTun) Name() (string, error) {
func (t *NativeTun) Start() error { func (t *NativeTun) Start() error {
if t.options.FileDescriptor == 0 { if t.options.FileDescriptor == 0 {
if !t.options.EXP_ExternalConfiguration && t.options.NetNs == "" { if !t.options.EXP_ExternalConfiguration {
t.options.InterfaceMonitor.RegisterMyInterface(t.options.Name) t.options.InterfaceMonitor.RegisterMyInterface(t.options.Name)
} }
err := runInNetworkNamespace(t.options.NetNs, t.start) err := t.start()
if err != nil { if err != nil {
return err return err
} }
@ -355,7 +354,7 @@ func (t *NativeTun) start() error {
return E.Cause(err, "set rules") return E.Cause(err, "set rules")
} }
if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { if t.options.DNSMode != DNSModeDisabled {
err = t.setSearchDomainForSystemdResolved() err = t.setSearchDomainForSystemdResolved()
if err != nil { if err != nil {
return E.Cause(err, "set search domain") return E.Cause(err, "set search domain")
@ -375,13 +374,11 @@ func (t *NativeTun) Close() error {
if t.options.EXP_ExternalConfiguration { if t.options.EXP_ExternalConfiguration {
return common.Close(common.PtrOrNil(t.tunFile)) return common.Close(common.PtrOrNil(t.tunFile))
} }
if t.options.DNSMode != DNSModeDisabled && t.options.NetNs == "" { if t.options.DNSMode != DNSModeDisabled {
t.unsetSearchDomainForSystemdResolved() t.unsetSearchDomainForSystemdResolved()
} }
return E.Errors(runInNetworkNamespace(t.options.NetNs, func() error {
t.unsetAddresses() t.unsetAddresses()
return E.Errors(t.unsetRoute(), t.unsetRules()) return E.Errors(t.unsetRoute(), t.unsetRules(), common.Close(common.PtrOrNil(t.tunFile)))
}), common.Close(common.PtrOrNil(t.tunFile)))
} }
func (t *NativeTun) Read(p []byte) (n int, err error) { func (t *NativeTun) Read(p []byte) (n int, err error) {
@ -628,7 +625,6 @@ func (t *NativeTun) UpdateRouteOptions(tunOptions Options) error {
t.options = tunOptions t.options = tunOptions
return nil return nil
} }
return runInNetworkNamespace(t.options.NetNs, func() error {
tunLink, err := netlink.LinkByName(t.options.Name) tunLink, err := netlink.LinkByName(t.options.Name)
if err != nil { if err != nil {
return E.Cause(err, "find tun interface") return E.Cause(err, "find tun interface")
@ -639,7 +635,6 @@ func (t *NativeTun) UpdateRouteOptions(tunOptions Options) error {
} }
t.options = tunOptions t.options = tunOptions
return t.setRoute(tunLink) return t.setRoute(tunLink)
})
} }
func (t *NativeTun) routes(tunLink netlink.Link) ([]netlink.Route, error) { func (t *NativeTun) routes(tunLink netlink.Link) ([]netlink.Route, error) {

View file

@ -1,287 +0,0 @@
package tun
import (
"context"
"net"
"net/netip"
"runtime"
"slices"
"sync"
"sync/atomic"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
"github.com/sagernet/sing/common/x/list"
)
const udpEgressBufferSize = 65535
type UDPEgressPoolOptions struct {
Logger logger.Logger
Network string
Control control.Func
InterfaceFinder control.InterfaceFinder
InterfaceMonitor DefaultInterfaceMonitor
ExcludeInterface string
IsExempt func() bool
}
type UDPEgressPool struct {
logger logger.Logger
network string
control control.Func
interfaceFinder control.InterfaceFinder
interfaceMonitor DefaultInterfaceMonitor
excludeInterface string
isExempt func() bool
access sync.Mutex
port uint16
anchorInterfaceIndex int
receiveDone chan struct{}
members map[udpEgressSpec]*udpEgressMember
state atomic.Pointer[[]*udpEgressMember]
packetChan chan udpEgressPacket
memberReaders sync.WaitGroup
finderElement *list.Element[control.InterfaceUpdateCallback]
}
type udpEgressSpec struct {
interfaceIndex int
interfaceName string
prefix netip.Prefix
}
type udpEgressMember struct {
prefix netip.Prefix
conn *net.UDPConn
}
type udpEgressPacket struct {
buffer *buf.Buffer
source netip.AddrPort
}
func NewUDPEgressPool(options UDPEgressPoolOptions) *UDPEgressPool {
return &UDPEgressPool{
logger: options.Logger,
network: options.Network,
control: options.Control,
interfaceFinder: options.InterfaceFinder,
interfaceMonitor: options.InterfaceMonitor,
excludeInterface: options.ExcludeInterface,
isExempt: options.IsExempt,
anchorInterfaceIndex: -1,
members: make(map[udpEgressSpec]*udpEgressMember),
packetChan: make(chan udpEgressPacket, 128),
}
}
func (p *UDPEgressPool) Close() {
p.SetEgressPort(0)
p.access.Lock()
defer p.access.Unlock()
if p.finderElement != nil {
p.interfaceFinder.UnregisterInterfaceUpdateCallback(p.finderElement)
p.finderElement = nil
}
}
func (p *UDPEgressPool) SetEgressPort(port uint16) bool {
p.access.Lock()
defer p.access.Unlock()
if p.port == port {
return p.state.Load() != nil
}
if p.receiveDone != nil {
close(p.receiveDone)
p.receiveDone = nil
}
p.port = 0
p.state.Store(nil)
for spec, member := range p.members {
delete(p.members, spec)
member.conn.Close()
}
p.memberReaders.Wait()
for {
select {
case packet := <-p.packetChan:
packet.buffer.Release()
default:
goto drained
}
}
drained:
p.anchorInterfaceIndex = -1
if port == 0 {
return false
}
p.port = port
defaultInterface := p.interfaceMonitor.DefaultInterface()
if defaultInterface != nil {
p.anchorInterfaceIndex = defaultInterface.Index
}
p.receiveDone = make(chan struct{})
if p.finderElement == nil {
p.finderElement = p.interfaceFinder.RegisterInterfaceUpdateCallback(func(interfaces []control.Interface) {
p.access.Lock()
defer p.access.Unlock()
p.rebuildLocked()
})
}
p.rebuildLocked()
return p.state.Load() != nil
}
func (p *UDPEgressPool) LookupEgress(destination netip.AddrPort) *net.UDPConn {
members := p.state.Load()
if members == nil {
return nil
}
address := destination.Addr().Unmap()
for _, member := range *members {
if member.prefix.Contains(address) {
return member.conn
}
}
return nil
}
func (p *UDPEgressPool) ReceiveEgress(buffer []byte) (int, netip.AddrPort, error) {
p.access.Lock()
receiveDone := p.receiveDone
p.access.Unlock()
if receiveDone == nil {
return 0, netip.AddrPort{}, net.ErrClosed
}
select {
case <-receiveDone:
return 0, netip.AddrPort{}, net.ErrClosed
default:
}
select {
case packet := <-p.packetChan:
copied := copy(buffer, packet.buffer.Bytes())
packet.buffer.Release()
return copied, packet.source, nil
case <-receiveDone:
return 0, netip.AddrPort{}, net.ErrClosed
}
}
func (p *UDPEgressPool) rebuildLocked() {
if p.port == 0 {
return
}
specs := make(map[udpEgressSpec]struct{})
if !p.isExempt() {
for _, networkInterface := range p.interfaceFinder.Interfaces() {
if networkInterface.Flags&net.FlagUp == 0 ||
networkInterface.Flags&net.FlagLoopback != 0 ||
networkInterface.Flags&net.FlagPointToPoint != 0 ||
networkInterface.Flags&net.FlagBroadcast == 0 ||
networkInterface.Index == p.anchorInterfaceIndex ||
networkInterface.Name == p.excludeInterface {
continue
}
for _, prefix := range networkInterface.Addresses {
if !prefix.Addr().IsGlobalUnicast() {
continue
}
if p.network == "udp4" && !prefix.Addr().Is4() {
continue
}
if p.network == "udp6" && prefix.Addr().Is4() {
continue
}
specs[udpEgressSpec{
interfaceIndex: networkInterface.Index,
interfaceName: networkInterface.Name,
prefix: prefix,
}] = struct{}{}
}
}
}
for spec, member := range p.members {
_, loaded := specs[spec]
if loaded {
continue
}
delete(p.members, spec)
member.conn.Close()
}
for spec := range specs {
_, loaded := p.members[spec]
if loaded {
continue
}
memberConn, err := p.listenMember(spec)
if err != nil {
p.logger.Warn(E.Cause(err, "listen egress member on ", spec.interfaceName, " (", spec.prefix.Addr(), ")"))
continue
}
member := &udpEgressMember{
prefix: spec.prefix.Masked(),
conn: memberConn,
}
p.members[spec] = member
p.memberReaders.Add(1)
go p.readMember(member, p.receiveDone)
}
members := make([]*udpEgressMember, 0, len(p.members))
for _, member := range p.members {
members = append(members, member)
}
slices.SortFunc(members, func(firstMember, secondMember *udpEgressMember) int {
return secondMember.prefix.Bits() - firstMember.prefix.Bits()
})
if len(members) == 0 {
p.state.Store(nil)
} else {
p.state.Store(&members)
}
}
func (p *UDPEgressPool) listenMember(spec udpEgressSpec) (*net.UDPConn, error) {
var listenConfig net.ListenConfig
if runtime.GOOS == "darwin" || runtime.GOOS == "ios" {
listenConfig.Control = control.ReuseAddrOnly()
}
listenConfig.Control = control.Append(listenConfig.Control, control.DisableUDPNetReset())
listenConfig.Control = control.Append(listenConfig.Control, control.BindToInterface(p.interfaceFinder, spec.interfaceName, spec.interfaceIndex))
listenConfig.Control = control.Append(listenConfig.Control, p.control)
var network string
if spec.prefix.Addr().Is4() {
network = "udp4"
} else {
network = "udp6"
}
packetConn, err := listenConfig.ListenPacket(context.Background(), network, netip.AddrPortFrom(spec.prefix.Addr(), p.port).String())
if err != nil {
return nil, err
}
return packetConn.(*net.UDPConn), nil
}
func (p *UDPEgressPool) readMember(member *udpEgressMember, doneChan <-chan struct{}) {
defer p.memberReaders.Done()
for {
buffer := buf.NewSize(udpEgressBufferSize)
dataLength, source, err := member.conn.ReadFromUDPAddrPort(buffer.FreeBytes())
if err != nil {
buffer.Release()
return
}
buffer.Extend(dataLength)
select {
case p.packetChan <- udpEgressPacket{buffer: buffer, source: source}:
case <-doneChan:
buffer.Release()
return
default:
buffer.Release()
}
}
}

View file

@ -1,124 +0,0 @@
package tun
import (
"net"
"net/netip"
"sync"
"time"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
)
type UDPEgressConn struct {
anchor *net.UDPConn
pool *UDPEgressPool
packetChan chan udpEgressConnPacket
doneChan chan struct{}
closeOnce sync.Once
readWait sync.WaitGroup
}
type udpEgressConnPacket struct {
buffer *buf.Buffer
source netip.AddrPort
err error
}
func NewUDPEgressConn(anchor *net.UDPConn, pool *UDPEgressPool) *UDPEgressConn {
conn := &UDPEgressConn{
anchor: anchor,
pool: pool,
packetChan: make(chan udpEgressConnPacket, 64),
doneChan: make(chan struct{}),
}
conn.readWait.Add(2)
go conn.read(anchor.ReadFromUDPAddrPort)
go conn.read(pool.ReceiveEgress)
return conn
}
func (c *UDPEgressConn) read(readPacket func([]byte) (int, netip.AddrPort, error)) {
defer c.readWait.Done()
for {
buffer := buf.NewSize(udpEgressBufferSize)
dataLength, source, err := readPacket(buffer.FreeBytes())
if err != nil {
buffer.Release()
if E.IsClosed(err) {
return
}
select {
case c.packetChan <- udpEgressConnPacket{err: err}:
case <-c.doneChan:
return
}
continue
}
buffer.Extend(dataLength)
select {
case c.packetChan <- udpEgressConnPacket{buffer: buffer, source: source}:
case <-c.doneChan:
buffer.Release()
return
}
}
}
func (c *UDPEgressConn) ReadFromUDPAddrPort(buffer []byte) (int, netip.AddrPort, error) {
select {
case packet := <-c.packetChan:
if packet.err != nil {
return 0, netip.AddrPort{}, packet.err
}
copied := copy(buffer, packet.buffer.Bytes())
packet.buffer.Release()
return copied, packet.source, nil
case <-c.doneChan:
return 0, netip.AddrPort{}, net.ErrClosed
}
}
func (c *UDPEgressConn) WriteToUDPAddrPort(buffer []byte, destination netip.AddrPort) (int, error) {
memberConn := c.pool.LookupEgress(destination)
if memberConn != nil {
return memberConn.WriteToUDPAddrPort(buffer, destination)
}
return c.anchor.WriteToUDPAddrPort(buffer, destination)
}
func (c *UDPEgressConn) LocalAddr() net.Addr {
return c.anchor.LocalAddr()
}
func (c *UDPEgressConn) SetDeadline(t time.Time) error {
return c.anchor.SetDeadline(t)
}
func (c *UDPEgressConn) SetReadDeadline(t time.Time) error {
return c.anchor.SetReadDeadline(t)
}
func (c *UDPEgressConn) SetWriteDeadline(t time.Time) error {
return c.anchor.SetWriteDeadline(t)
}
func (c *UDPEgressConn) Close() error {
c.closeOnce.Do(func() {
close(c.doneChan)
c.anchor.Close()
c.pool.Close()
c.readWait.Wait()
for {
select {
case packet := <-c.packetChan:
if packet.buffer != nil {
packet.buffer.Release()
}
default:
return
}
}
})
return nil
}

View file

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

View file

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