diff --git a/proxy/shadowsocks_2022/inbound.go b/proxy/shadowsocks_2022/inbound.go index ca07fa3ae..9883e00e8 100644 --- a/proxy/shadowsocks_2022/inbound.go +++ b/proxy/shadowsocks_2022/inbound.go @@ -2,7 +2,6 @@ package shadowsocks_2022 import ( "context" - "io" "time" "github.com/xtls/xray-core/common" @@ -13,9 +12,6 @@ import ( "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/session" - "github.com/xtls/xray-core/common/signal" - "github.com/xtls/xray-core/common/task" - "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" @@ -101,35 +97,29 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout return errors.New("unable to set read deadline").Base(err) } + // 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4 + headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize + headerBuf := make([]byte, headerLen) + n, err := conn.Read(headerBuf) + if err != nil || n < headerLen { + ResetTCPConn(conn) + return errors.New("failed to read complete handshake header") + } + var salt [32]byte + copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength]) saltSlice := salt[:i.method.KeySaltLength] - if _, err := io.ReadFull(conn, saltSlice); err != nil { - return err - } + fixedChunk := headerBuf[i.method.KeySaltLength:] - if !i.saltFilter.Check(salt) { - return ErrSaltNotUnique - } - - sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength) - aead, err := i.method.NewAEAD(sessionKey) + reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter) if err != nil { + ResetTCPConn(conn) return err } - reader := NewStreamReader(conn, aead) - - reqHeader, err := ReadClientRequestHeader(conn, reader) - if err != nil { - return err - } - conn.SetReadDeadline(time.Time{}) dest := reqHeader.Destination - writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil) - if err != nil { - return err - } + writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice) ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ From: conn.RemoteAddr(), @@ -152,35 +142,11 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout } } - sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level)) - ctx, cancel := context.WithCancel(ctx) - timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) - ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) - - requestDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) - return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) - } - - responseDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) - return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) - } - - responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) - return task.Run(ctx, requestDone, responseDoneAndCloseWriter) + return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link) } func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { - udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() - defer func() { - udpConns.Range(func(key uint64, entry *udpConnEntry) bool { - entry.timer.SetTimeout(0) - return true - }) - }() - - reader := buf.NewReader(conn) + reader := buf.NewPacketReader(conn) for { mb, err := reader.ReadMultiBuffer() if err != nil { @@ -190,75 +156,30 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch for _, b := range mb { decoded, err := i.udpCodec.DecodePacket(b.Bytes()) - if err != nil { - b.Release() + b.Release() + if err != nil || decoded.HeaderType != HeaderTypeClient { continue } - entry, ok := udpConns.Load(decoded.SessionID) - if !ok { - sessCtx, cancel := context.WithCancel(ctx) - sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ - From: conn.RemoteAddr(), - To: decoded.Destination, - Status: log.AccessAccepted, - Email: i.user.Email, - }) - - link, err := dispatcher.Dispatch(sessCtx, decoded.Destination) - if err != nil { - cancel() - b.Release() - continue - } - - newEntry := &udpConnEntry{ - link: link, - cancel: cancel, - } - sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level)) - newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { - udpConns.Delete(decoded.SessionID) - common.Interrupt(link.Reader) - common.Interrupt(link.Writer) - cancel() - }, sessionPolicy.Timeouts.ConnectionIdle) - - actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry) - if loaded { - // Another goroutine/packet beat us to storing, terminate our redundant link - newEntry.timer.SetTimeout(0) - entry = actual - } else { - entry = newEntry - go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) { - defer func() { - cEntry.timer.SetTimeout(0) - }() - for { - resMb, err := cEntry.link.Reader.ReadMultiBuffer() - if err != nil { - return - } - cEntry.timer.Update() - for _, rb := range resMb { - encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes()) - rb.Release() - if err != nil { - continue - } - _, _ = conn.Write(encPacket) - } - } - }(decoded.SessionID, decoded.Destination, entry) + sessionItem := i.udpCodec.GetSession(decoded.SessionID) + if sessionItem.User == nil { + sessionItem.Lock() + if sessionItem.User == nil { + sessionItem.User = i.user } + sessionItem.Unlock() + } + link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) { + return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload) + }) + if err != nil { + continue } - entry.timer.Update() payloadBuf := buf.New() payloadBuf.Write(decoded.Payload) - b.Release() - _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf}) + payloadBuf.UDP = &decoded.Destination + _ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf}) } } } diff --git a/proxy/shadowsocks_2022/inbound_multi.go b/proxy/shadowsocks_2022/inbound_multi.go index f46bf72a7..746bd27f0 100644 --- a/proxy/shadowsocks_2022/inbound_multi.go +++ b/proxy/shadowsocks_2022/inbound_multi.go @@ -4,7 +4,6 @@ import ( "context" "crypto/cipher" "encoding/binary" - "io" "strconv" "strings" "sync" @@ -19,8 +18,6 @@ import ( "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/session" - "github.com/xtls/xray-core/common/signal" - "github.com/xtls/xray-core/common/task" "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/common/uuid" "github.com/xtls/xray-core/core" @@ -207,64 +204,46 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat return errors.New("unable to set read deadline").Base(err) } - // 1. Read Request Salt (16 or 32 bytes) + // 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4 + headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + headerBuf := make([]byte, headerLen) + n, err := conn.Read(headerBuf) + if err != nil || n < headerLen { + ResetTCPConn(conn) + return errors.New("failed to read complete handshake header") + } + var salt [32]byte + copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength]) saltSlice := salt[:i.method.KeySaltLength] - if _, err := io.ReadFull(conn, saltSlice); err != nil { - return err - } + eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize] + fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:] - if !i.saltFilter.Check(salt) { - return ErrSaltNotUnique - } - - // 2. Read Extended Identity Header (16 bytes) - var eih [AESBlockSize]byte - if _, err := io.ReadFull(conn, eih[:]); err != nil { - return err - } - - // Decrypt EIH with IdentitySubKey derived from masterPSK and salt - identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength) - block, err := i.method.NewBlock(identitySubkey) + decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih) if err != nil { + ResetTCPConn(conn) return err } - var decryptedHash [AESBlockSize]byte - block.Decrypt(decryptedHash[:], eih[:]) - // Lookup user user, ok := i.usersByHash.Load(decryptedHash) - if !ok || user == nil { + if !ok { + ResetTCPConn(conn) return ErrInvalidRequest } userPSK := user.Account.(*MemoryAccount).Key - // 3. Derive Session Subkey using matched user's PSK - sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength) - aead, err := i.method.NewAEAD(sessionKey) + reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter) if err != nil { + ResetTCPConn(conn) return err } - reader := NewStreamReader(conn, aead) - - // 4 & 5. Read Client Request Header - reqHeader, err := ReadClientRequestHeader(conn, reader) - if err != nil { - return err - } - conn.SetReadDeadline(time.Time{}) dest := reqHeader.Destination - // 6. Send Server Response Handshake - writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil) - if err != nil { - return err - } + writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice) - // 7. Dispatch Connection to Xray routing with matched User + // Dispatch Connection to Xray routing with matched User inbound := session.InboundFromContext(ctx) inbound.User = user @@ -289,35 +268,11 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat } } - sessionPolicy = i.policyManager.ForLevel(user.Level) - ctx, cancel := context.WithCancel(ctx) - timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) - ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) - - requestDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) - return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) - } - - responseDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) - return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) - } - - responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) - return task.Run(ctx, requestDone, responseDoneAndCloseWriter) + return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link) } func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { - udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() - defer func() { - udpConns.Range(func(key uint64, entry *udpConnEntry) bool { - entry.timer.SetTimeout(0) - return true - }) - }() - - reader := buf.NewReader(conn) + reader := buf.NewPacketReader(conn) for { mb, err := reader.ReadMultiBuffer() if err != nil { @@ -341,164 +296,61 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, sessionID := binary.BigEndian.Uint64(rawHeader[:8]) packetID := binary.BigEndian.Uint64(rawHeader[8:16]) - // Replay protection & session lookup sessionItem := i.udpSessions.GetOrCreate(sessionID) - sessionItem.Lock() - if !sessionItem.Window.Check(packetID) { - sessionItem.Unlock() + if !sessionItem.CheckPacketID(packetID) { b.Release() continue } var userPSK []byte var currentUser *protocol.MemoryUser - if sessionItem.User != nil { - currentUser = sessionItem.User - userPSK = sessionItem.UserPSK - sessionItem.Unlock() - } else { - sessionItem.Unlock() + sessionItem.Lock() + currentUser = sessionItem.User + userPSK = sessionItem.UserPSK + sessionItem.Unlock() + + if currentUser == nil { // Decrypt EIH - var decryptedHash [16]byte - i.udpMasterCipher.Decrypt(decryptedHash[:], packetBytes[16:32]) - for k := 0; k < 16; k++ { - decryptedHash[k] ^= rawHeader[k] - } + decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32]) user, ok := i.usersByHash.Load(decryptedHash) - if !ok || user == nil { + if !ok { b.Release() continue } currentUser = user userPSK = user.Account.(*MemoryAccount).Key - - sessionItem.Lock() - sessionItem.User = user - sessionItem.UserPSK = userPSK - sessionItem.Unlock() } - // Decrypt Body (with AEAD caching per session) - bodyAead := sessionItem.GetRemoteCipher() - if bodyAead == nil { - bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength) - var err error - bodyAead, err = i.method.NewAEAD(bodyKey) - if err != nil { - b.Release() - continue - } - sessionItem.SetRemoteCipher(bodyAead) - } - - bodyNonce := rawHeader[4:16] - bodyCipher := packetBytes[32:] - bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil) + decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:]) b.Release() - if err != nil || len(bodyPlain) < 1+8+2 { - continue - } - - sessionItem.Lock() - sessionItem.Window.Add(packetID) - sessionItem.Unlock() - - if bodyPlain[0] != HeaderTypeClient { - continue - } - epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) - diff := time.Now().Unix() - int64(epoch) - if diff < -30 || diff > 30 { - continue - } - - paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11])) - offset := 11 + paddingLen - if len(bodyPlain) < offset { - continue - } - - dest, addrLen, err := parseAddressPort(bodyPlain[offset:]) if err != nil { continue } - payload := bodyPlain[offset+addrLen:] + sessionItem.Lock() + if sessionItem.User == nil { + sessionItem.User = currentUser + sessionItem.UserPSK = userPSK + } + sessionItem.Unlock() - entry, ok := udpConns.Load(sessionID) - if !ok { - sessCtx, cancel := context.WithCancel(ctx) - inbound := session.InboundFromContext(sessCtx) - inbound.User = currentUser - - sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ - From: conn.RemoteAddr(), - To: dest, - Status: log.AccessAccepted, - Email: currentUser.Email, - }) - - link, err := dispatcher.Dispatch(sessCtx, dest) - if err != nil { - cancel() - continue - } - - newEntry := &udpConnEntry{ - link: link, - cancel: cancel, - } - sessionPolicy := i.policyManager.ForLevel(currentUser.Level) - newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { - udpConns.Delete(sessionID) - common.Interrupt(link.Reader) - common.Interrupt(link.Writer) - cancel() - }, sessionPolicy.Timeouts.ConnectionIdle) - - actual, loaded := udpConns.LoadOrStore(sessionID, newEntry) - if loaded { - newEntry.timer.SetTimeout(0) - entry = actual - } else { - entry = newEntry - go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) { - defer func() { - cEntry.timer.SetTimeout(0) - }() - for { - resMb, err := cEntry.link.Reader.ReadMultiBuffer() - if err != nil { - return - } - cEntry.timer.Update() - for _, rb := range resMb { - encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes()) - rb.Release() - if err != nil { - continue - } - _, _ = conn.Write(encPacket) - } - } - }(sessionID, userPSK, dest, entry) - } + link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) { + return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload) + }) + if err != nil { + continue } - entry.timer.Update() pBuf := buf.New() - pBuf.Write(payload) - _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) + pBuf.Write(decoded.Payload) + pBuf.UDP = &decoded.Destination + _ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) } } } func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { - sessionItem := i.udpSessions.GetOrCreate(clientSessionID) - if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil { - return nil, err - } - return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload) + return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload) } diff --git a/proxy/shadowsocks_2022/inbound_relay.go b/proxy/shadowsocks_2022/inbound_relay.go index 8f29db12e..914cb91f0 100644 --- a/proxy/shadowsocks_2022/inbound_relay.go +++ b/proxy/shadowsocks_2022/inbound_relay.go @@ -4,7 +4,6 @@ import ( "context" "crypto/cipher" "encoding/binary" - "io" "strconv" "time" @@ -15,9 +14,6 @@ import ( "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/session" - "github.com/xtls/xray-core/common/signal" - "github.com/xtls/xray-core/common/task" - "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/common/uuid" "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/features/policy" @@ -35,18 +31,17 @@ type relayDest struct { destination net.Destination email string level uint32 - key []byte blockCipher cipher.Block } type RelayInbound struct { - networks []net.Network - method *CipherMethod - relayPSK []byte - relayBlock cipher.Block - destinations map[[AESBlockSize]byte]*relayDest - rawDestinations []*RelayDestination - policyManager policy.Manager + networks []net.Network + method *CipherMethod + relayPSK []byte + relayBlock cipher.Block + destinations map[[AESBlockSize]byte]*relayDest + udpSessions *UDPSessionManager + policyManager policy.Manager } func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) { @@ -78,13 +73,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou v := core.MustFromContext(ctx) i := &RelayInbound{ - networks: networks, - method: method, - relayPSK: relayPSK, - relayBlock: relayBlock, - destinations: make(map[[AESBlockSize]byte]*relayDest), - rawDestinations: config.Destinations, - policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + networks: networks, + method: method, + relayPSK: relayPSK, + relayBlock: relayBlock, + destinations: make(map[[AESBlockSize]byte]*relayDest), + udpSessions: NewUDPSessionManager(500 * time.Second), + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), } for idx, d := range config.Destinations { @@ -108,7 +103,6 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)), email: d.Email, level: uint32(d.Level), - key: destKey, blockCipher: destBlock, } } @@ -139,28 +133,36 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher return errors.New("unable to set read deadline").Base(err) } - // Read Salt + Outer EIH + // Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4 needed := i.method.KeySaltLength + AESBlockSize - var headerBuf [48]byte - headerSlice := headerBuf[:needed] - if _, err := io.ReadFull(conn, headerSlice); err != nil { - return err - } - - salt := headerSlice[:i.method.KeySaltLength] - eih := headerSlice[i.method.KeySaltLength:] - - identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength) - block, err := i.method.NewBlock(identitySubkey) + requestHeader := buf.New() + n, err := requestHeader.ReadFrom(conn) if err != nil { + requestHeader.Release() + ResetTCPConn(conn) return err } + if int(n) < needed { + requestHeader.Release() + ResetTCPConn(conn) + return ErrInvalidRequest + } - var decryptedHash [AESBlockSize]byte - block.Decrypt(decryptedHash[:], eih) + headerSlice := requestHeader.Bytes() + salt := headerSlice[:i.method.KeySaltLength] + eih := headerSlice[i.method.KeySaltLength:needed] + + decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih) + if err != nil { + requestHeader.Release() + ResetTCPConn(conn) + return err + } targetDest, ok := i.destinations[decryptedHash] if !ok { + requestHeader.Release() + ResetTCPConn(conn) return ErrInvalidRequest } conn.SetReadDeadline(time.Time{}) @@ -182,45 +184,26 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher link, err := dispatcher.Dispatch(ctx, targetDest.destination) if err != nil { + requestHeader.Release() return err } - // Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH - saltBuf := buf.New() - saltBuf.Write(salt) - if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil { + // Unwrap outer EIH: send client salt and remaining handshake bytes to next hop + // in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3). + var saltCopy [32]byte + copy(saltCopy[:i.method.KeySaltLength], salt) + copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength]) + requestHeader.Advance(AESBlockSize) + + if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil { return err } - sessionPolicy = i.policyManager.ForLevel(targetDest.level) - ctx, cancel := context.WithCancel(ctx) - timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) - ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) - - requestDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) - return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer)) - } - - responseDone := func() error { - defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) - return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer)) - } - - responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) - return task.Run(ctx, requestDone, responseDoneAndCloseWriter) + return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link) } func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { - udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() - defer func() { - udpConns.Range(func(key uint64, entry *udpConnEntry) bool { - entry.timer.SetTimeout(0) - return true - }) - }() - - reader := buf.NewReader(conn) + reader := buf.NewPacketReader(conn) for { mb, err := reader.ReadMultiBuffer() if err != nil { @@ -238,11 +221,7 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis var packetHeader [AESBlockSize]byte i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize]) - var eiHeader [AESBlockSize]byte - i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize]) - for idx := 0; idx < AESBlockSize; idx++ { - eiHeader[idx] ^= packetHeader[idx] - } + eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize]) targetDest, ok := i.destinations[eiHeader] if !ok { @@ -263,68 +242,24 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis dest := targetDest.destination dest.Network = net.Network_UDP - entry, ok := udpConns.Load(sessionID) - if !ok { - sessCtx, cancel := context.WithCancel(ctx) - inbound := session.InboundFromContext(sessCtx) - inbound.User = &protocol.MemoryUser{ - Email: targetDest.email, - Level: targetDest.level, - } - - sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ - From: conn.RemoteAddr(), - To: dest, - Status: log.AccessAccepted, - Email: targetDest.email, - }) - - link, err := dispatcher.Dispatch(sessCtx, dest) - if err != nil { - cancel() - b.Release() - continue - } - - newEntry := &udpConnEntry{ - link: link, - cancel: cancel, - } - sessionPolicy := i.policyManager.ForLevel(targetDest.level) - newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { - udpConns.Delete(sessionID) - common.Interrupt(link.Reader) - common.Interrupt(link.Writer) - cancel() - }, sessionPolicy.Timeouts.ConnectionIdle) - - actual, loaded := udpConns.LoadOrStore(sessionID, newEntry) - if loaded { - newEntry.timer.SetTimeout(0) - entry = actual - } else { - entry = newEntry - go func(cEntry *udpConnEntry) { - defer func() { - cEntry.timer.SetTimeout(0) - }() - for { - resMb, err := cEntry.link.Reader.ReadMultiBuffer() - if err != nil { - return - } - cEntry.timer.Update() - for _, rb := range resMb { - _, _ = conn.Write(rb.Bytes()) - rb.Release() - } - } - }(entry) + sessionItem := i.udpSessions.GetOrCreate(sessionID) + if sessionItem.User == nil { + sessionItem.Lock() + if sessionItem.User == nil { + sessionItem.User = &protocol.MemoryUser{ + Email: targetDest.email, + Level: targetDest.level, + } } + sessionItem.Unlock() + } + link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil) + if err != nil { + b.Release() + continue } - entry.timer.Update() - _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b}) + _ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b}) } } } diff --git a/proxy/shadowsocks_2022/kdf.go b/proxy/shadowsocks_2022/kdf.go index 3ebc18379..0f70c49df 100644 --- a/proxy/shadowsocks_2022/kdf.go +++ b/proxy/shadowsocks_2022/kdf.go @@ -61,3 +61,14 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte { copy(out[:], h[:AESBlockSize]) return out } + +func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) { + identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength) + block, err := method.NewBlock(identitySubkey) + if err != nil { + return [AESBlockSize]byte{}, err + } + var decryptedHash [AESBlockSize]byte + block.Decrypt(decryptedHash[:], eih) + return decryptedHash, nil +} diff --git a/proxy/shadowsocks_2022/outbound.go b/proxy/shadowsocks_2022/outbound.go index ff1be9f06..67e0acd28 100644 --- a/proxy/shadowsocks_2022/outbound.go +++ b/proxy/shadowsocks_2022/outbound.go @@ -4,7 +4,6 @@ import ( "context" "crypto/rand" "io" - "time" "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/buf" @@ -46,6 +45,10 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) { return nil, errors.New("invalid key: ", config.Key).Base(err) } + if method.IsChaCha && len(pskList) > 1 { + return nil, errors.New("multi-key is not supported for chacha20-poly1305") + } + finalPSK := pskList[len(pskList)-1] udpCodec, err := NewUDPPacketCodec(method, pskList) if err != nil { @@ -126,18 +129,30 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int requestDone := func() error { defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) - bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn)) - bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil) + + var initialPayload []byte + var firstBuf *buf.Buffer + var remainingMB buf.MultiBuffer + if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok { + if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() { + remainingMB, firstBuf = buf.SplitFirst(mb) + initialPayload = firstBuf.Bytes() + } + } + + bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload) + if firstBuf != nil { + firstBuf.Release() + } if err != nil { + buf.ReleaseMulti(remainingMB) return errors.New("failed to write request").Base(err) } - if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { - return errors.New("failed to write A request payload").Base(err) - } - - if err := bufferedWriter.SetBuffered(false); err != nil { - return err + if !remainingMB.IsEmpty() { + if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil { + return err + } } return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)) @@ -163,13 +178,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int } if network == net.Network_UDP { + session, err := o.udpCodec.NewClientSession() + if err != nil { + return errors.New("failed to create client udp session").Base(err) + } + requestDone := func() error { defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) writer := &UDPWriter{ Writer: conn, Destination: destination, - Codec: o.udpCodec, + Session: session, } if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil { @@ -182,8 +202,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) reader := &UDPReader{ - Reader: conn, - Codec: o.udpCodec, + Reader: conn, + Session: session, } if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil { diff --git a/proxy/shadowsocks_2022/packet.go b/proxy/shadowsocks_2022/packet.go index fc906aeb8..f81754575 100644 --- a/proxy/shadowsocks_2022/packet.go +++ b/proxy/shadowsocks_2022/packet.go @@ -16,16 +16,13 @@ import ( ) type UDPCodec struct { - method *CipherMethod - pskList [][]byte - psk []byte - blockCipher cipher.Block - blockCiphers []cipher.Block - chachaCipher cipher.AEAD - clientBodyCipher cipher.AEAD - clientSessionID uint64 - nextPacketID atomic.Uint64 - sessions *UDPSessionManager + method *CipherMethod + pskList [][]byte + psk []byte + blockCipher cipher.Block + blockCiphers []cipher.Block + chachaCipher cipher.AEAD + sessions *UDPSessionManager } type ( @@ -51,13 +48,16 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) { } func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) { + if method.IsChaCha && len(pskList) > 1 { + return nil, errors.New("multi-key is not supported for chacha20-poly1305") + } finalPSK := pskList[len(pskList)-1] c, err := newUDPCodec(method, finalPSK) if err != nil { return nil, err } c.pskList = pskList - if len(pskList) > 1 && !method.IsChaCha { + if len(pskList) > 1 { c.blockCiphers = make([]cipher.Block, len(pskList)) for i, psk := range pskList { c.blockCiphers[i], err = method.NewBlock(psk) @@ -66,20 +66,6 @@ func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error } } } - - var sessID [8]byte - if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil { - return nil, err - } - c.clientSessionID = binary.BigEndian.Uint64(sessID[:]) - - if !method.IsChaCha { - clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], method.KeySaltLength) - c.clientBodyCipher, err = method.NewAEAD(clientBodyKey) - if err != nil { - return nil, err - } - } return c, nil } @@ -92,138 +78,37 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur return c, nil } -func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) { - packetID := c.nextPacketID.Add(1) - sessID := c.clientSessionID +func (c *UDPCodec) Sessions() *UDPSessionManager { + return c.sessions +} - // Padding determination (e.g. DNS port 53 disguise) - var paddingLen int - if dest.Port == 53 && len(payload) < MaxPaddingLength { - paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1 +func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession { + if c.sessions == nil { + return nil } - - addrPortLen := AddrPortLength(dest) - - if c.method.IsChaCha { - // ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B) - totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize - if totalLen > buf.Size { - return nil, ErrPacketTooLarge - } - - outBuf := buf.New() - - var nonce [PacketNonceSize]byte - if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil { - outBuf.Release() - return nil, err - } - outBuf.Write(nonce[:]) - - var hdr [16 + 1 + 8 + 2]byte - binary.BigEndian.PutUint64(hdr[0:8], sessID) - binary.BigEndian.PutUint64(hdr[8:16], packetID) - hdr[16] = HeaderTypeClient - binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix())) - binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen)) - outBuf.Write(hdr[:]) - if paddingLen > 0 { - outBuf.Write(zeroPadding[:paddingLen]) - } - - if err := WriteAddressPort(outBuf, dest); err != nil { - outBuf.Release() - return nil, err - } - outBuf.Write(payload) - - plainBytes := outBuf.Bytes()[PacketNonceSize:] - outBuf.Extend(int32(c.chachaCipher.Overhead())) - c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil) - return outBuf, nil - } - - // AES mode: - var sessBytes [8]byte - binary.BigEndian.PutUint64(sessBytes[:], sessID) - - var rawHeader [16]byte - copy(rawHeader[:8], sessBytes[:]) - binary.BigEndian.PutUint64(rawHeader[8:16], packetID) - - eihCount := 0 - if len(c.pskList) > 1 { - eihCount = len(c.pskList) - 1 - } - - totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize - if totalLen > buf.Size { - return nil, ErrPacketTooLarge - } - - outBuf := buf.New() - - if len(c.pskList) > 1 { - // Multi-user / Relay mode: - // 1. Header (16B) encrypted with first hop's block cipher - var encryptedHeader [16]byte - c.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:]) - outBuf.Write(encryptedHeader[:]) - - // 2. Multi-hop EIHs for intermediate hops - for i := 0; i < len(c.pskList)-1; i++ { - nextPSK := c.pskList[i+1] - pskHash := DeriveUserPSKHash(nextPSK) - var eihPlain [16]byte - for k := 0; k < 16; k++ { - eihPlain[k] = pskHash[k] ^ rawHeader[k] - } - var encryptedEIH [16]byte - c.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:]) - outBuf.Write(encryptedEIH[:]) - } - } else { - // Single-user mode: - var encryptedHeader [16]byte - c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) - outBuf.Write(encryptedHeader[:]) - } - - bodyAead := c.clientBodyCipher - - var hdr [1 + 8 + 2]byte - hdr[0] = HeaderTypeClient - binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix())) - binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen)) - outBuf.Write(hdr[:]) - if paddingLen > 0 { - outBuf.Write(zeroPadding[:paddingLen]) - } - - if err := WriteAddressPort(outBuf, dest); err != nil { - outBuf.Release() - return nil, err - } - outBuf.Write(payload) - - headerOffset := 16 + eihCount*16 - plainBytes := outBuf.Bytes()[headerOffset:] - bodyNonce := rawHeader[4:16] - outBuf.Extend(int32(bodyAead.Overhead())) - bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil) - return outBuf, nil + return c.sessions.GetOrCreate(sessionID) } type DecodedUDPPacket struct { - SessionID uint64 - PacketID uint64 - HeaderType byte - Timestamp uint64 - Destination net.Destination - Payload []byte + SessionID uint64 + PacketID uint64 + HeaderType byte + Timestamp uint64 + ClientSessionID uint64 + Destination net.Destination + Payload []byte } -func parseAddressPort(data []byte) (net.Destination, int, error) { +func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte { + var decryptedHash [AESBlockSize]byte + block.Decrypt(decryptedHash[:], eih) + for k := 0; k < AESBlockSize; k++ { + decryptedHash[k] ^= rawHeader[k] + } + return decryptedHash +} + +func ParseAddressPort(data []byte) (net.Destination, int, error) { if len(data) < 1 { return net.Destination{}, 0, ErrPacketTooShort } @@ -264,6 +149,9 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU } headerType := bodyPlain[0] + if headerType != HeaderTypeClient && headerType != HeaderTypeServer { + return DecodedUDPPacket{}, ErrBadHeaderType + } epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) if diff > 30 { @@ -271,11 +159,13 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU } offset := 9 + var clientSessionID uint64 if headerType == HeaderTypeServer { if len(bodyPlain) < offset+8+2 { return DecodedUDPPacket{}, ErrPacketTooShort } - offset += 8 // skip clientSessionID + clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8]) + offset += 8 } paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2])) @@ -286,19 +176,20 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU } offset += paddingLen - dest, addrLen, err := parseAddressPort(bodyPlain[offset:]) + dest, addrLen, err := ParseAddressPort(bodyPlain[offset:]) if err != nil { return DecodedUDPPacket{}, err } payload := bodyPlain[offset+addrLen:] return DecodedUDPPacket{ - SessionID: sessionID, - PacketID: packetID, - HeaderType: headerType, - Timestamp: epoch, - Destination: dest, - Payload: payload, + SessionID: sessionID, + PacketID: packetID, + HeaderType: headerType, + Timestamp: epoch, + ClientSessionID: clientSessionID, + Destination: dest, + Payload: payload, }, nil } @@ -313,7 +204,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { } nonce := data[:PacketNonceSize] ciphertext := data[PacketNonceSize:] - plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil) + plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil) if err != nil { return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err) } @@ -324,17 +215,22 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { sessionID := binary.BigEndian.Uint64(plain[:8]) packetID := binary.BigEndian.Uint64(plain[8:16]) - if c.sessions != nil { - sessionItem := c.sessions.GetOrCreate(sessionID) - sessionItem.Lock() - if !sessionItem.Window.CheckAndAdd(packetID) { - sessionItem.Unlock() - return DecodedUDPPacket{}, ErrPacketIdNotUnique - } - sessionItem.Unlock() + sessionItem := c.sessions.GetOrCreate(sessionID) + if !sessionItem.CheckPacketID(packetID) { + return DecodedUDPPacket{}, ErrPacketIdNotUnique } - return parsePlainUDPPacket(sessionID, packetID, plain[16:]) + decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:]) + if err != nil { + return DecodedUDPPacket{}, err + } + + if decoded.HeaderType != HeaderTypeClient { + return DecodedUDPPacket{}, ErrBadHeaderType + } + + sessionItem.AddPacketID(packetID) + return decoded, nil } // AES mode @@ -343,54 +239,52 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { sessionID := binary.BigEndian.Uint64(rawHeader[:8]) packetID := binary.BigEndian.Uint64(rawHeader[8:16]) - var bodyAead cipher.AEAD - var sessionItem *ServerUDPSession + sessionItem := c.sessions.GetOrCreate(sessionID) + if !sessionItem.CheckPacketID(packetID) { + return DecodedUDPPacket{}, ErrPacketIdNotUnique + } - if c.sessions != nil { - sessionItem = c.sessions.GetOrCreate(sessionID) - sessionItem.Lock() - if !sessionItem.Window.Check(packetID) { - sessionItem.Unlock() - return DecodedUDPPacket{}, ErrPacketIdNotUnique - } - sessionItem.Unlock() + return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:]) +} - bodyAead = sessionItem.GetRemoteCipher() - if bodyAead == nil { - bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) - var err error - bodyAead, err = c.method.NewAEAD(bodyKey) - if err != nil { - return DecodedUDPPacket{}, err - } - sessionItem.SetRemoteCipher(bodyAead) - } - } else { - bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) +func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) { + bodyAead := s.clientBodyCipher + isNewCipher := false + if bodyAead == nil { + bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength) var err error - bodyAead, err = c.method.NewAEAD(bodyKey) + bodyAead, err = method.NewAEAD(bodyKey) if err != nil { return DecodedUDPPacket{}, err } + isNewCipher = true } bodyNonce := rawHeader[4:16] - bodyCipher := data[16:] - bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil) + bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil) if err != nil { return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err) } - if sessionItem != nil { - sessionItem.Lock() - sessionItem.Window.Add(packetID) - sessionItem.Unlock() + decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain) + if err != nil { + return DecodedUDPPacket{}, err } - return parsePlainUDPPacket(sessionID, packetID, bodyPlain) + if decoded.HeaderType != HeaderTypeClient { + return DecodedUDPPacket{}, ErrBadHeaderType + } + + s.AddPacketID(packetID) + + if isNewCipher { + s.clientBodyCipher = bodyAead + } + + return decoded, nil } -func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error { +func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error { s.Lock() defer s.Unlock() if s.ServerSessionID != 0 { @@ -407,23 +301,29 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock c } } if method.IsChaCha { - s.ServerChaCha = chachaCipher - } else { - s.ServerBlockCipher = headerBlock - bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength) - bodyAead, err := method.NewAEAD(bodyKey) - if err != nil { - s.ServerSessionID = 0 - return err - } - s.ServerCipher = bodyAead + var err error + s.serverChaCha, err = method.NewUDPCipher(psk) + return err + } + + var err error + s.serverHeaderBlock, err = method.NewBlock(psk) + if err != nil { + s.ServerSessionID = 0 + return err + } + bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength) + s.serverBodyCipher, err = method.NewAEAD(bodyKey) + if err != nil { + s.ServerSessionID = 0 + return err } return nil } func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { serverSessionID := s.ServerSessionID - serverPacketID := s.ServerPacketID.Add(1) + serverPacketID := s.ServerPacketID.Add(1) - 1 if method.IsChaCha { var nonce [PacketNonceSize]byte @@ -448,7 +348,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio } plainBuf.Write(payload) - sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil) + sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil) res := make([]byte, PacketNonceSize+len(sealed)) copy(res[:PacketNonceSize], nonce[:]) copy(res[PacketNonceSize:], sealed) @@ -461,7 +361,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID) var encryptedHeader [16]byte - s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) + s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:]) bodyBuf := buf.New() defer bodyBuf.Release() @@ -479,7 +379,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio bodyBuf.Write(payload) bodyNonce := rawHeader[4:16] - sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil) + sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil) res := make([]byte, 16+len(sealedBody)) copy(res[:16], encryptedHeader[:]) @@ -488,17 +388,327 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio } func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { - sessionItem := c.sessions.GetOrCreate(clientSessionID) - if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil { + return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload) +} + +type serverSessionState struct { + sessionID uint64 + window *SlidingWindow + cipher cipher.AEAD + lastSeen atomic.Int64 +} + +func (st *serverSessionState) check(packetID uint64) bool { + if st.window == nil { + st.window = new(SlidingWindow) + } + return st.window.Check(packetID) +} + +func (st *serverSessionState) add(packetID uint64) { + if st.window == nil { + st.window = new(SlidingWindow) + } + st.window.Add(packetID) +} + +type ClientUDPSession struct { + codec *UDPCodec + clientSessionID uint64 + nextPacketID atomic.Uint64 + clientBodyCipher cipher.AEAD + current atomic.Pointer[serverSessionState] + old atomic.Pointer[serverSessionState] +} + +func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) { + var sessID [8]byte + if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil { return nil, err } - return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload) + clientSessionID := binary.BigEndian.Uint64(sessID[:]) + + var clientBodyCipher cipher.AEAD + var err error + if !c.method.IsChaCha { + finalPSK := c.psk + clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength) + clientBodyCipher, err = c.method.NewAEAD(clientBodyKey) + if err != nil { + return nil, err + } + } + + return &ClientUDPSession{ + codec: c, + clientSessionID: clientSessionID, + clientBodyCipher: clientBodyCipher, + }, nil +} + +func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) { + cur := s.current.Load() + if cur != nil && cur.sessionID == sessionID { + return cur, nil + } + + old := s.old.Load() + if old != nil && old.sessionID == sessionID { + if now-old.lastSeen.Load() > 60 { + s.old.CompareAndSwap(old, nil) + return nil, errors.New("old server session expired") + } + return old, nil + } + + // New server session: + // Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old. + if old != nil && now-old.lastSeen.Load() < 60 { + return nil, errors.New("newer server session rejected: old session is less than 1 minute old") + } + + var bodyAead cipher.AEAD + if !s.codec.method.IsChaCha { + var sessBytes [8]byte + binary.BigEndian.PutUint64(sessBytes[:], sessionID) + bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength) + var err error + bodyAead, err = s.codec.method.NewAEAD(bodyKey) + if err != nil { + return nil, err + } + } + + newState := &serverSessionState{ + sessionID: sessionID, + cipher: bodyAead, + } + newState.lastSeen.Store(now) + + if cur == nil { + s.current.CompareAndSwap(nil, newState) + return s.current.Load(), nil + } + + s.old.Store(cur) + s.current.Store(newState) + return newState, nil +} + +func (s *ClientUDPSession) ClientSessionID() uint64 { + return s.clientSessionID +} + +func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) { + packetID := s.nextPacketID.Add(1) - 1 + sessID := s.clientSessionID + + var paddingLen int + if dest.Port == 53 && len(payload) < MaxPaddingLength { + paddingLen = mrand.IntN(MaxPaddingLength) + 1 + } + + addrPortLen := AddrPortLength(dest) + + if s.codec.method.IsChaCha { + totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize + if totalLen > buf.Size { + return nil, ErrPacketTooLarge + } + + outBuf := buf.New() + + var nonce [PacketNonceSize]byte + if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(nonce[:]) + + var hdr [16 + 1 + 8 + 2]byte + binary.BigEndian.PutUint64(hdr[0:8], sessID) + binary.BigEndian.PutUint64(hdr[8:16], packetID) + hdr[16] = HeaderTypeClient + binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen)) + outBuf.Write(hdr[:]) + if paddingLen > 0 { + outBuf.Write(zeroPadding[:paddingLen]) + } + + if err := WriteAddressPort(outBuf, dest); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(payload) + + plainBytes := outBuf.Bytes()[PacketNonceSize:] + outBuf.Extend(int32(s.codec.chachaCipher.Overhead())) + s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil) + return outBuf, nil + } + + // AES mode + var sessBytes [8]byte + binary.BigEndian.PutUint64(sessBytes[:], sessID) + + var rawHeader [16]byte + copy(rawHeader[:8], sessBytes[:]) + binary.BigEndian.PutUint64(rawHeader[8:16], packetID) + + eihCount := 0 + if len(s.codec.pskList) > 1 { + eihCount = len(s.codec.pskList) - 1 + } + + totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize + if totalLen > buf.Size { + return nil, ErrPacketTooLarge + } + + outBuf := buf.New() + + if len(s.codec.pskList) > 1 { + var encryptedHeader [16]byte + s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:]) + outBuf.Write(encryptedHeader[:]) + + for i := 0; i < len(s.codec.pskList)-1; i++ { + nextPSK := s.codec.pskList[i+1] + pskHash := DeriveUserPSKHash(nextPSK) + var eihPlain [16]byte + for k := 0; k < 16; k++ { + eihPlain[k] = pskHash[k] ^ rawHeader[k] + } + var encryptedEIH [16]byte + s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:]) + outBuf.Write(encryptedEIH[:]) + } + } else { + var encryptedHeader [16]byte + s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) + outBuf.Write(encryptedHeader[:]) + } + + bodyAead := s.clientBodyCipher + + var hdr [1 + 8 + 2]byte + hdr[0] = HeaderTypeClient + binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen)) + outBuf.Write(hdr[:]) + if paddingLen > 0 { + outBuf.Write(zeroPadding[:paddingLen]) + } + + if err := WriteAddressPort(outBuf, dest); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(payload) + + headerOffset := 16 + eihCount*16 + plainBytes := outBuf.Bytes()[headerOffset:] + bodyNonce := rawHeader[4:16] + outBuf.Extend(int32(bodyAead.Overhead())) + bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil) + return outBuf, nil +} + +func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) { + if len(data) < PacketMinimalHeaderSize { + return DecodedUDPPacket{}, ErrPacketTooShort + } + + if s.codec.method.IsChaCha { + if len(data) < PacketNonceSize+AEADTagSize { + return DecodedUDPPacket{}, ErrPacketTooShort + } + nonce := data[:PacketNonceSize] + ciphertext := data[PacketNonceSize:] + plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil) + if err != nil { + return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err) + } + if len(plain) < 16+1+8+2 { + return DecodedUDPPacket{}, ErrPacketTooShort + } + + sessionID := binary.BigEndian.Uint64(plain[:8]) + packetID := binary.BigEndian.Uint64(plain[8:16]) + + now := time.Now().Unix() + st, err := s.getServerSession(sessionID, now) + if err != nil { + return DecodedUDPPacket{}, err + } + if !st.check(packetID) { + return DecodedUDPPacket{}, ErrPacketIdNotUnique + } + + decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:]) + if err != nil { + return DecodedUDPPacket{}, err + } + + if decoded.HeaderType != HeaderTypeServer { + return DecodedUDPPacket{}, ErrBadHeaderType + } + if decoded.ClientSessionID != s.clientSessionID { + return DecodedUDPPacket{}, errors.New("client session ID mismatch") + } + + st.add(packetID) + st.lastSeen.Store(now) + + return decoded, nil + } + + // AES mode + var rawHeader [16]byte + s.codec.blockCipher.Decrypt(rawHeader[:], data[:16]) + sessionID := binary.BigEndian.Uint64(rawHeader[:8]) + packetID := binary.BigEndian.Uint64(rawHeader[8:16]) + + now := time.Now().Unix() + st, err := s.getServerSession(sessionID, now) + if err != nil { + return DecodedUDPPacket{}, err + } + if !st.check(packetID) { + return DecodedUDPPacket{}, ErrPacketIdNotUnique + } + bodyAead := st.cipher + + bodyNonce := rawHeader[4:16] + bodyCipher := data[16:] + bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil) + if err != nil { + return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err) + } + + decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain) + if err != nil { + return DecodedUDPPacket{}, err + } + + if decoded.HeaderType != HeaderTypeServer { + return DecodedUDPPacket{}, ErrBadHeaderType + } + if decoded.ClientSessionID != s.clientSessionID { + return DecodedUDPPacket{}, errors.New("client session ID mismatch") + } + + st.add(packetID) + st.lastSeen.Store(now) + + return decoded, nil } type UDPWriter struct { Writer io.Writer Destination net.Destination - Codec *UDPPacketCodec + Session *ClientUDPSession } func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { @@ -512,7 +722,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { if b.UDP != nil { dest = *b.UDP } - pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes()) + pktBuf, err := w.Session.EncodePacket(dest, b.Bytes()) b.Release() if err != nil { buf.ReleaseMulti(mb) @@ -529,8 +739,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { } type UDPReader struct { - Reader io.Reader - Codec *UDPPacketCodec + Reader io.Reader + Session *ClientUDPSession } func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) { @@ -542,7 +752,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) { return nil, err } - decoded, err := r.Codec.DecodePacket(buffer.Bytes()) + decoded, err := r.Session.DecodePacket(buffer.Bytes()) if err != nil { buffer.Release() continue diff --git a/proxy/shadowsocks_2022/relay_test.go b/proxy/shadowsocks_2022/relay_test.go index c981e5d73..096852ca6 100644 --- a/proxy/shadowsocks_2022/relay_test.go +++ b/proxy/shadowsocks_2022/relay_test.go @@ -2,9 +2,11 @@ package shadowsocks_2022_test import ( "context" + "crypto/rand" "encoding/base64" "encoding/binary" "errors" + "io" gonet "net" "sync" "sync/atomic" @@ -269,3 +271,106 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error { } return nil } + +func TestRelayTCPHandshakeForwarding(t *testing.T) { + methods := []string{MethodAES128GCM, MethodAES256GCM} + for _, methodName := range methods { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + + relayKey := make([]byte, method.KeySaltLength) + destKey := make([]byte, method.KeySaltLength) + _, _ = io.ReadFull(rand.Reader, relayKey) + _, _ = io.ReadFull(rand.Reader, destKey) + + targetPort := uint32(54321) + relayConfig := &RelayServerConfig{ + Method: methodName, + Key: base64.StdEncoding.EncodeToString(relayKey), + Destinations: []*RelayDestination{ + { + Key: base64.StdEncoding.EncodeToString(destKey), + Address: net.NewIPOrDomain(net.LocalHostIP), + Port: targetPort, + Email: "test@xray.com", + }, + }, + } + + testCtx := newTestContext() + inbound, err := NewRelayServer(testCtx, relayConfig) + common.Must(err) + + targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort)) + + downstreamR, downstreamW := gonet.Pipe() + defer downstreamR.Close() + defer downstreamW.Close() + + disp := &dummyDispatcher{ + onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) { + inLink := &transport.Link{ + Reader: buf.NewReader(downstreamR), + Writer: &customWriter{ + write: func(mb buf.MultiBuffer) error { + defer buf.ReleaseMulti(mb) + for _, b := range mb { + if _, err := downstreamW.Write(b.Bytes()); err != nil { + return err + } + } + return nil + }, + }, + } + return inLink, nil + }, + } + + clientConn, relayConn := gonet.Pipe() + defer clientConn.Close() + defer relayConn.Close() + + go func() { + _ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp) + }() + + clientSalt := make([]byte, method.KeySaltLength) + _, _ = io.ReadFull(rand.Reader, clientSalt) + pskList := [][]byte{relayKey, destKey} + + go func() { + _, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload")) + if err != nil { + t.Errorf("WriteTCPRequest failed: %v", err) + } + }() + + // Downstream server must be able to read Salt + Fixed chunk in a single Read call! + headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize + headerBuf := make([]byte, headerLen) + n, err := downstreamR.Read(headerBuf) + if err != nil { + t.Fatalf("downstream failed to read handshake: %v", err) + } + if n < headerLen { + t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n) + } + + // Verify downstream can decode the fixed chunk and subsequent payload + sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + common.Must(err) + + reader := NewStreamReader(downstreamR, aead) + reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:]) + if err != nil { + t.Fatalf("downstream failed to parse client request header: %v", err) + } + if string(reqHeader.EarlyData) != "relay payload" { + t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData)) + } + }) + } +} diff --git a/proxy/shadowsocks_2022/replay.go b/proxy/shadowsocks_2022/replay.go index 7908371ba..241eb053c 100644 --- a/proxy/shadowsocks_2022/replay.go +++ b/proxy/shadowsocks_2022/replay.go @@ -6,8 +6,11 @@ import ( "sync/atomic" "time" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/signal" "github.com/xtls/xray-core/common/utils" + "github.com/xtls/xray-core/transport" ) const ( @@ -74,30 +77,42 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool { type ServerUDPSession struct { sync.Mutex - SessionID uint64 - RemoteCipher atomic.Pointer[cipher.AEAD] - Window SlidingWindow - User *protocol.MemoryUser - UserPSK []byte - LastActive atomic.Int64 // Unix timestamp in seconds + SessionID uint64 + Window *SlidingWindow + User *protocol.MemoryUser + UserPSK []byte + LastActive atomic.Int64 // Unix timestamp in seconds + + clientBodyCipher cipher.AEAD ServerSessionID uint64 ServerPacketID atomic.Uint64 - ServerCipher cipher.AEAD - ServerBlockCipher cipher.Block - ServerChaCha cipher.AEAD + serverBodyCipher cipher.AEAD + serverHeaderBlock cipher.Block + serverChaCha cipher.AEAD + + manager *UDPSessionManager + link atomic.Pointer[transport.Link] + timer *signal.ActivityTimer + currentConn atomic.Value // stores stat.Connection } -func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD { - ptr := s.RemoteCipher.Load() - if ptr == nil { - return nil +func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool { + s.Lock() + defer s.Unlock() + if s.Window == nil { + s.Window = new(SlidingWindow) } - return *ptr + return s.Window.Check(packetID) } -func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) { - s.RemoteCipher.Store(&c) +func (s *ServerUDPSession) AddPacketID(packetID uint64) { + s.Lock() + defer s.Unlock() + if s.Window == nil { + s.Window = new(SlidingWindow) + } + s.Window.Add(packetID) } type UDPSessionManager struct { @@ -122,6 +137,7 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession { s := &ServerUDPSession{ SessionID: sessionID, + manager: m, } s.LastActive.Store(now) @@ -148,6 +164,7 @@ func (m *UDPSessionManager) cleanup(now int64) { m.sessions.Range(func(k uint64, v *ServerUDPSession) bool { if now-v.LastActive.Load() > timeoutSec { m.sessions.Delete(k) + v.Close() } return true }) @@ -156,3 +173,11 @@ func (m *UDPSessionManager) cleanup(now int64) { func (m *UDPSessionManager) Delete(sessionID uint64) { m.sessions.Delete(sessionID) } + +func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + sessionItem := m.GetOrCreate(clientSessionID) + if err := sessionItem.EnsureServerState(method, psk); err != nil { + return nil, err + } + return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload) +} diff --git a/proxy/shadowsocks_2022/shadowsocks_2022.go b/proxy/shadowsocks_2022/shadowsocks_2022.go index 505f33c2d..ec8562e6a 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022.go @@ -2,18 +2,161 @@ package shadowsocks_2022 import ( "context" - "sync" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/log" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/session" "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/features/policy" + "github.com/xtls/xray-core/features/routing" + "github.com/xtls/xray-core/proxy" "github.com/xtls/xray-core/transport" + "github.com/xtls/xray-core/transport/internet/stat" ) -type udpConnEntry struct { - sync.Mutex - link *transport.Link - timer *signal.ActivityTimer - cancel context.CancelFunc +func (s *ServerUDPSession) UpdateConn(conn stat.Connection) { + if s.currentConn.Load() == nil { + s.currentConn.Store(conn) + } + if s.timer != nil { + s.timer.Update() + } +} + +func (s *ServerUDPSession) WriteToClient(b []byte) error { + connVal := s.currentConn.Load() + if connVal == nil { + return errors.New("client connection closed") + } + conn, ok := connVal.(stat.Connection) + if !ok || conn == nil { + return errors.New("client connection closed") + } + _, err := conn.Write(b) + return err +} + +func (s *ServerUDPSession) Close() { + if s.timer != nil { + s.timer.SetTimeout(0) + } + if link := s.link.Load(); link != nil { + common.Interrupt(link.Reader) + common.Interrupt(link.Writer) + } +} + +func (s *ServerUDPSession) EnsureLink( + ctx context.Context, + conn stat.Connection, + dest net.Destination, + dispatcher routing.Dispatcher, + policyManager policy.Manager, + responseEncoder func(dest net.Destination, payload []byte) ([]byte, error), +) (*transport.Link, error) { + s.UpdateConn(conn) + + if link := s.link.Load(); link != nil { + return link, nil + } + + s.Lock() + defer s.Unlock() + + if link := s.link.Load(); link != nil { + return link, nil + } + + sessCtx, cancel := context.WithCancel(ctx) + inbound := session.InboundFromContext(sessCtx) + if inbound != nil && s.User != nil { + inbound.User = s.User + } + var email string + var level uint32 + if s.User != nil { + email = s.User.Email + level = s.User.Level + } + sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: email, + }) + + link, err := dispatcher.Dispatch(sessCtx, dest) + if err != nil { + cancel() + return nil, err + } + + s.link.Store(link) + sessionPolicy := policyManager.ForLevel(level) + s.timer = signal.CancelAfterInactivity(sessCtx, func() { + if s.manager != nil { + s.manager.Delete(s.SessionID) + } + s.Close() + cancel() + }, sessionPolicy.Timeouts.ConnectionIdle) + + go handleUDPResponse(s, link, dest, responseEncoder) + return link, nil +} + +// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close +// when handshake or header validation fails. +func ResetTCPConn(conn net.Conn) { + rawConn, _, _ := proxy.UnwrapRawConn(conn) + if tcpConn, ok := rawConn.(*net.TCPConn); ok { + _ = tcpConn.SetLinger(0) + } +} + +func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) { + defer func() { + if s.timer != nil { + s.timer.SetTimeout(0) + } + }() + for { + resMb, err := link.Reader.ReadMultiBuffer() + if err != nil { + return + } + if s.timer != nil { + s.timer.Update() + } + for i, rb := range resMb { + b := rb.Bytes() + if encode != nil { + replyDest := fallbackDest + if rb.UDP != nil { + replyDest = *rb.UDP + } + encPacket, err := encode(replyDest, b) + rb.Release() + if err != nil { + continue + } + if err := s.WriteToClient(encPacket); err != nil { + buf.ReleaseMulti(resMb[i+1:]) + return + } + } else { + err := s.WriteToClient(b) + rb.Release() + if err != nil { + buf.ReleaseMulti(resMb[i+1:]) + return + } + } + } + } } const ( diff --git a/proxy/shadowsocks_2022/shadowsocks_2022_test.go b/proxy/shadowsocks_2022/shadowsocks_2022_test.go index bcec78d50..382696efb 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022_test.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022_test.go @@ -182,57 +182,48 @@ func TestTCPStream(t *testing.T) { common.Must(err) IncreaseNonce(reader.Nonce()) - vBuf := buf.New() - vBuf.Write(plainVar) - receivedDest, err = ReadAddressPort(vBuf) + dest, addrLen, err := ParseAddressPort(plainVar) common.Must(err) + receivedDest = net.TCPDestination(dest.Address, dest.Port) + plainVar = plainVar[addrLen:] + padLen := int(binary.BigEndian.Uint16(plainVar[:2])) + receivedPayload = plainVar[2+padLen:] - // Skip padding - var padBytes [2]byte - _, _ = vBuf.Read(padBytes[:]) - padLen := int(padBytes[0])<<8 | int(padBytes[1]) - vBuf.Advance(int32(padLen)) + // Server sends response stream with receivedPayload as first payload + writer := NewServerStreamWriter(serverConn, method, rawKey, salt) + pBuf := buf.New() + pBuf.Write(receivedPayload) + _ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) - receivedPayload = make([]byte, vBuf.Len()) - copy(receivedPayload, vBuf.Bytes()) - vBuf.Release() - - // Server sends response handshake - serverSalt := make([]byte, method.KeySaltLength) - _, _ = rand.Read(serverSalt) - respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength) - respAead, err := method.NewAEAD(respKey) - writer := NewStreamWriter(serverConn, respAead) - _, _ = serverConn.Write(serverSalt) - - fixedResp := make([]byte, 1+8+method.KeySaltLength+2) - fixedResp[0] = HeaderTypeServer - binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix())) - copy(fixedResp[9:9+method.KeySaltLength], salt) - binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0) - - fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil) - IncreaseNonce(writer.Nonce()) - _, _ = serverConn.Write(fixedChunk) - - // Echo stream data + // Read and echo additional stream data mb, err := reader.ReadMultiBuffer() common.Must(err) _ = writer.WriteMultiBuffer(mb) + _ = writer.Close() }() // Client goroutine go func() { defer wg.Done() - clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload) + clientSalt := make([]byte, method.KeySaltLength) + common.Must2(io.ReadFull(rand.Reader, clientSalt)) + writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload) common.Must(err) - reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt) + reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt) common.Must(err) + // The first ReadMultiBuffer drains initialPayload from reader cache + mbInit, err := reader.ReadMultiBuffer() + common.Must(err) + if !bytes.Equal(mbInit[0].Bytes(), testPayload) { + t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload) + } + buf.ReleaseMulti(mbInit) + // Send additional stream data streamData := []byte("stream chunk test") - _ = writer.WriteChunk(streamData) + _ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)}) mb, err := reader.ReadMultiBuffer() common.Must(err) @@ -277,7 +268,9 @@ func TestUDPCodec(t *testing.T) { serverCodec, err := NewUDPServerCodec(method, psk, time.Minute) common.Must(err) - pktBuf, err := clientCodec.EncodeClientPacket(dest, payload) + session, err := clientCodec.NewClientSession() + common.Must(err) + pktBuf, err := session.EncodePacket(dest, payload) common.Must(err) defer pktBuf.Release() @@ -423,3 +416,83 @@ func TestLargeStreamTransfer(t *testing.T) { t.Fatal("received data does not match sent data") } } + +func TestClientUDPSessionMultiDestination(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawKey := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawKey) + + clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey}) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute) + common.Must(err) + + session, err := clientCodec.NewClientSession() + common.Must(err) + + dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53)) + dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53)) + + payload1 := []byte("query-google-dns") + payload2 := []byte("query-cloudflare-dns") + + // Client sends to dest1 and dest2 using SAME session + pkt1, err := session.EncodePacket(dest1, payload1) + common.Must(err) + defer pkt1.Release() + pkt2, err := session.EncodePacket(dest2, payload2) + common.Must(err) + defer pkt2.Release() + + // Server decodes both + dec1, err := serverCodec.DecodePacket(pkt1.Bytes()) + common.Must(err) + dec2, err := serverCodec.DecodePacket(pkt2.Bytes()) + common.Must(err) + + if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() { + t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID) + } + if dec1.Destination.String() != dest1.String() { + t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination) + } + if dec2.Destination.String() != dest2.String() { + t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination) + } + if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) { + t.Fatal("payload mismatch") + } + + // Server replies to dest1 and dest2 + respPayload1 := []byte("reply-google-dns") + respPayload2 := []byte("reply-cloudflare-dns") + + respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1) + common.Must(err) + respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2) + common.Must(err) + + // Client decodes replies + clientDec1, err := session.DecodePacket(respPkt1) + common.Must(err) + if clientDec1.Destination.String() != dest1.String() { + t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination) + } + if !bytes.Equal(clientDec1.Payload, respPayload1) { + t.Fatal("reply payload 1 mismatch") + } + + clientDec2, err := session.DecodePacket(respPkt2) + common.Must(err) + if clientDec2.Destination.String() != dest2.String() { + t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination) + } + if !bytes.Equal(clientDec2.Payload, respPayload2) { + t.Fatal("reply payload 2 mismatch") + } + }) + } +} diff --git a/proxy/shadowsocks_2022/stream.go b/proxy/shadowsocks_2022/stream.go index 14afd66dc..acde16155 100644 --- a/proxy/shadowsocks_2022/stream.go +++ b/proxy/shadowsocks_2022/stream.go @@ -1,18 +1,25 @@ package shadowsocks_2022 import ( + "context" "crypto/cipher" "crypto/rand" "encoding/binary" "io" "math" mrand "math/rand/v2" + "sync" "time" + "github.com/xtls/xray-core/common/antireplay" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/common/task" + "github.com/xtls/xray-core/features/policy" + "github.com/xtls/xray-core/transport" ) var addrParser = protocol.NewAddressParser( @@ -38,15 +45,6 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error { return addrParser.WriteAddressPort(w, dest.Address, dest.Port) } -// ReadAddressPort reads a destination address and port in SOCKS5 format -func ReadAddressPort(r io.Reader) (net.Destination, error) { - addr, port, err := addrParser.ReadAddressPort(nil, r) - if err != nil { - return net.Destination{}, err - } - return net.TCPDestination(addr, port), nil -} - // AddrPortLength returns the serialized length of a destination in SOCKS5 format func AddrPortLength(dest net.Destination) int { switch dest.Address.Family() { @@ -243,13 +241,8 @@ type ClientRequestHeader struct { EarlyData []byte } -func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) { - var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte - if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil { - return nil, err - } - - plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil) +func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) { + plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil) if err != nil { return nil, errors.New("failed to decrypt client request header").Base(err) } @@ -278,7 +271,7 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque } else { varChunkCipher = make([]byte, needed) } - if _, err := io.ReadFull(conn, varChunkCipher); err != nil { + if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil { return nil, err } @@ -288,7 +281,7 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque } IncreaseNonce(reader.Nonce()) - dest, addrLen, err := parseAddressPort(plainVar) + dest, addrLen, err := ParseAddressPort(plainVar) if err != nil { return nil, err } @@ -307,8 +300,15 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque offset += paddingLen var earlyData []byte + var payloadLen int if len(plainVar) > offset { earlyData = plainVar[offset:] + payloadLen = len(earlyData) + } + + // SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0. + if paddingLen == 0 && payloadLen == 0 { + return nil, errors.New("request without payload and padding is not allowed") } return &ClientRequestHeader{ @@ -317,34 +317,6 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque }, nil } -// ClientHandshake writes the full client request header to w -func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) { - salt := make([]byte, method.KeySaltLength) - if _, err := io.ReadFull(rand.Reader, salt); err != nil { - return nil, nil, err - } - writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload) - if err != nil { - return nil, nil, err - } - return salt, writer.(*StreamWriter), nil -} - -// ClientVerifyServerResponse reads and verifies the server's handshake response -func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) { - reader, err := ReadTCPResponse(r, method, psk, clientSalt) - if err != nil { - return nil, nil, err - } - sr := reader.(*StreamReader) - var initialPayload []byte - if sr.cached > 0 { - initialPayload = make([]byte, sr.cached) - copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached]) - } - return sr, initialPayload, nil -} - // WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer. func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) { finalPSK := pskList[len(pskList)-1] @@ -356,7 +328,16 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n writer := NewStreamWriter(w, aead) - handshakeBuf := buf.New() + payloadLen := len(payload) + var paddingLen int + if payloadLen < MaxPaddingLength { + paddingLen = mrand.IntN(MaxPaddingLength) + 1 + } + addrPortLen := AddrPortLength(dest) + varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen + + totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize) + handshakeBuf := buf.NewWithSize(totalHandshakeLen) defer handshakeBuf.Release() handshakeBuf.Write(clientSalt) @@ -374,14 +355,6 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n handshakeBuf.Write(encryptedEIH[:]) } - payloadLen := len(payload) - var paddingLen int - if payloadLen < MaxPaddingLength { - paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1 - } - addrPortLen := AddrPortLength(dest) - varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen - var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte fixedHeaderPlaintext[0] = HeaderTypeClient binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix())) @@ -391,7 +364,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n IncreaseNonce(writer.nonce[:]) handshakeBuf.Write(fixedChunk) - varHeaderBuf := buf.New() + varHeaderBuf := buf.NewWithSize(int32(varHeaderLen)) defer varHeaderBuf.Release() if err := WriteAddressPort(varHeaderBuf, dest); err != nil { @@ -423,12 +396,21 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n // ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream. func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) { - var serverSalt [32]byte - serverSaltSlice := serverSalt[:method.KeySaltLength] - if _, err := io.ReadFull(r, serverSaltSlice); err != nil { - return nil, err + fixedPlainLen := 1 + 8 + method.KeySaltLength + 2 + chunkCipherLen := fixedPlainLen + AEADTagSize + headerLen := method.KeySaltLength + chunkCipherLen + + // Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4 + var headerBuf [128]byte + headerSlice := headerBuf[:headerLen] + n, err := r.Read(headerSlice) + if err != nil || n < headerLen { + return nil, errors.New("failed to read complete server response header") } + serverSaltSlice := headerSlice[:method.KeySaltLength] + chunkSlice := headerSlice[method.KeySaltLength:headerLen] + sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) aead, err := method.NewAEAD(sessionKey) if err != nil { @@ -437,14 +419,6 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [ reader := NewStreamReader(r, aead) - fixedPlainLen := 1 + 8 + method.KeySaltLength + 2 - chunkCipherLen := fixedPlainLen + AEADTagSize - var chunkBuf [64]byte - chunkSlice := chunkBuf[:chunkCipherLen] - if _, err := io.ReadFull(r, chunkSlice); err != nil { - return nil, err - } - decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil) if err != nil { return nil, errors.New("failed to decrypt server response header").Base(err) @@ -486,46 +460,190 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [ return reader, nil } -// WriteTCPResponse writes the server handshake response and returns a body writer for server stream. -func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) { +// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4. +type ServerStreamWriter struct { + mu sync.Mutex + w io.Writer + method *CipherMethod + psk []byte + clientSalt []byte + streamWriter *StreamWriter +} + +func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter { + return &ServerStreamWriter{ + w: w, + method: method, + psk: psk, + clientSalt: clientSalt, + } +} + +func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) { var serverSalt [32]byte - serverSaltSlice := serverSalt[:method.KeySaltLength] + serverSaltSlice := serverSalt[:s.method.KeySaltLength] if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil { return nil, err } - respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) - respAead, err := method.NewAEAD(respKey) + respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength) + respAead, err := s.method.NewAEAD(respKey) if err != nil { return nil, err } - writer := NewStreamWriter(w, respAead) + sw := NewStreamWriter(s.w, respAead) - respBuf := buf.New() - defer respBuf.Release() + totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize) + outBuf := buf.NewWithSize(totalHeaderLen) + defer outBuf.Release() - respBuf.Write(serverSaltSlice) + outBuf.Write(serverSaltSlice) var fixedRespPlain [1 + 8 + 32 + 2]byte - fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2] + fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2] fixedRespSlice[0] = HeaderTypeServer binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix())) - copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt) - binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload))) + copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt) + binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload))) - fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil) - IncreaseNonce(writer.nonce[:]) - respBuf.Write(fixedRespChunk) + fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil) + IncreaseNonce(sw.nonce[:]) + outBuf.Write(fixedRespChunk) - if len(initialPayload) > 0 { - initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil) - IncreaseNonce(writer.nonce[:]) - respBuf.Write(initialChunk) + if len(payload) > 0 { + payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil) + IncreaseNonce(sw.nonce[:]) + outBuf.Write(payloadChunk) } - if _, err := w.Write(respBuf.Bytes()); err != nil { + if _, err := s.w.Write(outBuf.Bytes()); err != nil { return nil, err } + return sw, nil +} - return writer, nil +func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + if mb.IsEmpty() { + return nil + } + + if s.streamWriter == nil { + s.mu.Lock() + if s.streamWriter == nil { + firstBuf := mb[0] + firstBytes := firstBuf.Bytes() + chunkSize := len(firstBytes) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + firstPayload := firstBytes[:chunkSize] + sw, err := s.sendHeaderWithFirstPayload(firstPayload) + if err != nil { + s.mu.Unlock() + buf.ReleaseMulti(mb) + return err + } + s.streamWriter = sw + + firstBuf.Advance(int32(chunkSize)) + if firstBuf.IsEmpty() { + firstBuf.Release() + mb = mb[1:] + } + } + s.mu.Unlock() + if len(mb) == 0 { + return nil + } + } + + return s.streamWriter.WriteMultiBuffer(mb) +} + +func (s *ServerStreamWriter) Write(p []byte) (int, error) { + n := len(p) + if s.streamWriter == nil { + s.mu.Lock() + if s.streamWriter == nil { + chunkSize := len(p) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + firstPayload := p[:chunkSize] + sw, err := s.sendHeaderWithFirstPayload(firstPayload) + if err != nil { + s.mu.Unlock() + return 0, err + } + s.streamWriter = sw + p = p[chunkSize:] + } + s.mu.Unlock() + if len(p) == 0 { + return n, nil + } + } + + _, err := s.streamWriter.Write(p) + return n, err +} + +func (s *ServerStreamWriter) Close() error { + if s.streamWriter == nil { + s.mu.Lock() + defer s.mu.Unlock() + if s.streamWriter == nil { + sw, err := s.sendHeaderWithFirstPayload(nil) + if err != nil { + return err + } + s.streamWriter = sw + } + } + return nil +} + +// InitServerStream decrypts the client request header, verifies the timestamp and replay filter, +// and returns a StreamReader for subsequent stream chunks. +func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) { + sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, nil, err + } + + reader := NewStreamReader(conn, aead) + + reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk) + if err != nil { + return nil, nil, err + } + _ = conn.SetReadDeadline(time.Time{}) + + if !saltFilter.Check(salt) { + return nil, nil, ErrSaltNotUnique + } + return reader, reqHeader, nil +} + +func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error { + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + if c, ok := writer.(io.Closer); ok { + defer c.Close() + } + return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + return task.Run(ctx, requestDone, responseDoneAndCloseWriter) }