Files
sing-box-extended-mirror/route/conn.go
T

432 lines
14 KiB
Go
Raw Normal View History

2024-11-20 11:32:02 +08:00
package route
import (
"context"
"io"
"net"
"net/netip"
2025-03-15 08:09:04 +08:00
"os"
2025-04-25 16:27:56 +08:00
"strings"
2024-11-27 18:08:19 +08:00
"sync"
2024-11-20 11:32:02 +08:00
"sync/atomic"
2024-11-24 14:45:40 +08:00
"time"
2024-11-20 11:32:02 +08:00
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/common/sniff"
2025-01-26 09:01:00 +08:00
"github.com/sagernet/sing-box/common/tlsfragment"
2024-11-24 14:45:40 +08:00
C "github.com/sagernet/sing-box/constant"
2024-11-20 11:32:02 +08:00
"github.com/sagernet/sing/common"
2025-03-15 08:09:04 +08:00
"github.com/sagernet/sing/common/buf"
2024-11-20 11:32:02 +08:00
"github.com/sagernet/sing/common/bufio"
2024-11-24 14:45:40 +08:00
"github.com/sagernet/sing/common/canceler"
2024-11-20 11:32:02 +08:00
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
2024-11-27 18:08:19 +08:00
"github.com/sagernet/sing/common/x/list"
2024-11-20 11:32:02 +08:00
)
var _ adapter.ConnectionManager = (*ConnectionManager)(nil)
type ConnectionManager struct {
2024-11-27 18:08:19 +08:00
logger logger.ContextLogger
access sync.Mutex
connections list.List[io.Closer]
2024-11-20 11:32:02 +08:00
}
func NewConnectionManager(logger logger.ContextLogger) *ConnectionManager {
return &ConnectionManager{
2024-11-27 18:08:19 +08:00
logger: logger,
2024-11-20 11:32:02 +08:00
}
}
func (m *ConnectionManager) Start(stage adapter.StartStage) error {
2024-11-27 18:08:19 +08:00
return nil
2024-11-20 11:32:02 +08:00
}
2026-02-26 12:54:00 +08:00
func (m *ConnectionManager) Count() int {
return m.connections.Len()
}
func (m *ConnectionManager) CloseAll() {
2024-11-27 18:08:19 +08:00
m.access.Lock()
2026-02-26 12:54:00 +08:00
var closers []io.Closer
for element := m.connections.Front(); element != nil; {
nextElement := element.Next()
closers = append(closers, element.Value)
m.connections.Remove(element)
element = nextElement
2024-11-27 18:08:19 +08:00
}
2026-02-26 12:54:00 +08:00
m.access.Unlock()
for _, closer := range closers {
common.Close(closer)
}
}
func (m *ConnectionManager) Close() error {
m.CloseAll()
2024-11-27 18:08:19 +08:00
return nil
2024-11-20 11:32:02 +08:00
}
2026-02-26 12:54:00 +08:00
func (m *ConnectionManager) TrackConn(conn net.Conn) net.Conn {
m.access.Lock()
element := m.connections.PushBack(conn)
m.access.Unlock()
return &trackedConn{
Conn: conn,
manager: m,
element: element,
}
}
func (m *ConnectionManager) TrackPacketConn(conn net.PacketConn) net.PacketConn {
m.access.Lock()
element := m.connections.PushBack(conn)
m.access.Unlock()
return &trackedPacketConn{
PacketConn: conn,
manager: m,
element: element,
}
}
2024-11-20 11:32:02 +08:00
func (m *ConnectionManager) NewConnection(ctx context.Context, this N.Dialer, conn net.Conn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) {
ctx = adapter.WithContext(ctx, &metadata)
var (
remoteConn net.Conn
err error
)
2024-12-15 00:45:41 +08:00
if len(metadata.DestinationAddresses) > 0 || metadata.Destination.IsIP() {
2024-11-20 11:32:02 +08:00
remoteConn, err = dialer.DialSerialNetwork(ctx, this, N.NetworkTCP, metadata.Destination, metadata.DestinationAddresses, metadata.NetworkStrategy, metadata.NetworkType, metadata.FallbackNetworkType, metadata.FallbackDelay)
} else {
remoteConn, err = this.DialContext(ctx, N.NetworkTCP, metadata.Destination)
}
if err != nil {
2025-04-25 16:27:56 +08:00
var remoteString string
if len(metadata.DestinationAddresses) > 0 {
remoteString = "[" + strings.Join(common.Map(metadata.DestinationAddresses, netip.Addr.String), ",") + "]"
} else {
remoteString = metadata.Destination.String()
}
var dialerString string
if outbound, isOutbound := this.(adapter.Outbound); isOutbound {
dialerString = " using outbound/" + outbound.Type() + "[" + outbound.Tag() + "]"
}
err = E.Cause(err, "open connection to ", remoteString, dialerString)
2024-11-20 11:32:02 +08:00
N.CloseOnHandshakeFailure(conn, onClose, err)
2024-11-27 18:08:19 +08:00
m.logger.ErrorContext(ctx, err)
2024-11-20 11:32:02 +08:00
return
}
err = N.ReportConnHandshakeSuccess(conn, remoteConn)
if err != nil {
2024-11-27 18:08:19 +08:00
err = E.Cause(err, "report handshake success")
2024-11-20 11:32:02 +08:00
remoteConn.Close()
N.CloseOnHandshakeFailure(conn, onClose, err)
2024-11-27 18:08:19 +08:00
m.logger.ErrorContext(ctx, err)
2024-11-20 11:32:02 +08:00
return
}
2025-06-12 08:58:07 +08:00
if metadata.TLSFragment || metadata.TLSRecordFragment {
remoteConn = tf.NewConn(remoteConn, ctx, metadata.TLSFragment, metadata.TLSRecordFragment, metadata.TLSFragmentFallbackDelay)
2025-01-26 09:01:00 +08:00
}
serverFirst := sniff.Skip(&metadata)
2024-11-20 11:32:02 +08:00
var done atomic.Bool
if m.kickWriteHandshake(ctx, conn, remoteConn, serverFirst, false, &done, onClose) {
2026-02-03 02:51:05 +08:00
return
}
if m.kickWriteHandshake(ctx, remoteConn, conn, serverFirst, true, &done, onClose) {
2026-02-03 02:51:05 +08:00
return
}
2024-11-20 11:32:02 +08:00
go m.connectionCopy(ctx, conn, remoteConn, false, &done, onClose)
go m.connectionCopy(ctx, remoteConn, conn, true, &done, onClose)
}
func (m *ConnectionManager) NewPacketConnection(ctx context.Context, this N.Dialer, conn N.PacketConn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) {
ctx = adapter.WithContext(ctx, &metadata)
var (
remotePacketConn net.PacketConn
remoteConn net.Conn
destinationAddress netip.Addr
err error
)
if metadata.UDPConnect {
2024-12-15 00:45:41 +08:00
parallelDialer, isParallelDialer := this.(dialer.ParallelInterfaceDialer)
2024-11-20 11:32:02 +08:00
if len(metadata.DestinationAddresses) > 0 {
2024-12-15 00:45:41 +08:00
if isParallelDialer {
2024-11-20 11:32:02 +08:00
remoteConn, err = dialer.DialSerialNetwork(ctx, parallelDialer, N.NetworkUDP, metadata.Destination, metadata.DestinationAddresses, metadata.NetworkStrategy, metadata.NetworkType, metadata.FallbackNetworkType, metadata.FallbackDelay)
} else {
remoteConn, err = N.DialSerial(ctx, this, N.NetworkUDP, metadata.Destination, metadata.DestinationAddresses)
}
2024-12-15 00:45:41 +08:00
} else if metadata.Destination.IsIP() {
if isParallelDialer {
remoteConn, err = dialer.DialSerialNetwork(ctx, parallelDialer, N.NetworkUDP, metadata.Destination, metadata.DestinationAddresses, metadata.NetworkStrategy, metadata.NetworkType, metadata.FallbackNetworkType, metadata.FallbackDelay)
} else {
remoteConn, err = this.DialContext(ctx, N.NetworkUDP, metadata.Destination)
}
2024-11-20 11:32:02 +08:00
} else {
remoteConn, err = this.DialContext(ctx, N.NetworkUDP, metadata.Destination)
}
if err != nil {
2025-04-25 16:27:56 +08:00
var remoteString string
if len(metadata.DestinationAddresses) > 0 {
remoteString = "[" + strings.Join(common.Map(metadata.DestinationAddresses, netip.Addr.String), ",") + "]"
} else {
remoteString = metadata.Destination.String()
}
var dialerString string
if outbound, isOutbound := this.(adapter.Outbound); isOutbound {
dialerString = " using outbound/" + outbound.Type() + "[" + outbound.Tag() + "]"
}
err = E.Cause(err, "open packet connection to ", remoteString, dialerString)
2024-11-20 11:32:02 +08:00
N.CloseOnHandshakeFailure(conn, onClose, err)
2025-04-25 16:27:56 +08:00
m.logger.ErrorContext(ctx, err)
2024-11-20 11:32:02 +08:00
return
}
remotePacketConn = bufio.NewUnbindPacketConn(remoteConn)
connRemoteAddr := M.AddrFromNet(remoteConn.RemoteAddr())
if connRemoteAddr != metadata.Destination.Addr {
destinationAddress = connRemoteAddr
}
} else {
if len(metadata.DestinationAddresses) > 0 {
remotePacketConn, destinationAddress, err = dialer.ListenSerialNetworkPacket(ctx, this, metadata.Destination, metadata.DestinationAddresses, metadata.NetworkStrategy, metadata.NetworkType, metadata.FallbackNetworkType, metadata.FallbackDelay)
2026-03-02 11:30:06 +08:00
} else if packetDialer, withDestination := this.(dialer.PacketDialerWithDestination); withDestination {
remotePacketConn, destinationAddress, err = packetDialer.ListenPacketWithDestination(ctx, metadata.Destination)
2024-11-20 11:32:02 +08:00
} else {
remotePacketConn, err = this.ListenPacket(ctx, metadata.Destination)
}
if err != nil {
2025-04-25 16:27:56 +08:00
var dialerString string
if outbound, isOutbound := this.(adapter.Outbound); isOutbound {
dialerString = " using outbound/" + outbound.Type() + "[" + outbound.Tag() + "]"
}
err = E.Cause(err, "listen packet connection using ", dialerString)
2024-11-20 11:32:02 +08:00
N.CloseOnHandshakeFailure(conn, onClose, err)
2025-04-25 16:27:56 +08:00
m.logger.ErrorContext(ctx, err)
2024-11-20 11:32:02 +08:00
return
}
}
err = N.ReportPacketConnHandshakeSuccess(conn, remotePacketConn)
if err != nil {
conn.Close()
remotePacketConn.Close()
m.logger.ErrorContext(ctx, "report handshake success: ", err)
return
}
if destinationAddress.IsValid() {
var originDestination M.Socksaddr
if metadata.RouteOriginalDestination.IsValid() {
originDestination = metadata.RouteOriginalDestination
} else {
originDestination = metadata.Destination
}
2025-03-11 14:12:59 +08:00
if natConn, loaded := common.Cast[bufio.NATPacketConn](conn); loaded {
natConn.UpdateDestination(destinationAddress)
2026-03-02 11:30:06 +08:00
} else {
destination := M.SocksaddrFrom(destinationAddress, metadata.Destination.Port)
if metadata.Destination != destination {
if metadata.UDPDisableDomainUnmapping {
remotePacketConn = bufio.NewUnidirectionalNATPacketConn(bufio.NewPacketConn(remotePacketConn), destination, originDestination)
} else {
remotePacketConn = bufio.NewNATPacketConn(bufio.NewPacketConn(remotePacketConn), destination, originDestination)
}
} else if metadata.RouteOriginalDestination.IsValid() && metadata.RouteOriginalDestination != metadata.Destination {
remotePacketConn = bufio.NewDestinationNATPacketConn(bufio.NewPacketConn(remotePacketConn), metadata.Destination, metadata.RouteOriginalDestination)
2024-11-20 11:32:02 +08:00
}
}
2025-02-02 23:17:31 +08:00
} else if metadata.RouteOriginalDestination.IsValid() && metadata.RouteOriginalDestination != metadata.Destination {
2025-03-09 15:20:55 +08:00
remotePacketConn = bufio.NewDestinationNATPacketConn(bufio.NewPacketConn(remotePacketConn), metadata.Destination, metadata.RouteOriginalDestination)
2024-11-20 11:32:02 +08:00
}
2024-11-24 14:45:40 +08:00
var udpTimeout time.Duration
if metadata.UDPTimeout > 0 {
udpTimeout = metadata.UDPTimeout
} else {
protocol := metadata.Protocol
if protocol == "" {
protocol = C.PortProtocols[metadata.Destination.Port]
}
if protocol != "" {
udpTimeout = C.ProtocolTimeouts[protocol]
}
}
if udpTimeout > 0 {
ctx, conn = canceler.NewPacketConn(ctx, conn, udpTimeout)
}
2024-11-20 11:32:02 +08:00
destination := bufio.NewPacketConn(remotePacketConn)
var done atomic.Bool
go m.packetConnectionCopy(ctx, conn, destination, false, &done, onClose)
go m.packetConnectionCopy(ctx, destination, conn, true, &done, onClose)
}
2025-03-15 08:09:04 +08:00
func (m *ConnectionManager) connectionCopy(ctx context.Context, source net.Conn, destination net.Conn, direction bool, done *atomic.Bool, onClose N.CloseHandlerFunc) {
2026-02-03 02:51:05 +08:00
_, err := bufio.CopyWithIncreateBuffer(destination, source, bufio.DefaultIncreaseBufferAfter, bufio.DefaultBatchSize)
2024-11-27 18:08:19 +08:00
if err != nil {
2025-03-15 08:09:04 +08:00
common.Close(source, destination)
2024-11-27 18:08:19 +08:00
} else if duplexDst, isDuplex := destination.(N.WriteCloser); isDuplex {
err = duplexDst.CloseWrite()
2024-11-20 11:32:02 +08:00
if err != nil {
2025-03-15 08:09:04 +08:00
common.Close(source, destination)
2024-11-20 11:32:02 +08:00
}
2024-11-27 18:08:19 +08:00
} else {
2025-03-15 08:09:04 +08:00
destination.Close()
2024-11-20 11:32:02 +08:00
}
2024-11-27 18:08:19 +08:00
if done.Swap(true) {
2026-02-26 12:54:00 +08:00
if onClose != nil {
onClose(err)
}
2025-03-15 08:09:04 +08:00
common.Close(source, destination)
2024-11-27 18:08:19 +08:00
}
if !direction {
if err == nil {
m.logger.DebugContext(ctx, "connection upload finished")
} else if !E.IsClosedOrCanceled(err) {
m.logger.ErrorContext(ctx, "connection upload closed: ", err)
} else {
m.logger.TraceContext(ctx, "connection upload closed")
}
} else {
if err == nil {
m.logger.DebugContext(ctx, "connection download finished")
} else if !E.IsClosedOrCanceled(err) {
m.logger.ErrorContext(ctx, "connection download closed: ", err)
} else {
m.logger.TraceContext(ctx, "connection download closed")
}
2024-11-20 11:32:02 +08:00
}
2024-11-27 18:08:19 +08:00
}
func (m *ConnectionManager) kickWriteHandshake(ctx context.Context, source net.Conn, destination net.Conn, serverFirst bool, direction bool, done *atomic.Bool, onClose N.CloseHandlerFunc) bool {
2026-02-03 02:51:05 +08:00
if !N.NeedHandshakeForWrite(destination) {
return false
2025-03-15 08:09:04 +08:00
}
2025-09-14 17:28:43 +08:00
var (
err error
2026-02-03 02:51:05 +08:00
wrotePayload bool
2025-09-14 17:28:43 +08:00
)
if serverFirst {
2026-02-03 02:51:05 +08:00
_ = destination.SetWriteDeadline(time.Now().Add(C.ReadPayloadTimeout))
_, err = destination.Write(nil)
2026-02-03 02:51:05 +08:00
_ = destination.SetWriteDeadline(time.Time{})
} else {
var cachedBuffer *buf.Buffer
sourceReader, readCounters := N.UnwrapCountReader(source, nil)
destinationWriter, writeCounters := N.UnwrapCountWriter(destination, nil)
if cachedReader, ok := sourceReader.(N.CachedReader); ok {
cachedBuffer = cachedReader.ReadCached()
}
if cachedBuffer != nil {
wrotePayload = true
dataLen := cachedBuffer.Len()
_, err = destinationWriter.Write(cachedBuffer.Bytes())
cachedBuffer.Release()
if err == nil {
for _, counter := range readCounters {
counter(int64(dataLen))
}
for _, counter := range writeCounters {
counter(int64(dataLen))
}
}
} else {
_ = destination.SetWriteDeadline(time.Now().Add(C.ReadPayloadTimeout))
_, err = destinationWriter.Write(nil)
_ = destination.SetWriteDeadline(time.Time{})
}
2025-09-09 19:20:15 +08:00
}
2026-02-03 02:51:05 +08:00
if err == nil {
return false
2025-03-15 08:09:04 +08:00
}
2026-02-03 02:51:05 +08:00
if !wrotePayload && (E.IsMulti(err, os.ErrInvalid, context.DeadlineExceeded, io.EOF) || E.IsTimeout(err)) {
return false
}
if !done.Swap(true) {
2026-02-26 12:54:00 +08:00
if onClose != nil {
onClose(err)
}
2026-02-03 02:51:05 +08:00
}
common.Close(source, destination)
if !direction {
m.logger.ErrorContext(ctx, "connection upload handshake: ", err)
} else {
m.logger.ErrorContext(ctx, "connection download handshake: ", err)
}
return true
2025-03-15 08:09:04 +08:00
}
2024-11-27 18:08:19 +08:00
func (m *ConnectionManager) packetConnectionCopy(ctx context.Context, source N.PacketReader, destination N.PacketWriter, direction bool, done *atomic.Bool, onClose N.CloseHandlerFunc) {
_, err := bufio.CopyPacket(destination, source)
2024-11-20 11:32:02 +08:00
if !direction {
2025-03-09 15:20:55 +08:00
if err == nil {
m.logger.DebugContext(ctx, "packet upload finished")
} else if E.IsClosedOrCanceled(err) {
2024-11-20 11:32:02 +08:00
m.logger.TraceContext(ctx, "packet upload closed")
} else {
m.logger.DebugContext(ctx, "packet upload closed: ", err)
}
} else {
2025-03-09 15:20:55 +08:00
if err == nil {
m.logger.DebugContext(ctx, "packet download finished")
} else if E.IsClosedOrCanceled(err) {
2024-11-20 11:32:02 +08:00
m.logger.TraceContext(ctx, "packet download closed")
} else {
m.logger.DebugContext(ctx, "packet download closed: ", err)
}
}
if !done.Swap(true) {
2026-02-26 12:54:00 +08:00
if onClose != nil {
onClose(err)
}
2024-11-20 11:32:02 +08:00
}
common.Close(source, destination)
}
2026-02-26 12:54:00 +08:00
type trackedConn struct {
net.Conn
manager *ConnectionManager
element *list.Element[io.Closer]
}
func (c *trackedConn) Close() error {
c.manager.access.Lock()
c.manager.connections.Remove(c.element)
c.manager.access.Unlock()
return c.Conn.Close()
}
func (c *trackedConn) Upstream() any {
return c.Conn
}
func (c *trackedConn) ReaderReplaceable() bool {
return true
}
func (c *trackedConn) WriterReplaceable() bool {
return true
}
type trackedPacketConn struct {
net.PacketConn
manager *ConnectionManager
element *list.Element[io.Closer]
}
func (c *trackedPacketConn) Close() error {
c.manager.access.Lock()
c.manager.connections.Remove(c.element)
c.manager.access.Unlock()
return c.PacketConn.Close()
}
func (c *trackedPacketConn) Upstream() any {
return bufio.NewPacketConn(c.PacketConn)
}
func (c *trackedPacketConn) ReaderReplaceable() bool {
return true
}
func (c *trackedPacketConn) WriterReplaceable() bool {
return true
}