Compare commits

..
3 Commits
Author SHA1 Message Date
Fangliding 7f23673023 Do not create Finalmask if not needed 2026-09-29 15:25:46 +08:00
Fangliding 3a9412c128 fmt 2026-09-29 15:25:45 +08:00
Fangliding 6140ff6844 Move PacketConnWrapper to common/net 2026-09-29 15:25:40 +08:00
25 changed files with 891 additions and 1427 deletions
+20
View File
@@ -0,0 +1,20 @@
package net
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
type PacketConnWrapper struct {
PacketConn
Dest Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() Addr {
return c.Dest
}
+4 -4
View File
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil { if statConn != nil {
counter = statConn.ReadCounter counter = statConn.ReadCounter
} }
if c, ok := iConn.(*internet.PacketConnWrapper); ok { if c, ok := iConn.(*net.PacketConnWrapper); ok {
isOverridden := false isOverridden := false
if UDPOverride.Address != nil || UDPOverride.Port != 0 { if UDPOverride.Address != nil || UDPOverride.Port != 0 {
isOverridden = true isOverridden = true
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
} }
type PacketReader struct { type PacketReader struct {
*internet.PacketConnWrapper *net.PacketConnWrapper
stats.Counter stats.Counter
Handler *Handler Handler *Handler
DefaultRule *FinalRule DefaultRule *FinalRule
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil { if statConn != nil {
counter = statConn.WriteCounter counter = statConn.WriteCounter
} }
if c, ok := iConn.(*internet.PacketConnWrapper); ok { if c, ok := iConn.(*net.PacketConnWrapper); ok {
// If DialDest is a domain, it will be resolved in dialer // If DialDest is a domain, it will be resolved in dialer
// check this behavior and add it to map // check this behavior and add it to map
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]() resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
} }
type PacketWriter struct { type PacketWriter struct {
*internet.PacketConnWrapper *net.PacketConnWrapper
stats.Counter stats.Counter
*Handler *Handler
DefaultRule *FinalRule DefaultRule *FinalRule
+1 -1
View File
@@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
} }
defer conn.Close() defer conn.Close()
uc := &wireguard.UDPConnClient{ uc := &wireguard.UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn, PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr), Dest: conn.RemoteAddr().(*net.UDPAddr),
} }
reader = uc reader = uc
+118 -38
View File
@@ -2,6 +2,7 @@ package shadowsocks_2022
import ( import (
"context" "context"
"io"
"time" "time"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
@@ -12,6 +13,9 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/core"
"github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/features/routing"
@@ -97,29 +101,35 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
return errors.New("unable to set read deadline").Base(err) 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 var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength] saltSlice := salt[:i.method.KeySaltLength]
fixedChunk := headerBuf[i.method.KeySaltLength:] if _, err := io.ReadFull(conn, saltSlice); err != nil {
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
if err != nil {
ResetTCPConn(conn)
return err return err
} }
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
return err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice) writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
if err != nil {
return err
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(), From: conn.RemoteAddr(),
@@ -136,17 +146,42 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
} }
if len(reqHeader.EarlyData) > 0 { if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData) earlyBuf := buf.New()
if err := link.Writer.WriteMultiBuffer(mb); err != nil { earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err return err
} }
} }
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link) 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)
} }
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn) 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)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -156,30 +191,75 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
for _, b := range mb { for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes()) decoded, err := i.udpCodec.DecodePacket(b.Bytes())
b.Release()
if err != nil || decoded.HeaderType != HeaderTypeClient {
continue
}
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 { if err != nil {
b.Release()
continue 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)
}
}
entry.timer.Update()
payloadBuf := buf.New() payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload) payloadBuf.Write(decoded.Payload)
payloadBuf.UDP = &decoded.Destination b.Release()
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf}) _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
} }
} }
} }
+201 -48
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"crypto/cipher" "crypto/cipher"
"encoding/binary" "encoding/binary"
"io"
"strconv" "strconv"
"strings" "strings"
"sync" "sync"
@@ -18,6 +19,8 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/utils"
"github.com/xtls/xray-core/common/uuid" "github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
@@ -204,46 +207,64 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
return errors.New("unable to set read deadline").Base(err) return errors.New("unable to set read deadline").Base(err)
} }
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4 // 1. Read Request Salt (16 or 32 bytes)
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 var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength] saltSlice := salt[:i.method.KeySaltLength]
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize] if _, err := io.ReadFull(conn, saltSlice); err != nil {
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
if err != nil {
ResetTCPConn(conn)
return err return err
} }
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)
if err != nil {
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih[:])
// Lookup user // Lookup user
user, ok := i.usersByHash.Load(decryptedHash) user, ok := i.usersByHash.Load(decryptedHash)
if !ok { if !ok || user == nil {
ResetTCPConn(conn)
return ErrInvalidRequest return ErrInvalidRequest
} }
userPSK := user.Account.(*MemoryAccount).Key userPSK := user.Account.(*MemoryAccount).Key
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter) // 3. Derive Session Subkey using matched user's PSK
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil { if err != nil {
ResetTCPConn(conn)
return err 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 dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice) // 6. Send Server Response Handshake
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
if err != nil {
return err
}
// Dispatch Connection to Xray routing with matched User // 7. Dispatch Connection to Xray routing with matched User
inbound := session.InboundFromContext(ctx) inbound := session.InboundFromContext(ctx)
inbound.User = user inbound.User = user
@@ -262,17 +283,42 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
} }
if len(reqHeader.EarlyData) > 0 { if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData) earlyBuf := buf.New()
if err := link.Writer.WriteMultiBuffer(mb); err != nil { earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err return err
} }
} }
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link) 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)
} }
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn) 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)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -296,61 +342,168 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
sessionID := binary.BigEndian.Uint64(rawHeader[:8]) sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16]) packetID := binary.BigEndian.Uint64(rawHeader[8:16])
// Replay protection & session lookup
sessionItem := i.udpSessions.GetOrCreate(sessionID) sessionItem := i.udpSessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) { sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
b.Release() b.Release()
continue continue
} }
var userPSK []byte var userPSK []byte
var currentUser *protocol.MemoryUser var currentUser *protocol.MemoryUser
sessionItem.Lock() if sessionItem.User != nil {
currentUser = sessionItem.User currentUser = sessionItem.User
userPSK = sessionItem.UserPSK userPSK = sessionItem.UserPSK
sessionItem.Unlock() sessionItem.Unlock()
} else {
if currentUser == nil { sessionItem.Unlock()
// Decrypt EIH // Decrypt EIH
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32]) identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
idBlock, err := i.method.NewBlock(identitySubkey)
if err != nil {
b.Release()
continue
}
var decryptedHash [16]byte
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash) user, ok := i.usersByHash.Load(decryptedHash)
if !ok { if !ok || user == nil {
b.Release() b.Release()
continue continue
} }
currentUser = user currentUser = user
userPSK = user.Account.(*MemoryAccount).Key userPSK = user.Account.(*MemoryAccount).Key
sessionItem.Lock()
sessionItem.User = user
sessionItem.UserPSK = userPSK
sessionItem.Unlock()
} }
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:]) // 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)
b.Release() b.Release()
if err != nil { if err != nil || len(bodyPlain) < 1+8+2 {
continue continue
} }
sessionItem.Lock() sessionItem.Lock()
if sessionItem.User == nil { sessionItem.Window.Add(packetID)
sessionItem.User = currentUser
sessionItem.UserPSK = userPSK
}
sessionItem.Unlock() sessionItem.Unlock()
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) { if bodyPlain[0] != HeaderTypeClient {
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload) 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 { if err != nil {
continue continue
} }
payload := bodyPlain[offset+addrLen:]
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)
}
}
entry.timer.Update()
pBuf := buf.New() pBuf := buf.New()
pBuf.Write(decoded.Payload) pBuf.Write(payload)
pBuf.UDP = &decoded.Destination _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
} }
} }
} }
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload) 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)
} }
+124 -59
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"crypto/cipher" "crypto/cipher"
"encoding/binary" "encoding/binary"
"io"
"strconv" "strconv"
"time" "time"
@@ -14,6 +15,9 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/common/uuid"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/policy"
@@ -31,17 +35,18 @@ type relayDest struct {
destination net.Destination destination net.Destination
email string email string
level uint32 level uint32
key []byte
blockCipher cipher.Block blockCipher cipher.Block
} }
type RelayInbound struct { type RelayInbound struct {
networks []net.Network networks []net.Network
method *CipherMethod method *CipherMethod
relayPSK []byte relayPSK []byte
relayBlock cipher.Block relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest destinations map[[AESBlockSize]byte]*relayDest
udpSessions *UDPSessionManager rawDestinations []*RelayDestination
policyManager policy.Manager policyManager policy.Manager
} }
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) { func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -73,13 +78,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
v := core.MustFromContext(ctx) v := core.MustFromContext(ctx)
i := &RelayInbound{ i := &RelayInbound{
networks: networks, networks: networks,
method: method, method: method,
relayPSK: relayPSK, relayPSK: relayPSK,
relayBlock: relayBlock, relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest), destinations: make(map[[AESBlockSize]byte]*relayDest),
udpSessions: NewUDPSessionManager(500 * time.Second), rawDestinations: config.Destinations,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
} }
for idx, d := range config.Destinations { for idx, d := range config.Destinations {
@@ -103,6 +108,7 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)), destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
email: d.Email, email: d.Email,
level: uint32(d.Level), level: uint32(d.Level),
key: destKey,
blockCipher: destBlock, blockCipher: destBlock,
} }
} }
@@ -133,36 +139,28 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
return errors.New("unable to set read deadline").Base(err) return errors.New("unable to set read deadline").Base(err)
} }
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4 // Read Salt + Outer EIH
needed := i.method.KeySaltLength + AESBlockSize needed := i.method.KeySaltLength + AESBlockSize
requestHeader := buf.New() var headerBuf [48]byte
n, err := requestHeader.ReadFrom(conn) headerSlice := headerBuf[:needed]
if err != nil { if _, err := io.ReadFull(conn, headerSlice); err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err return err
} }
if int(n) < needed {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
headerSlice := requestHeader.Bytes()
salt := headerSlice[:i.method.KeySaltLength] salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:needed] eih := headerSlice[i.method.KeySaltLength:]
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih) identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil { if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err return err
} }
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
targetDest, ok := i.destinations[decryptedHash] targetDest, ok := i.destinations[decryptedHash]
if !ok { if !ok {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest return ErrInvalidRequest
} }
conn.SetReadDeadline(time.Time{}) conn.SetReadDeadline(time.Time{})
@@ -184,26 +182,45 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
link, err := dispatcher.Dispatch(ctx, targetDest.destination) link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil { if err != nil {
requestHeader.Release()
return err return err
} }
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop // Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3). saltBuf := buf.New()
var saltCopy [32]byte saltBuf.Write(salt)
copy(saltCopy[:i.method.KeySaltLength], salt) if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
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 return err
} }
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link) 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)
} }
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn) 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)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -221,7 +238,11 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
var packetHeader [AESBlockSize]byte var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize]) i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize]) var eiHeader [AESBlockSize]byte
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
for idx := 0; idx < AESBlockSize; idx++ {
eiHeader[idx] ^= packetHeader[idx]
}
targetDest, ok := i.destinations[eiHeader] targetDest, ok := i.destinations[eiHeader]
if !ok { if !ok {
@@ -242,24 +263,68 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
dest := targetDest.destination dest := targetDest.destination
dest.Network = net.Network_UDP dest.Network = net.Network_UDP
sessionItem := i.udpSessions.GetOrCreate(sessionID) entry, ok := udpConns.Load(sessionID)
if sessionItem.User == nil { if !ok {
sessionItem.Lock() sessCtx, cancel := context.WithCancel(ctx)
if sessionItem.User == nil { inbound := session.InboundFromContext(sessCtx)
sessionItem.User = &protocol.MemoryUser{ inbound.User = &protocol.MemoryUser{
Email: targetDest.email, Email: targetDest.email,
Level: targetDest.level, 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.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
if err != nil {
b.Release()
continue
} }
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b}) entry.timer.Update()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
} }
} }
} }
-11
View File
@@ -61,14 +61,3 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
copy(out[:], h[:AESBlockSize]) copy(out[:], h[:AESBlockSize])
return out 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
}
+13 -33
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"crypto/rand" "crypto/rand"
"io" "io"
"time"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/buf"
@@ -45,12 +46,8 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
return nil, errors.New("invalid key: ", config.Key).Base(err) 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] finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, pskList) udpCodec, err := NewUDPPacketCodec(method, finalPSK)
if err != nil { if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err) return nil, errors.New("failed to create udp packet codec").Base(err)
} }
@@ -129,30 +126,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
requestDone := func() error { requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
var initialPayload []byte bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
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 { if err != nil {
buf.ReleaseMulti(remainingMB)
return errors.New("failed to write request").Base(err) return errors.New("failed to write request").Base(err)
} }
if !remainingMB.IsEmpty() { if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil { return errors.New("failed to write A request payload").Base(err)
return err }
}
if err := bufferedWriter.SetBuffered(false); err != nil {
return err
} }
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)) return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
@@ -178,18 +163,13 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
} }
if network == net.Network_UDP { 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 { requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{ writer := &UDPWriter{
Writer: conn, Writer: conn,
Destination: destination, Destination: destination,
Session: session, Codec: o.udpCodec,
} }
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
@@ -202,8 +182,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{ reader := &UDPReader{
Reader: conn, Reader: conn,
Session: session, Codec: o.udpCodec,
} }
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
+187 -441
View File
@@ -16,13 +16,14 @@ import (
) )
type UDPCodec struct { type UDPCodec struct {
method *CipherMethod method *CipherMethod
pskList [][]byte psk []byte
psk []byte blockCipher cipher.Block
blockCipher cipher.Block chachaCipher cipher.AEAD
blockCiphers []cipher.Block clientBodyCipher cipher.AEAD
chachaCipher cipher.AEAD clientSessionID uint64
sessions *UDPSessionManager nextPacketID atomic.Uint64
sessions *UDPSessionManager
} }
type ( type (
@@ -47,23 +48,22 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
return c, nil return c, nil
} }
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) { func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
if method.IsChaCha && len(pskList) > 1 { c, err := newUDPCodec(method, psk)
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 { if err != nil {
return nil, err return nil, err
} }
c.pskList = pskList var sessID [8]byte
if len(pskList) > 1 { if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
c.blockCiphers = make([]cipher.Block, len(pskList)) return nil, err
for i, psk := range pskList { }
c.blockCiphers[i], err = method.NewBlock(psk) c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
if err != nil {
return nil, err if !method.IsChaCha {
} clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
} }
} }
return c, nil return c, nil
@@ -78,37 +78,108 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
return c, nil return c, nil
} }
func (c *UDPCodec) Sessions() *UDPSessionManager { func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
return c.sessions packetID := c.nextPacketID.Add(1)
} sessID := c.clientSessionID
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession { // Padding determination (e.g. DNS port 53 disguise)
if c.sessions == nil { var paddingLen int
return nil if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
} }
return c.sessions.GetOrCreate(sessionID)
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:
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
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)
plainBytes := outBuf.Bytes()[16:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
} }
type DecodedUDPPacket struct { type DecodedUDPPacket struct {
SessionID uint64 SessionID uint64
PacketID uint64 PacketID uint64
HeaderType byte HeaderType byte
Timestamp uint64 Timestamp uint64
ClientSessionID uint64 Destination net.Destination
Destination net.Destination Payload []byte
Payload []byte
} }
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte { func parseAddressPort(data []byte) (net.Destination, int, error) {
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 { if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort return net.Destination{}, 0, ErrPacketTooShort
} }
@@ -149,9 +220,6 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
headerType := bodyPlain[0] headerType := bodyPlain[0]
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 { if diff > 30 {
@@ -159,13 +227,11 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
offset := 9 offset := 9
var clientSessionID uint64
if headerType == HeaderTypeServer { if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 { if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort return DecodedUDPPacket{}, ErrPacketTooShort
} }
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8]) offset += 8 // skip clientSessionID
offset += 8
} }
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2])) paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
@@ -176,20 +242,19 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
offset += paddingLen offset += paddingLen
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:]) dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil { if err != nil {
return DecodedUDPPacket{}, err return DecodedUDPPacket{}, err
} }
payload := bodyPlain[offset+addrLen:] payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{ return DecodedUDPPacket{
SessionID: sessionID, SessionID: sessionID,
PacketID: packetID, PacketID: packetID,
HeaderType: headerType, HeaderType: headerType,
Timestamp: epoch, Timestamp: epoch,
ClientSessionID: clientSessionID, Destination: dest,
Destination: dest, Payload: payload,
Payload: payload,
}, nil }, nil
} }
@@ -204,7 +269,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
} }
nonce := data[:PacketNonceSize] nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:] ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil) plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
if err != nil { if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err) return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
} }
@@ -215,22 +280,17 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(plain[:8]) sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16]) packetID := binary.BigEndian.Uint64(plain[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID) if c.sessions != nil {
if !sessionItem.CheckPacketID(packetID) { sessionItem := c.sessions.GetOrCreate(sessionID)
return DecodedUDPPacket{}, ErrPacketIdNotUnique sessionItem.Lock()
if !sessionItem.Window.CheckAndAdd(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
} }
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:]) return 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 // AES mode
@@ -239,52 +299,54 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(rawHeader[:8]) sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16]) packetID := binary.BigEndian.Uint64(rawHeader[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID) var bodyAead cipher.AEAD
if !sessionItem.CheckPacketID(packetID) { var sessionItem *ServerUDPSession
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:]) if c.sessions != nil {
} sessionItem = c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) { bodyAead = sessionItem.GetRemoteCipher()
bodyAead := s.clientBodyCipher if bodyAead == nil {
isNewCipher := false bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
if bodyAead == nil { var err error
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength) 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)
var err error var err error
bodyAead, err = method.NewAEAD(bodyKey) bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil { if err != nil {
return DecodedUDPPacket{}, err return DecodedUDPPacket{}, err
} }
isNewCipher = true
} }
bodyNonce := rawHeader[4:16] bodyNonce := rawHeader[4:16]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil) bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
if err != nil { if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err) return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
} }
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain) if sessionItem != nil {
if err != nil { sessionItem.Lock()
return DecodedUDPPacket{}, err sessionItem.Window.Add(packetID)
sessionItem.Unlock()
} }
if decoded.HeaderType != HeaderTypeClient { return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
return DecodedUDPPacket{}, ErrBadHeaderType
}
s.AddPacketID(packetID)
if isNewCipher {
s.clientBodyCipher = bodyAead
}
return decoded, nil
} }
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error { func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
s.Lock() s.Lock()
defer s.Unlock() defer s.Unlock()
if s.ServerSessionID != 0 { if s.ServerSessionID != 0 {
@@ -301,29 +363,23 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) e
} }
} }
if method.IsChaCha { if method.IsChaCha {
var err error s.ServerChaCha = chachaCipher
s.serverChaCha, err = method.NewUDPCipher(psk) } else {
return err s.ServerBlockCipher = headerBlock
} bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
bodyAead, err := method.NewAEAD(bodyKey)
var err error if err != nil {
s.serverHeaderBlock, err = method.NewBlock(psk) s.ServerSessionID = 0
if err != nil { return err
s.ServerSessionID = 0 }
return err s.ServerCipher = bodyAead
}
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
} }
return nil return nil
} }
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1) - 1 serverPacketID := s.ServerPacketID.Add(1)
if method.IsChaCha { if method.IsChaCha {
var nonce [PacketNonceSize]byte var nonce [PacketNonceSize]byte
@@ -348,7 +404,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
} }
plainBuf.Write(payload) 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)) res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:]) copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed) copy(res[PacketNonceSize:], sealed)
@@ -361,7 +417,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID) binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte var encryptedHeader [16]byte
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:]) s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New() bodyBuf := buf.New()
defer bodyBuf.Release() defer bodyBuf.Release()
@@ -379,7 +435,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
bodyBuf.Write(payload) bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16] bodyNonce := rawHeader[4:16]
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil) sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody)) res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:]) copy(res[:16], encryptedHeader[:])
@@ -388,327 +444,17 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
} }
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload) sessionItem := c.sessions.GetOrCreate(clientSessionID)
} if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
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 nil, err
} }
clientSessionID := binary.BigEndian.Uint64(sessID[:]) return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
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 { type UDPWriter struct {
Writer io.Writer Writer io.Writer
Destination net.Destination Destination net.Destination
Session *ClientUDPSession Codec *UDPPacketCodec
} }
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
@@ -722,7 +468,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if b.UDP != nil { if b.UDP != nil {
dest = *b.UDP dest = *b.UDP
} }
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes()) pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
b.Release() b.Release()
if err != nil { if err != nil {
buf.ReleaseMulti(mb) buf.ReleaseMulti(mb)
@@ -739,8 +485,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
} }
type UDPReader struct { type UDPReader struct {
Reader io.Reader Reader io.Reader
Session *ClientUDPSession Codec *UDPPacketCodec
} }
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) { func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
@@ -752,7 +498,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err return nil, err
} }
decoded, err := r.Session.DecodePacket(buffer.Bytes()) decoded, err := r.Codec.DecodePacket(buffer.Bytes())
if err != nil { if err != nil {
buffer.Release() buffer.Release()
continue continue
-105
View File
@@ -2,11 +2,9 @@ package shadowsocks_2022_test
import ( import (
"context" "context"
"crypto/rand"
"encoding/base64" "encoding/base64"
"encoding/binary" "encoding/binary"
"errors" "errors"
"io"
gonet "net" gonet "net"
"sync" "sync"
"sync/atomic" "sync/atomic"
@@ -271,106 +269,3 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
} }
return nil 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))
}
})
}
}
+16 -41
View File
@@ -6,11 +6,8 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "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/common/utils"
"github.com/xtls/xray-core/transport"
) )
const ( const (
@@ -77,42 +74,30 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
type ServerUDPSession struct { type ServerUDPSession struct {
sync.Mutex sync.Mutex
SessionID uint64 SessionID uint64
Window *SlidingWindow RemoteCipher atomic.Pointer[cipher.AEAD]
User *protocol.MemoryUser Window SlidingWindow
UserPSK []byte User *protocol.MemoryUser
LastActive atomic.Int64 // Unix timestamp in seconds UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
clientBodyCipher cipher.AEAD
ServerSessionID uint64 ServerSessionID uint64
ServerPacketID atomic.Uint64 ServerPacketID atomic.Uint64
serverBodyCipher cipher.AEAD ServerCipher cipher.AEAD
serverHeaderBlock cipher.Block ServerBlockCipher cipher.Block
serverChaCha cipher.AEAD ServerChaCha cipher.AEAD
manager *UDPSessionManager
link atomic.Pointer[transport.Link]
timer *signal.ActivityTimer
currentConn atomic.Value // stores stat.Connection
} }
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool { func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
s.Lock() ptr := s.RemoteCipher.Load()
defer s.Unlock() if ptr == nil {
if s.Window == nil { return nil
s.Window = new(SlidingWindow)
} }
return s.Window.Check(packetID) return *ptr
} }
func (s *ServerUDPSession) AddPacketID(packetID uint64) { func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
s.Lock() s.RemoteCipher.Store(&c)
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
s.Window.Add(packetID)
} }
type UDPSessionManager struct { type UDPSessionManager struct {
@@ -137,7 +122,6 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
s := &ServerUDPSession{ s := &ServerUDPSession{
SessionID: sessionID, SessionID: sessionID,
manager: m,
} }
s.LastActive.Store(now) s.LastActive.Store(now)
@@ -164,7 +148,6 @@ func (m *UDPSessionManager) cleanup(now int64) {
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool { m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec { if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k) m.sessions.Delete(k)
v.Close()
} }
return true return true
}) })
@@ -173,11 +156,3 @@ func (m *UDPSessionManager) cleanup(now int64) {
func (m *UDPSessionManager) Delete(sessionID uint64) { func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID) 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)
}
+6 -149
View File
@@ -2,161 +2,18 @@ package shadowsocks_2022
import ( import (
"context" "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/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/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"
"github.com/xtls/xray-core/transport/internet/stat"
) )
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) { type udpConnEntry struct {
if s.currentConn.Load() == nil { sync.Mutex
s.currentConn.Store(conn) link *transport.Link
} timer *signal.ActivityTimer
if s.timer != nil { cancel context.CancelFunc
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 ( const (
+36 -172
View File
@@ -182,48 +182,57 @@ func TestTCPStream(t *testing.T) {
common.Must(err) common.Must(err)
IncreaseNonce(reader.Nonce()) IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar) vBuf := buf.New()
vBuf.Write(plainVar)
receivedDest, err = ReadAddressPort(vBuf)
common.Must(err) common.Must(err)
receivedDest = net.TCPDestination(dest.Address, dest.Port)
plainVar = plainVar[addrLen:]
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
receivedPayload = plainVar[2+padLen:]
// Server sends response stream with receivedPayload as first payload // Skip padding
writer := NewServerStreamWriter(serverConn, method, rawKey, salt) var padBytes [2]byte
pBuf := buf.New() _, _ = vBuf.Read(padBytes[:])
pBuf.Write(receivedPayload) padLen := int(padBytes[0])<<8 | int(padBytes[1])
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) vBuf.Advance(int32(padLen))
// Read and echo additional stream data 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
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
common.Must(err) common.Must(err)
_ = writer.WriteMultiBuffer(mb) _ = writer.WriteMultiBuffer(mb)
_ = writer.Close()
}() }()
// Client goroutine // Client goroutine
go func() { go func() {
defer wg.Done() defer wg.Done()
clientSalt := make([]byte, method.KeySaltLength) clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
common.Must2(io.ReadFull(rand.Reader, clientSalt))
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
common.Must(err) common.Must(err)
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt) reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
common.Must(err) 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 // Send additional stream data
streamData := []byte("stream chunk test") streamData := []byte("stream chunk test")
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)}) _ = writer.WriteChunk(streamData)
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
common.Must(err) common.Must(err)
@@ -263,14 +272,12 @@ func TestUDPCodec(t *testing.T) {
psk := make([]byte, method.KeySaltLength) psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk) _, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk}) clientCodec, err := NewUDPPacketCodec(method, psk)
common.Must(err) common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute) serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err) common.Must(err)
session, err := clientCodec.NewClientSession() pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
common.Must(err)
pktBuf, err := session.EncodePacket(dest, payload)
common.Must(err) common.Must(err)
defer pktBuf.Release() defer pktBuf.Release()
@@ -353,146 +360,3 @@ func TestMultiUserManager(t *testing.T) {
t.Fatal("user1 should have been removed") t.Fatal("user1 should have been removed")
} }
} }
func TestLargeStreamTransfer(t *testing.T) {
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
sessionKey := make([]byte, 16)
_, _ = rand.Read(sessionKey)
clientAead, err := method.NewAEAD(sessionKey)
common.Must(err)
serverAead, err := method.NewAEAD(sessionKey)
common.Must(err)
r, w := io.Pipe()
defer r.Close()
defer w.Close()
writer := NewStreamWriter(w, clientAead)
reader := NewStreamReader(r, serverAead)
const totalSize = 100 * 1024 // 100 KB
data := make([]byte, totalSize)
_, _ = rand.Read(data)
errCh := make(chan error, 1)
go func() {
// Write using Write (which splits by MaxPacketSize = 65535)
_, werr := writer.Write(data)
if werr != nil {
errCh <- werr
return
}
_ = w.Close()
errCh <- nil
}()
var received []byte
for {
mb, rerr := reader.ReadMultiBuffer()
if !mb.IsEmpty() {
for _, b := range mb {
received = append(received, b.Bytes()...)
}
buf.ReleaseMulti(mb)
}
if rerr != nil {
if rerr == io.EOF {
break
}
t.Fatalf("ReadMultiBuffer error: %v", rerr)
}
}
if werr := <-errCh; werr != nil {
t.Fatalf("writer error: %v", werr)
}
if len(received) != totalSize {
t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize)
}
if !bytes.Equal(received, data) {
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")
}
})
}
}
+114 -234
View File
@@ -1,25 +1,18 @@
package shadowsocks_2022 package shadowsocks_2022
import ( import (
"context"
"crypto/cipher" "crypto/cipher"
"crypto/rand" "crypto/rand"
"encoding/binary" "encoding/binary"
"io" "io"
"math" "math"
mrand "math/rand/v2" mrand "math/rand/v2"
"sync"
"time" "time"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "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( var addrParser = protocol.NewAddressParser(
@@ -45,6 +38,15 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port) 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 // AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int { func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() { switch dest.Address.Family() {
@@ -117,16 +119,8 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb) defer buf.ReleaseMulti(mb)
for _, b := range mb { for _, b := range mb {
p := b.Bytes() if err := w.WriteChunk(b.Bytes()); err != nil {
for len(p) > 0 { return err
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return err
}
p = p[chunkSize:]
} }
} }
return nil return nil
@@ -174,7 +168,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize { if payloadLen == 0 {
return 0, ErrInvalidRequest return 0, ErrInvalidRequest
} }
@@ -200,10 +194,11 @@ func (r *StreamReader) Read(p []byte) (int, error) {
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) { func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 { if r.cached > 0 {
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached]) b := buf.New()
b.Write(r.buffer[r.offset : r.offset+r.cached])
r.cached = 0 r.cached = 0
r.offset = 0 r.offset = 0
return mb, nil return buf.MultiBuffer{b}, nil
} }
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
@@ -217,7 +212,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize { if payloadLen == 0 {
return nil, ErrInvalidRequest return nil, ErrInvalidRequest
} }
@@ -232,8 +227,9 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
} }
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
mb := buf.MergeBytes(nil, decryptedPayload) b := buf.New()
return mb, nil b.Write(decryptedPayload)
return buf.MultiBuffer{b}, nil
} }
type ClientRequestHeader struct { type ClientRequestHeader struct {
@@ -241,8 +237,13 @@ type ClientRequestHeader struct {
EarlyData []byte EarlyData []byte
} }
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) { func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil) 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)
if err != nil { if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err) return nil, errors.New("failed to decrypt client request header").Base(err)
} }
@@ -271,7 +272,7 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
} else { } else {
varChunkCipher = make([]byte, needed) varChunkCipher = make([]byte, needed)
} }
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil { if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
return nil, err return nil, err
} }
@@ -281,34 +282,31 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
} }
IncreaseNonce(reader.Nonce()) IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar) b := buf.New()
b.Write(plainVar)
defer b.Release()
dest, err := ReadAddressPort(b)
if err != nil { if err != nil {
return nil, err return nil, err
} }
dest.Network = net.Network_TCP
offset := addrLen var padLenBytes [2]byte
if len(plainVar) < offset+2 { if _, err := b.Read(padLenBytes[:]); err != nil {
return nil, ErrPacketTooShort return nil, err
} }
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2])) paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
offset += 2 if int(b.Len()) < paddingLen {
if len(plainVar) < offset+paddingLen {
return nil, ErrNoPadding return nil, ErrNoPadding
} }
offset += paddingLen if paddingLen > 0 {
b.Advance(int32(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. var earlyData []byte
if paddingLen == 0 && payloadLen == 0 { if b.Len() > 0 {
return nil, errors.New("request without payload and padding is not allowed") earlyData = make([]byte, b.Len())
copy(earlyData, b.Bytes())
} }
return &ClientRequestHeader{ return &ClientRequestHeader{
@@ -317,6 +315,34 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
}, nil }, 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. // 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) { func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1] finalPSK := pskList[len(pskList)-1]
@@ -328,16 +354,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
writer := NewStreamWriter(w, aead) writer := NewStreamWriter(w, aead)
payloadLen := len(payload) handshakeBuf := buf.New()
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() defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt) handshakeBuf.Write(clientSalt)
@@ -355,6 +372,14 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
handshakeBuf.Write(encryptedEIH[:]) 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 var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix())) binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
@@ -364,7 +389,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
IncreaseNonce(writer.nonce[:]) IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk) handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen)) varHeaderBuf := buf.New()
defer varHeaderBuf.Release() defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil { if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
@@ -396,21 +421,12 @@ 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. // 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) { func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2 var serverSalt [32]byte
chunkCipherLen := fixedPlainLen + AEADTagSize serverSaltSlice := serverSalt[:method.KeySaltLength]
headerLen := method.KeySaltLength + chunkCipherLen if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
return nil, err
// 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) sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey) aead, err := method.NewAEAD(sessionKey)
if err != nil { if err != nil {
@@ -419,6 +435,14 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
reader := NewStreamReader(r, aead) 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) decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil { if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err) return nil, errors.New("failed to decrypt server response header").Base(err)
@@ -460,190 +484,46 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
return reader, nil return reader, nil
} }
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4. // WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
type ServerStreamWriter struct { func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
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 var serverSalt [32]byte
serverSaltSlice := serverSalt[:s.method.KeySaltLength] serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil { if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err return nil, err
} }
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength) respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
respAead, err := s.method.NewAEAD(respKey) respAead, err := method.NewAEAD(respKey)
if err != nil { if err != nil {
return nil, err return nil, err
} }
sw := NewStreamWriter(s.w, respAead) writer := NewStreamWriter(w, respAead)
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize) respBuf := buf.New()
outBuf := buf.NewWithSize(totalHeaderLen) defer respBuf.Release()
defer outBuf.Release()
outBuf.Write(serverSaltSlice) respBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2] fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix())) binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt) copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload))) binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil) fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
IncreaseNonce(sw.nonce[:]) IncreaseNonce(writer.nonce[:])
outBuf.Write(fixedRespChunk) respBuf.Write(fixedRespChunk)
if len(payload) > 0 { if len(initialPayload) > 0 {
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil) initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
IncreaseNonce(sw.nonce[:]) IncreaseNonce(writer.nonce[:])
outBuf.Write(payloadChunk) respBuf.Write(initialChunk)
} }
if _, err := s.w.Write(outBuf.Bytes()); err != nil { if _, err := w.Write(respBuf.Bytes()); err != nil {
return nil, err return nil, err
} }
return sw, nil
}
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { return writer, nil
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)
} }
+3 -4
View File
@@ -27,7 +27,6 @@ import (
"github.com/xtls/xray-core/features/stats" "github.com/xtls/xray-core/features/stats"
"github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
"golang.zx2c4.com/wireguard/device" "golang.zx2c4.com/wireguard/device"
) )
@@ -200,7 +199,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
} }
defer conn.Close() defer conn.Close()
c := &UDPConnClient{ c := &UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn, PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr), Dest: conn.RemoteAddr().(*net.UDPAddr),
} }
reader = c reader = c
@@ -264,14 +263,14 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn pktConn = conn.(*net.PacketConnWrapper).PacketConn
} else { } else {
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings) conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
switch c := conn.(type) { switch c := conn.(type) {
case *internet.PacketConnWrapper: case *net.PacketConnWrapper:
pktConn = c.PacketConn pktConn = c.PacketConn
case *cnc.Connection: case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c} pktConn = &internet.FakePacketConn{Conn: c}
+2 -2
View File
@@ -21,7 +21,7 @@ import (
"syscall" "syscall"
"time" "time"
"github.com/xtls/xray-core/transport/internet" xnet "github.com/xtls/xray-core/common/net"
"golang.zx2c4.com/wireguard/tun" "golang.zx2c4.com/wireguard/tun"
"golang.org/x/net/dns/dnsmessage" "golang.org/x/net/dns/dnsmessage"
@@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &internet.PacketConnWrapper{ return &xnet.PacketConnWrapper{
PacketConn: conn, PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr), Dest: net.UDPAddrFromAddrPort(raddr),
}, nil }, nil
+2 -2
View File
@@ -16,7 +16,7 @@ import (
"github.com/vishvananda/netlink" "github.com/vishvananda/netlink"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet" xnet "github.com/xtls/xray-core/common/net"
"golang.zx2c4.com/wireguard/tun" "golang.zx2c4.com/wireguard/tun"
) )
@@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &internet.PacketConnWrapper{ return &xnet.PacketConnWrapper{
PacketConn: conn, PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr), Dest: net.UDPAddrFromAddrPort(raddr),
}, nil }, nil
+4 -22
View File
@@ -82,7 +82,7 @@ func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Con
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
}, },
} }
for i := range fm.tcpMasks { for i := range fm.tcpMasks {
@@ -144,7 +144,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
} }
for i := range fm.udpMasks { for i := range fm.udpMasks {
if i > 0 { if i > 0 {
@@ -171,7 +171,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
}, },
} }
var sizes []int var sizes []int
@@ -208,7 +208,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if addr == nil { if addr == nil {
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}} addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
} }
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
} }
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) { func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
@@ -272,24 +272,6 @@ const (
UDPSize = 4096 UDPSize = 4096
) )
type PacketConnWrapper struct {
net.PacketConn
udpAddr net.Addr
}
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
return c.udpAddr
}
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
n, _, err = c.PacketConn.ReadFrom(b)
return
}
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
return c.PacketConn.WriteTo(b, c.udpAddr)
}
type headerManagerConn struct { type headerManagerConn struct {
net.PacketConn net.PacketConn
+1 -1
View File
@@ -380,7 +380,7 @@ func TestPacketConnReadWrite(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { clientConn.Close() }) t.Cleanup(func() { clientConn.Close() })
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn client := clientConn.(*net.PacketConnWrapper).PacketConn
_ = client.SetDeadline(time.Now().Add(time.Second)) _ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second))
+2 -2
View File
@@ -73,7 +73,7 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (
if err != nil { if err != nil {
return nil, err return nil, err
} }
cur := conn.(*finalmask.PacketConnWrapper).PacketConn cur := conn.(*net.PacketConnWrapper).PacketConn
addr := conn.RemoteAddr().(*net.UDPAddr) addr := conn.RemoteAddr().(*net.UDPAddr)
client := &udpHopConn{ client := &udpHopConn{
dialer: dialer, dialer: dialer,
@@ -150,7 +150,7 @@ func (c *udpHopConn) hop() {
_ = c.pre.Close() _ = c.pre.Close()
} }
c.pre = c.cur c.pre = c.cur
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn c.cur = conn.(*net.PacketConnWrapper).PacketConn
c.wg.Add(1) c.wg.Add(1)
go c.recv(c.cur) go c.recv(c.cur)
} }
+2 -2
View File
@@ -119,7 +119,7 @@ func (c *client) dial(ctx context.Context) error {
if err != nil { if err != nil {
return errors.New("failed to dial to dest").Base(err) return errors.New("failed to dial to dest").Base(err)
} }
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr() udpAddr = conn.RemoteAddr()
} else { } else {
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig) conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
@@ -127,7 +127,7 @@ func (c *client) dial(ctx context.Context) error {
return errors.New("failed to dial to dest").Base(err) return errors.New("failed to dial to dest").Base(err)
} }
switch c := conn.(type) { switch c := conn.(type) {
case *internet.PacketConnWrapper: case *net.PacketConnWrapper:
pktConn = c.PacketConn pktConn = c.PacketConn
udpAddr = c.RemoteAddr() udpAddr = c.RemoteAddr()
case *cnc.Connection: case *cnc.Connection:
+2 -3
View File
@@ -19,7 +19,6 @@ import (
"github.com/xtls/xray-core/common/net/cnc" "github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion" "github.com/xtls/xray-core/transport/internet/hysteria/congestion"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
"github.com/xtls/xray-core/transport/internet/masque/connectip" "github.com/xtls/xray-core/transport/internet/masque/connectip"
@@ -80,7 +79,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr() udpAddr = conn.RemoteAddr()
} else { } else {
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
@@ -88,7 +87,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
switch c := conn.(type) { switch c := conn.(type) {
case *internet.PacketConnWrapper: case *net.PacketConnWrapper:
pktConn = c.PacketConn pktConn = c.PacketConn
udpAddr = c.RemoteAddr() udpAddr = c.RemoteAddr()
case *cnc.Connection: case *cnc.Connection:
+30 -31
View File
@@ -54,11 +54,10 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
mss.SecurityType = s.SecurityType mss.SecurityType = s.SecurityType
mss.SecuritySettings = ess mss.SecuritySettings = ess
} }
if s != nil && (len(s.Tcpmasks) != 0 || len(s.Udpmasks) != 0) {
var tcpMasks []finalmask.TCPMask
var udpMasks []finalmask.UDPMask
var tcpMasks []finalmask.TCPMask
var udpMasks []finalmask.UDPMask
if s != nil {
for i := range s.Tcpmasks { for i := range s.Tcpmasks {
instance := common.Must2(s.Tcpmasks[i].GetInstance()) instance := common.Must2(s.Tcpmasks[i].GetInstance())
tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask)) tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask))
@@ -67,37 +66,37 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
instance := common.Must2(s.Udpmasks[i].GetInstance()) instance := common.Must2(s.Udpmasks[i].GetInstance())
udpMasks = append(udpMasks, instance.(finalmask.UDPMask)) udpMasks = append(udpMasks, instance.(finalmask.UDPMask))
} }
}
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return DialSystem(ctx, dest, mss.SocketSettings) return DialSystem(ctx, dest, mss.SocketSettings)
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return ListenSystem(ctx, addr, mss.SocketSettings)
}
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
if err != nil {
return nil, nil, err
} }
var newConn net.PacketConn listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
var udpAddr net.Addr return ListenSystem(ctx, addr, mss.SocketSettings)
switch c := conn.(type) {
case *PacketConnWrapper:
newConn = c.PacketConn
udpAddr = conn.RemoteAddr()
case *cnc.Connection:
newConn = &FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
default:
panic(reflect.TypeOf(c))
} }
return newConn, udpAddr, nil dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
if err != nil {
return nil, nil, err
}
var newConn net.PacketConn
var udpAddr net.Addr
switch c := conn.(type) {
case *net.PacketConnWrapper:
newConn = c.PacketConn
udpAddr = conn.RemoteAddr()
case *cnc.Connection:
newConn = &FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
default:
panic(reflect.TypeOf(c))
}
return newConn, udpAddr, nil
}
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
}
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
} }
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
}
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
if s != nil && s.QuicParams != nil { if s != nil && s.QuicParams != nil {
mss.QuicParams = s.QuicParams mss.QuicParams = s.QuicParams
+2 -3
View File
@@ -25,7 +25,6 @@ import (
"github.com/xtls/xray-core/common/signal/done" "github.com/xtls/xray-core/common/signal/done"
"github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/browser_dialer" "github.com/xtls/xray-core/transport/internet/browser_dialer"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion" "github.com/xtls/xray-core/transport/internet/hysteria/congestion"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
"github.com/xtls/xray-core/transport/internet/reality" "github.com/xtls/xray-core/transport/internet/reality"
@@ -200,7 +199,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr() udpAddr = conn.RemoteAddr()
} else { } else {
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
@@ -208,7 +207,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
return nil, errors.New("failed to dial to dest").Base(err) return nil, errors.New("failed to dial to dest").Base(err)
} }
switch c := conn.(type) { switch c := conn.(type) {
case *internet.PacketConnWrapper: case *net.PacketConnWrapper:
pktConn = c.PacketConn pktConn = c.PacketConn
udpAddr = c.RemoteAddr() udpAddr = c.RemoteAddr()
case *cnc.Connection: case *cnc.Connection:
+1 -19
View File
@@ -86,7 +86,7 @@ func (d *DefaultSystemDialer) Dial(ctx context.Context, src net.Address, dest ne
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &PacketConnWrapper{ return &net.PacketConnWrapper{
PacketConn: packetConn, PacketConn: packetConn,
Dest: destAddr, Dest: destAddr,
}, nil }, nil
@@ -148,24 +148,6 @@ func (d *DefaultSystemDialer) DestIpAddress() net.IP {
return nil return nil
} }
type PacketConnWrapper struct {
net.PacketConn
Dest net.Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
return c.Dest
}
type SystemDialerAdapter interface { type SystemDialerAdapter interface {
Dial(network string, address string) (net.Conn, error) Dial(network string, address string) (net.Conn, error)
} }