mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 05:46:39 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a26cce936b | ||
|
|
7dc35bac94 |
@@ -2,7 +2,6 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"io"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -13,9 +12,6 @@ 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"
|
||||||
@@ -101,35 +97,29 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
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]
|
||||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
fixedChunk := headerBuf[i.method.KeySaltLength:]
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if !i.saltFilter.Check(salt) {
|
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
return ErrSaltNotUnique
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionKey := DeriveSessionSubKey(i.psk, 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)
|
|
||||||
|
|
||||||
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
conn.SetReadDeadline(time.Time{})
|
|
||||||
dest := reqHeader.Destination
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
|
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
From: conn.RemoteAddr(),
|
From: conn.RemoteAddr(),
|
||||||
@@ -146,42 +136,17 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(reqHeader.EarlyData) > 0 {
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
earlyBuf := buf.New()
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
earlyBuf.Write(reqHeader.EarlyData)
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
|
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
|
||||||
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 {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
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 {
|
||||||
@@ -191,75 +156,30 @@ 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())
|
||||||
if err != nil {
|
b.Release()
|
||||||
b.Release()
|
if err != nil || decoded.HeaderType != HeaderTypeClient {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
entry, ok := udpConns.Load(decoded.SessionID)
|
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
|
||||||
if !ok {
|
if sessionItem.User == nil {
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
sessionItem.Lock()
|
||||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
if sessionItem.User == nil {
|
||||||
From: conn.RemoteAddr(),
|
sessionItem.User = i.user
|
||||||
To: decoded.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: i.user.Email,
|
|
||||||
})
|
|
||||||
|
|
||||||
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
b.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
newEntry := &udpConnEntry{
|
|
||||||
link: link,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
|
|
||||||
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
|
||||||
udpConns.Delete(decoded.SessionID)
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(link.Writer)
|
|
||||||
cancel()
|
|
||||||
}, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
|
|
||||||
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
|
|
||||||
if loaded {
|
|
||||||
// Another goroutine/packet beat us to storing, terminate our redundant link
|
|
||||||
newEntry.timer.SetTimeout(0)
|
|
||||||
entry = actual
|
|
||||||
} else {
|
|
||||||
entry = newEntry
|
|
||||||
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
|
|
||||||
defer func() {
|
|
||||||
cEntry.timer.SetTimeout(0)
|
|
||||||
}()
|
|
||||||
for {
|
|
||||||
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cEntry.timer.Update()
|
|
||||||
for _, rb := range resMb {
|
|
||||||
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
|
|
||||||
rb.Release()
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
_, _ = conn.Write(encPacket)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}(decoded.SessionID, decoded.Destination, entry)
|
|
||||||
}
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
entry.timer.Update()
|
|
||||||
payloadBuf := buf.New()
|
payloadBuf := buf.New()
|
||||||
payloadBuf.Write(decoded.Payload)
|
payloadBuf.Write(decoded.Payload)
|
||||||
b.Release()
|
payloadBuf.UDP = &decoded.Destination
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -19,8 +18,6 @@ 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"
|
||||||
@@ -207,64 +204,46 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 1. Read Request Salt (16 or 32 bytes)
|
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
|
||||||
|
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := conn.Read(headerBuf)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return errors.New("failed to read complete handshake header")
|
||||||
|
}
|
||||||
|
|
||||||
var salt [32]byte
|
var salt [32]byte
|
||||||
|
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||||
saltSlice := salt[:i.method.KeySaltLength]
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
|
||||||
return err
|
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
|
||||||
}
|
|
||||||
|
|
||||||
if !i.saltFilter.Check(salt) {
|
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
|
||||||
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 {
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
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 || user == nil {
|
if !ok {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return ErrInvalidRequest
|
return ErrInvalidRequest
|
||||||
}
|
}
|
||||||
userPSK := user.Account.(*MemoryAccount).Key
|
userPSK := user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
// 3. Derive Session Subkey using matched user's PSK
|
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
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
|
||||||
|
|
||||||
// 6. Send Server Response Handshake
|
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
|
||||||
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// 7. Dispatch Connection to Xray routing with matched User
|
// Dispatch Connection to Xray routing with matched User
|
||||||
inbound := session.InboundFromContext(ctx)
|
inbound := session.InboundFromContext(ctx)
|
||||||
inbound.User = user
|
inbound.User = user
|
||||||
|
|
||||||
@@ -283,42 +262,17 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(reqHeader.EarlyData) > 0 {
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
earlyBuf := buf.New()
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
earlyBuf.Write(reqHeader.EarlyData)
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(user.Level)
|
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
|
||||||
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 {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
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 {
|
||||||
@@ -342,168 +296,61 @@ 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)
|
||||||
|
|
||||||
sessionItem.Lock()
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
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
|
||||||
if sessionItem.User != nil {
|
sessionItem.Lock()
|
||||||
currentUser = sessionItem.User
|
currentUser = sessionItem.User
|
||||||
userPSK = sessionItem.UserPSK
|
userPSK = sessionItem.UserPSK
|
||||||
sessionItem.Unlock()
|
sessionItem.Unlock()
|
||||||
} else {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
// Decrypt EIH
|
|
||||||
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
|
if currentUser == nil {
|
||||||
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
|
// Decrypt EIH
|
||||||
|
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
|
||||||
|
|
||||||
user, ok := i.usersByHash.Load(decryptedHash)
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
if !ok || user == nil {
|
if !ok {
|
||||||
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()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decrypt Body (with AEAD caching per session)
|
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
|
||||||
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 || len(bodyPlain) < 1+8+2 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionItem.Lock()
|
|
||||||
sessionItem.Window.Add(packetID)
|
|
||||||
sessionItem.Unlock()
|
|
||||||
|
|
||||||
if bodyPlain[0] != HeaderTypeClient {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
|
||||||
diff := time.Now().Unix() - int64(epoch)
|
|
||||||
if diff < -30 || diff > 30 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
|
|
||||||
offset := 11 + paddingLen
|
|
||||||
if len(bodyPlain) < offset {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := bodyPlain[offset+addrLen:]
|
sessionItem.Lock()
|
||||||
|
if sessionItem.User == nil {
|
||||||
|
sessionItem.User = currentUser
|
||||||
|
sessionItem.UserPSK = userPSK
|
||||||
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
entry, ok := udpConns.Load(sessionID)
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
|
||||||
if !ok {
|
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
})
|
||||||
inbound := session.InboundFromContext(sessCtx)
|
if err != nil {
|
||||||
inbound.User = currentUser
|
continue
|
||||||
|
|
||||||
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(payload)
|
pBuf.Write(decoded.Payload)
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
pBuf.UDP = &decoded.Destination
|
||||||
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
|
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
|
||||||
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -15,9 +14,6 @@ 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"
|
||||||
@@ -35,18 +31,17 @@ 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
|
||||||
rawDestinations []*RelayDestination
|
udpSessions *UDPSessionManager
|
||||||
policyManager policy.Manager
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||||
@@ -78,13 +73,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
|||||||
|
|
||||||
v := core.MustFromContext(ctx)
|
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),
|
||||||
rawDestinations: config.Destinations,
|
udpSessions: NewUDPSessionManager(500 * time.Second),
|
||||||
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 {
|
||||||
@@ -108,7 +103,6 @@ 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,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -139,28 +133,36 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Read Salt + Outer EIH
|
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
|
||||||
needed := i.method.KeySaltLength + AESBlockSize
|
needed := i.method.KeySaltLength + AESBlockSize
|
||||||
var headerBuf [48]byte
|
requestHeader := buf.New()
|
||||||
headerSlice := headerBuf[:needed]
|
n, err := requestHeader.ReadFrom(conn)
|
||||||
if _, err := io.ReadFull(conn, headerSlice); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
salt := headerSlice[:i.method.KeySaltLength]
|
|
||||||
eih := headerSlice[i.method.KeySaltLength:]
|
|
||||||
|
|
||||||
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
|
|
||||||
block, err := i.method.NewBlock(identitySubkey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if int(n) < needed {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
var decryptedHash [AESBlockSize]byte
|
headerSlice := requestHeader.Bytes()
|
||||||
block.Decrypt(decryptedHash[:], eih)
|
salt := headerSlice[:i.method.KeySaltLength]
|
||||||
|
eih := headerSlice[i.method.KeySaltLength:needed]
|
||||||
|
|
||||||
|
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
|
||||||
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
targetDest, ok := i.destinations[decryptedHash]
|
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{})
|
||||||
@@ -182,45 +184,26 @@ 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 to next hop, stripping this hop's EIH
|
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
|
||||||
saltBuf := buf.New()
|
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
|
||||||
saltBuf.Write(salt)
|
var saltCopy [32]byte
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
|
copy(saltCopy[:i.method.KeySaltLength], salt)
|
||||||
|
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
|
||||||
|
requestHeader.Advance(AESBlockSize)
|
||||||
|
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
|
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
|
||||||
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 {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
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 {
|
||||||
@@ -238,11 +221,7 @@ 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])
|
||||||
|
|
||||||
var eiHeader [AESBlockSize]byte
|
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||||
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 {
|
||||||
@@ -263,68 +242,24 @@ 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
|
||||||
|
|
||||||
entry, ok := udpConns.Load(sessionID)
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
if !ok {
|
if sessionItem.User == nil {
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
sessionItem.Lock()
|
||||||
inbound := session.InboundFromContext(sessCtx)
|
if sessionItem.User == nil {
|
||||||
inbound.User = &protocol.MemoryUser{
|
sessionItem.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
|
||||||
}
|
}
|
||||||
|
|
||||||
entry.timer.Update()
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -61,3 +61,14 @@ 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
|
||||||
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ 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"
|
||||||
@@ -46,8 +45,12 @@ 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, finalPSK)
|
udpCodec, err := NewUDPPacketCodec(method, pskList)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -126,18 +129,30 @@ 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))
|
|
||||||
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
|
var initialPayload []byte
|
||||||
|
var firstBuf *buf.Buffer
|
||||||
|
var remainingMB buf.MultiBuffer
|
||||||
|
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
|
||||||
|
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
|
||||||
|
remainingMB, firstBuf = buf.SplitFirst(mb)
|
||||||
|
initialPayload = firstBuf.Bytes()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
|
||||||
|
if firstBuf != nil {
|
||||||
|
firstBuf.Release()
|
||||||
|
}
|
||||||
if err != nil {
|
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 err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
if !remainingMB.IsEmpty() {
|
||||||
return errors.New("failed to write A request payload").Base(err)
|
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
|
||||||
}
|
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))
|
||||||
@@ -163,13 +178,18 @@ 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,
|
||||||
Codec: o.udpCodec,
|
Session: session,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
@@ -182,8 +202,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
|||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
|
||||||
reader := &UDPReader{
|
reader := &UDPReader{
|
||||||
Reader: conn,
|
Reader: conn,
|
||||||
Codec: o.udpCodec,
|
Session: session,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
|
|||||||
+441
-187
@@ -16,14 +16,13 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type UDPCodec struct {
|
type UDPCodec struct {
|
||||||
method *CipherMethod
|
method *CipherMethod
|
||||||
psk []byte
|
pskList [][]byte
|
||||||
blockCipher cipher.Block
|
psk []byte
|
||||||
chachaCipher cipher.AEAD
|
blockCipher cipher.Block
|
||||||
clientBodyCipher cipher.AEAD
|
blockCiphers []cipher.Block
|
||||||
clientSessionID uint64
|
chachaCipher cipher.AEAD
|
||||||
nextPacketID atomic.Uint64
|
sessions *UDPSessionManager
|
||||||
sessions *UDPSessionManager
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type (
|
type (
|
||||||
@@ -48,22 +47,23 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
|
||||||
c, err := newUDPCodec(method, psk)
|
if method.IsChaCha && len(pskList) > 1 {
|
||||||
|
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||||
|
}
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
c, err := newUDPCodec(method, finalPSK)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
var sessID [8]byte
|
c.pskList = pskList
|
||||||
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
if len(pskList) > 1 {
|
||||||
return nil, err
|
c.blockCiphers = make([]cipher.Block, len(pskList))
|
||||||
}
|
for i, psk := range pskList {
|
||||||
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
|
c.blockCiphers[i], err = method.NewBlock(psk)
|
||||||
|
if err != nil {
|
||||||
if !method.IsChaCha {
|
return nil, err
|
||||||
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,108 +78,37 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
func (c *UDPCodec) Sessions() *UDPSessionManager {
|
||||||
packetID := c.nextPacketID.Add(1)
|
return c.sessions
|
||||||
sessID := c.clientSessionID
|
}
|
||||||
|
|
||||||
// Padding determination (e.g. DNS port 53 disguise)
|
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
|
||||||
var paddingLen int
|
if c.sessions == nil {
|
||||||
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
return nil
|
||||||
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
|
||||||
Destination net.Destination
|
ClientSessionID uint64
|
||||||
Payload []byte
|
Destination net.Destination
|
||||||
|
Payload []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseAddressPort(data []byte) (net.Destination, int, error) {
|
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
for k := 0; k < AESBlockSize; k++ {
|
||||||
|
decryptedHash[k] ^= rawHeader[k]
|
||||||
|
}
|
||||||
|
return decryptedHash
|
||||||
|
}
|
||||||
|
|
||||||
|
func ParseAddressPort(data []byte) (net.Destination, int, error) {
|
||||||
if len(data) < 1 {
|
if len(data) < 1 {
|
||||||
return net.Destination{}, 0, ErrPacketTooShort
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
}
|
}
|
||||||
@@ -220,6 +149,9 @@ 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 {
|
||||||
@@ -227,11 +159,13 @@ 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
|
||||||
}
|
}
|
||||||
offset += 8 // skip clientSessionID
|
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
|
||||||
|
offset += 8
|
||||||
}
|
}
|
||||||
|
|
||||||
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
||||||
@@ -242,19 +176,20 @@ 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,
|
||||||
Destination: dest,
|
ClientSessionID: clientSessionID,
|
||||||
Payload: payload,
|
Destination: dest,
|
||||||
|
Payload: payload,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -269,7 +204,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(ciphertext[:0], nonce, ciphertext, nil)
|
plain, err := c.chachaCipher.Open(nil, 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)
|
||||||
}
|
}
|
||||||
@@ -280,17 +215,22 @@ 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])
|
||||||
|
|
||||||
if c.sessions != nil {
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
sessionItem := c.sessions.GetOrCreate(sessionID)
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
sessionItem.Lock()
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
if !sessionItem.Window.CheckAndAdd(packetID) {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
|
||||||
}
|
|
||||||
sessionItem.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionItem.AddPacketID(packetID)
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// AES mode
|
// AES mode
|
||||||
@@ -299,54 +239,52 @@ 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])
|
||||||
|
|
||||||
var bodyAead cipher.AEAD
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
var sessionItem *ServerUDPSession
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
if c.sessions != nil {
|
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
|
||||||
sessionItem = c.sessions.GetOrCreate(sessionID)
|
}
|
||||||
sessionItem.Lock()
|
|
||||||
if !sessionItem.Window.Check(packetID) {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
|
||||||
}
|
|
||||||
sessionItem.Unlock()
|
|
||||||
|
|
||||||
bodyAead = sessionItem.GetRemoteCipher()
|
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
|
||||||
if bodyAead == nil {
|
bodyAead := s.clientBodyCipher
|
||||||
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
isNewCipher := false
|
||||||
var err error
|
if bodyAead == nil {
|
||||||
bodyAead, err = c.method.NewAEAD(bodyKey)
|
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
|
||||||
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 = c.method.NewAEAD(bodyKey)
|
bodyAead, err = 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]
|
||||||
bodyCipher := data[16:]
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
if sessionItem != nil {
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
sessionItem.Lock()
|
if err != nil {
|
||||||
sessionItem.Window.Add(packetID)
|
return DecodedUDPPacket{}, err
|
||||||
sessionItem.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
s.AddPacketID(packetID)
|
||||||
|
|
||||||
|
if isNewCipher {
|
||||||
|
s.clientBodyCipher = bodyAead
|
||||||
|
}
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
|
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
|
||||||
s.Lock()
|
s.Lock()
|
||||||
defer s.Unlock()
|
defer s.Unlock()
|
||||||
if s.ServerSessionID != 0 {
|
if s.ServerSessionID != 0 {
|
||||||
@@ -363,23 +301,29 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock c
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if method.IsChaCha {
|
if method.IsChaCha {
|
||||||
s.ServerChaCha = chachaCipher
|
var err error
|
||||||
} else {
|
s.serverChaCha, err = method.NewUDPCipher(psk)
|
||||||
s.ServerBlockCipher = headerBlock
|
return err
|
||||||
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
}
|
||||||
bodyAead, err := method.NewAEAD(bodyKey)
|
|
||||||
if err != nil {
|
var err error
|
||||||
s.ServerSessionID = 0
|
s.serverHeaderBlock, err = method.NewBlock(psk)
|
||||||
return err
|
if err != nil {
|
||||||
}
|
s.ServerSessionID = 0
|
||||||
s.ServerCipher = bodyAead
|
return err
|
||||||
|
}
|
||||||
|
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)
|
serverPacketID := s.ServerPacketID.Add(1) - 1
|
||||||
|
|
||||||
if method.IsChaCha {
|
if method.IsChaCha {
|
||||||
var nonce [PacketNonceSize]byte
|
var nonce [PacketNonceSize]byte
|
||||||
@@ -404,7 +348,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)
|
||||||
@@ -417,7 +361,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.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
|
||||||
bodyBuf := buf.New()
|
bodyBuf := buf.New()
|
||||||
defer bodyBuf.Release()
|
defer bodyBuf.Release()
|
||||||
@@ -435,7 +379,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
|||||||
bodyBuf.Write(payload)
|
bodyBuf.Write(payload)
|
||||||
|
|
||||||
bodyNonce := rawHeader[4:16]
|
bodyNonce := rawHeader[4:16]
|
||||||
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
||||||
|
|
||||||
res := make([]byte, 16+len(sealedBody))
|
res := make([]byte, 16+len(sealedBody))
|
||||||
copy(res[:16], encryptedHeader[:])
|
copy(res[:16], encryptedHeader[:])
|
||||||
@@ -444,17 +388,327 @@ 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) {
|
||||||
sessionItem := c.sessions.GetOrCreate(clientSessionID)
|
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
|
||||||
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
|
||||||
}
|
}
|
||||||
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
|
clientSessionID := binary.BigEndian.Uint64(sessID[:])
|
||||||
|
|
||||||
|
var clientBodyCipher cipher.AEAD
|
||||||
|
var err error
|
||||||
|
if !c.method.IsChaCha {
|
||||||
|
finalPSK := c.psk
|
||||||
|
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
|
||||||
|
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClientUDPSession{
|
||||||
|
codec: c,
|
||||||
|
clientSessionID: clientSessionID,
|
||||||
|
clientBodyCipher: clientBodyCipher,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
|
||||||
|
cur := s.current.Load()
|
||||||
|
if cur != nil && cur.sessionID == sessionID {
|
||||||
|
return cur, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
old := s.old.Load()
|
||||||
|
if old != nil && old.sessionID == sessionID {
|
||||||
|
if now-old.lastSeen.Load() > 60 {
|
||||||
|
s.old.CompareAndSwap(old, nil)
|
||||||
|
return nil, errors.New("old server session expired")
|
||||||
|
}
|
||||||
|
return old, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// New server session:
|
||||||
|
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
|
||||||
|
if old != nil && now-old.lastSeen.Load() < 60 {
|
||||||
|
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
|
||||||
|
}
|
||||||
|
|
||||||
|
var bodyAead cipher.AEAD
|
||||||
|
if !s.codec.method.IsChaCha {
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
|
||||||
|
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
newState := &serverSessionState{
|
||||||
|
sessionID: sessionID,
|
||||||
|
cipher: bodyAead,
|
||||||
|
}
|
||||||
|
newState.lastSeen.Store(now)
|
||||||
|
|
||||||
|
if cur == nil {
|
||||||
|
s.current.CompareAndSwap(nil, newState)
|
||||||
|
return s.current.Load(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s.old.Store(cur)
|
||||||
|
s.current.Store(newState)
|
||||||
|
return newState, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) ClientSessionID() uint64 {
|
||||||
|
return s.clientSessionID
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||||
|
packetID := s.nextPacketID.Add(1) - 1
|
||||||
|
sessID := s.clientSessionID
|
||||||
|
|
||||||
|
var paddingLen int
|
||||||
|
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
var nonce [PacketNonceSize]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(nonce[:])
|
||||||
|
|
||||||
|
var hdr [16 + 1 + 8 + 2]byte
|
||||||
|
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
||||||
|
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
||||||
|
hdr[16] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
||||||
|
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
|
||||||
|
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessID)
|
||||||
|
|
||||||
|
var rawHeader [16]byte
|
||||||
|
copy(rawHeader[:8], sessBytes[:])
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||||
|
|
||||||
|
eihCount := 0
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
eihCount = len(s.codec.pskList) - 1
|
||||||
|
}
|
||||||
|
|
||||||
|
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
|
||||||
|
for i := 0; i < len(s.codec.pskList)-1; i++ {
|
||||||
|
nextPSK := s.codec.pskList[i+1]
|
||||||
|
pskHash := DeriveUserPSKHash(nextPSK)
|
||||||
|
var eihPlain [16]byte
|
||||||
|
for k := 0; k < 16; k++ {
|
||||||
|
eihPlain[k] = pskHash[k] ^ rawHeader[k]
|
||||||
|
}
|
||||||
|
var encryptedEIH [16]byte
|
||||||
|
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
|
||||||
|
outBuf.Write(encryptedEIH[:])
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyAead := s.clientBodyCipher
|
||||||
|
|
||||||
|
var hdr [1 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
headerOffset := 16 + eihCount*16
|
||||||
|
plainBytes := outBuf.Bytes()[headerOffset:]
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||||
|
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||||
|
if len(data) < PacketMinimalHeaderSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
if len(data) < PacketNonceSize+AEADTagSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
nonce := data[:PacketNonceSize]
|
||||||
|
ciphertext := data[PacketNonceSize:]
|
||||||
|
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
|
}
|
||||||
|
if len(plain) < 16+1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var rawHeader [16]byte
|
||||||
|
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
|
||||||
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
bodyAead := st.cipher
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyCipher := data[16:]
|
||||||
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type UDPWriter struct {
|
type UDPWriter struct {
|
||||||
Writer io.Writer
|
Writer io.Writer
|
||||||
Destination net.Destination
|
Destination net.Destination
|
||||||
Codec *UDPPacketCodec
|
Session *ClientUDPSession
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
@@ -468,7 +722,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.Codec.EncodeClientPacket(dest, b.Bytes())
|
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
|
||||||
b.Release()
|
b.Release()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buf.ReleaseMulti(mb)
|
buf.ReleaseMulti(mb)
|
||||||
@@ -485,8 +739,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type UDPReader struct {
|
type UDPReader struct {
|
||||||
Reader io.Reader
|
Reader io.Reader
|
||||||
Codec *UDPPacketCodec
|
Session *ClientUDPSession
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
@@ -498,7 +752,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
|
decoded, err := r.Session.DecodePacket(buffer.Bytes())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buffer.Release()
|
buffer.Release()
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -2,9 +2,11 @@ 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"
|
||||||
@@ -269,3 +271,106 @@ 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))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,8 +6,11 @@ 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 (
|
||||||
@@ -74,30 +77,42 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
|
|||||||
|
|
||||||
type ServerUDPSession struct {
|
type ServerUDPSession struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
SessionID uint64
|
SessionID uint64
|
||||||
RemoteCipher atomic.Pointer[cipher.AEAD]
|
Window *SlidingWindow
|
||||||
Window SlidingWindow
|
User *protocol.MemoryUser
|
||||||
User *protocol.MemoryUser
|
UserPSK []byte
|
||||||
UserPSK []byte
|
LastActive atomic.Int64 // Unix timestamp in seconds
|
||||||
LastActive atomic.Int64 // Unix timestamp in seconds
|
|
||||||
|
clientBodyCipher cipher.AEAD
|
||||||
|
|
||||||
ServerSessionID uint64
|
ServerSessionID uint64
|
||||||
ServerPacketID atomic.Uint64
|
ServerPacketID atomic.Uint64
|
||||||
ServerCipher cipher.AEAD
|
serverBodyCipher cipher.AEAD
|
||||||
ServerBlockCipher cipher.Block
|
serverHeaderBlock 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) GetRemoteCipher() cipher.AEAD {
|
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
|
||||||
ptr := s.RemoteCipher.Load()
|
s.Lock()
|
||||||
if ptr == nil {
|
defer s.Unlock()
|
||||||
return nil
|
if s.Window == nil {
|
||||||
|
s.Window = new(SlidingWindow)
|
||||||
}
|
}
|
||||||
return *ptr
|
return s.Window.Check(packetID)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
|
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
|
||||||
s.RemoteCipher.Store(&c)
|
s.Lock()
|
||||||
|
defer s.Unlock()
|
||||||
|
if s.Window == nil {
|
||||||
|
s.Window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
s.Window.Add(packetID)
|
||||||
}
|
}
|
||||||
|
|
||||||
type UDPSessionManager struct {
|
type UDPSessionManager struct {
|
||||||
@@ -122,6 +137,7 @@ 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)
|
||||||
|
|
||||||
@@ -148,6 +164,7 @@ 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
|
||||||
})
|
})
|
||||||
@@ -156,3 +173,11 @@ 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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,18 +2,161 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type udpConnEntry struct {
|
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
|
||||||
sync.Mutex
|
if s.currentConn.Load() == nil {
|
||||||
link *transport.Link
|
s.currentConn.Store(conn)
|
||||||
timer *signal.ActivityTimer
|
}
|
||||||
cancel context.CancelFunc
|
if s.timer != nil {
|
||||||
|
s.timer.Update()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) WriteToClient(b []byte) error {
|
||||||
|
connVal := s.currentConn.Load()
|
||||||
|
if connVal == nil {
|
||||||
|
return errors.New("client connection closed")
|
||||||
|
}
|
||||||
|
conn, ok := connVal.(stat.Connection)
|
||||||
|
if !ok || conn == nil {
|
||||||
|
return errors.New("client connection closed")
|
||||||
|
}
|
||||||
|
_, err := conn.Write(b)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) Close() {
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.SetTimeout(0)
|
||||||
|
}
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
common.Interrupt(link.Reader)
|
||||||
|
common.Interrupt(link.Writer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EnsureLink(
|
||||||
|
ctx context.Context,
|
||||||
|
conn stat.Connection,
|
||||||
|
dest net.Destination,
|
||||||
|
dispatcher routing.Dispatcher,
|
||||||
|
policyManager policy.Manager,
|
||||||
|
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
|
||||||
|
) (*transport.Link, error) {
|
||||||
|
s.UpdateConn(conn)
|
||||||
|
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s.Lock()
|
||||||
|
defer s.Unlock()
|
||||||
|
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sessCtx, cancel := context.WithCancel(ctx)
|
||||||
|
inbound := session.InboundFromContext(sessCtx)
|
||||||
|
if inbound != nil && s.User != nil {
|
||||||
|
inbound.User = s.User
|
||||||
|
}
|
||||||
|
var email string
|
||||||
|
var level uint32
|
||||||
|
if s.User != nil {
|
||||||
|
email = s.User.Email
|
||||||
|
level = s.User.Level
|
||||||
|
}
|
||||||
|
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: email,
|
||||||
|
})
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(sessCtx, dest)
|
||||||
|
if err != nil {
|
||||||
|
cancel()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.link.Store(link)
|
||||||
|
sessionPolicy := policyManager.ForLevel(level)
|
||||||
|
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||||
|
if s.manager != nil {
|
||||||
|
s.manager.Delete(s.SessionID)
|
||||||
|
}
|
||||||
|
s.Close()
|
||||||
|
cancel()
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
go handleUDPResponse(s, link, dest, responseEncoder)
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
|
||||||
|
// when handshake or header validation fails.
|
||||||
|
func ResetTCPConn(conn net.Conn) {
|
||||||
|
rawConn, _, _ := proxy.UnwrapRawConn(conn)
|
||||||
|
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
|
||||||
|
_ = tcpConn.SetLinger(0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
|
||||||
|
defer func() {
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.SetTimeout(0)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
for {
|
||||||
|
resMb, err := link.Reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.Update()
|
||||||
|
}
|
||||||
|
for i, rb := range resMb {
|
||||||
|
b := rb.Bytes()
|
||||||
|
if encode != nil {
|
||||||
|
replyDest := fallbackDest
|
||||||
|
if rb.UDP != nil {
|
||||||
|
replyDest = *rb.UDP
|
||||||
|
}
|
||||||
|
encPacket, err := encode(replyDest, b)
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := s.WriteToClient(encPacket); err != nil {
|
||||||
|
buf.ReleaseMulti(resMb[i+1:])
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
err := s.WriteToClient(b)
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(resMb[i+1:])
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -182,57 +182,48 @@ func TestTCPStream(t *testing.T) {
|
|||||||
common.Must(err)
|
common.Must(err)
|
||||||
IncreaseNonce(reader.Nonce())
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
vBuf := buf.New()
|
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||||
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:]
|
||||||
|
|
||||||
// Skip padding
|
// Server sends response stream with receivedPayload as first payload
|
||||||
var padBytes [2]byte
|
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
|
||||||
_, _ = vBuf.Read(padBytes[:])
|
pBuf := buf.New()
|
||||||
padLen := int(padBytes[0])<<8 | int(padBytes[1])
|
pBuf.Write(receivedPayload)
|
||||||
vBuf.Advance(int32(padLen))
|
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||||
|
|
||||||
receivedPayload = make([]byte, vBuf.Len())
|
// Read and echo additional stream data
|
||||||
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, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
|
clientSalt := make([]byte, method.KeySaltLength)
|
||||||
|
common.Must2(io.ReadFull(rand.Reader, clientSalt))
|
||||||
|
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
|
||||||
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
|
reader, err := ReadTCPResponse(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.WriteChunk(streamData)
|
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
|
||||||
|
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
@@ -272,12 +263,14 @@ 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, psk)
|
clientCodec, err := NewUDPPacketCodec(method, [][]byte{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)
|
||||||
|
|
||||||
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
|
session, err := clientCodec.NewClientSession()
|
||||||
|
common.Must(err)
|
||||||
|
pktBuf, err := session.EncodePacket(dest, payload)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
defer pktBuf.Release()
|
defer pktBuf.Release()
|
||||||
|
|
||||||
@@ -360,3 +353,146 @@ 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")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+235
-115
@@ -1,18 +1,25 @@
|
|||||||
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(
|
||||||
@@ -38,15 +45,6 @@ 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() {
|
||||||
@@ -119,8 +117,16 @@ 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 {
|
||||||
if err := w.WriteChunk(b.Bytes()); err != nil {
|
p := b.Bytes()
|
||||||
return err
|
for len(p) > 0 {
|
||||||
|
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
|
||||||
@@ -168,7 +174,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 {
|
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||||
return 0, ErrInvalidRequest
|
return 0, ErrInvalidRequest
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -194,11 +200,10 @@ 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 {
|
||||||
b := buf.New()
|
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
|
||||||
b.Write(r.buffer[r.offset : r.offset+r.cached])
|
|
||||||
r.cached = 0
|
r.cached = 0
|
||||||
r.offset = 0
|
r.offset = 0
|
||||||
return buf.MultiBuffer{b}, nil
|
return mb, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||||
@@ -212,7 +217,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 {
|
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||||
return nil, ErrInvalidRequest
|
return nil, ErrInvalidRequest
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -227,9 +232,8 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
}
|
}
|
||||||
IncreaseNonce(r.nonce[:])
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
b := buf.New()
|
mb := buf.MergeBytes(nil, decryptedPayload)
|
||||||
b.Write(decryptedPayload)
|
return mb, nil
|
||||||
return buf.MultiBuffer{b}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ClientRequestHeader struct {
|
type ClientRequestHeader struct {
|
||||||
@@ -237,13 +241,8 @@ type ClientRequestHeader struct {
|
|||||||
EarlyData []byte
|
EarlyData []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
|
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
|
||||||
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
|
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -272,7 +271,7 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
} else {
|
} else {
|
||||||
varChunkCipher = make([]byte, needed)
|
varChunkCipher = make([]byte, needed)
|
||||||
}
|
}
|
||||||
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
|
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -282,31 +281,34 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
}
|
}
|
||||||
IncreaseNonce(reader.Nonce())
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
b := buf.New()
|
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||||
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
|
||||||
|
|
||||||
var padLenBytes [2]byte
|
offset := addrLen
|
||||||
if _, err := b.Read(padLenBytes[:]); err != nil {
|
if len(plainVar) < offset+2 {
|
||||||
return nil, err
|
return nil, ErrPacketTooShort
|
||||||
}
|
}
|
||||||
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
|
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
|
||||||
if int(b.Len()) < paddingLen {
|
offset += 2
|
||||||
|
|
||||||
|
if len(plainVar) < offset+paddingLen {
|
||||||
return nil, ErrNoPadding
|
return nil, ErrNoPadding
|
||||||
}
|
}
|
||||||
if paddingLen > 0 {
|
offset += paddingLen
|
||||||
b.Advance(int32(paddingLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
var earlyData []byte
|
var earlyData []byte
|
||||||
if b.Len() > 0 {
|
var payloadLen int
|
||||||
earlyData = make([]byte, b.Len())
|
if len(plainVar) > offset {
|
||||||
copy(earlyData, b.Bytes())
|
earlyData = plainVar[offset:]
|
||||||
|
payloadLen = len(earlyData)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0.
|
||||||
|
if paddingLen == 0 && payloadLen == 0 {
|
||||||
|
return nil, errors.New("request without payload and padding is not allowed")
|
||||||
}
|
}
|
||||||
|
|
||||||
return &ClientRequestHeader{
|
return &ClientRequestHeader{
|
||||||
@@ -315,34 +317,6 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
}, 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]
|
||||||
@@ -354,7 +328,16 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
|
|
||||||
writer := NewStreamWriter(w, aead)
|
writer := NewStreamWriter(w, aead)
|
||||||
|
|
||||||
handshakeBuf := buf.New()
|
payloadLen := len(payload)
|
||||||
|
var paddingLen int
|
||||||
|
if payloadLen < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||||
|
}
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
|
||||||
|
|
||||||
|
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
|
||||||
|
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
|
||||||
defer handshakeBuf.Release()
|
defer handshakeBuf.Release()
|
||||||
|
|
||||||
handshakeBuf.Write(clientSalt)
|
handshakeBuf.Write(clientSalt)
|
||||||
@@ -372,14 +355,6 @@ 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()))
|
||||||
@@ -389,7 +364,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.New()
|
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
|
||||||
defer varHeaderBuf.Release()
|
defer varHeaderBuf.Release()
|
||||||
|
|
||||||
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
||||||
@@ -421,12 +396,21 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
|
|
||||||
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
|
// 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) {
|
||||||
var serverSalt [32]byte
|
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
|
||||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
chunkCipherLen := fixedPlainLen + AEADTagSize
|
||||||
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
|
headerLen := method.KeySaltLength + chunkCipherLen
|
||||||
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 {
|
||||||
@@ -435,14 +419,6 @@ 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)
|
||||||
@@ -484,46 +460,190 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
|
|||||||
return reader, nil
|
return reader, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
|
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
|
||||||
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
|
type ServerStreamWriter struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
w io.Writer
|
||||||
|
method *CipherMethod
|
||||||
|
psk []byte
|
||||||
|
clientSalt []byte
|
||||||
|
streamWriter *StreamWriter
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
|
||||||
|
return &ServerStreamWriter{
|
||||||
|
w: w,
|
||||||
|
method: method,
|
||||||
|
psk: psk,
|
||||||
|
clientSalt: clientSalt,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
|
||||||
var serverSalt [32]byte
|
var serverSalt [32]byte
|
||||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
serverSaltSlice := serverSalt[:s.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(psk, serverSaltSlice, method.KeySaltLength)
|
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
|
||||||
respAead, err := method.NewAEAD(respKey)
|
respAead, err := s.method.NewAEAD(respKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
writer := NewStreamWriter(w, respAead)
|
sw := NewStreamWriter(s.w, respAead)
|
||||||
|
|
||||||
respBuf := buf.New()
|
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
|
||||||
defer respBuf.Release()
|
outBuf := buf.NewWithSize(totalHeaderLen)
|
||||||
|
defer outBuf.Release()
|
||||||
|
|
||||||
respBuf.Write(serverSaltSlice)
|
outBuf.Write(serverSaltSlice)
|
||||||
|
|
||||||
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
||||||
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
|
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
|
||||||
fixedRespSlice[0] = HeaderTypeServer
|
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+method.KeySaltLength], clientSalt)
|
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
|
||||||
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
|
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
|
||||||
|
|
||||||
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
|
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
|
||||||
IncreaseNonce(writer.nonce[:])
|
IncreaseNonce(sw.nonce[:])
|
||||||
respBuf.Write(fixedRespChunk)
|
outBuf.Write(fixedRespChunk)
|
||||||
|
|
||||||
if len(initialPayload) > 0 {
|
if len(payload) > 0 {
|
||||||
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
|
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
|
||||||
IncreaseNonce(writer.nonce[:])
|
IncreaseNonce(sw.nonce[:])
|
||||||
respBuf.Write(initialChunk)
|
outBuf.Write(payloadChunk)
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := w.Write(respBuf.Bytes()); err != nil {
|
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
return sw, nil
|
||||||
return writer, nil
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
if mb.IsEmpty() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
firstBuf := mb[0]
|
||||||
|
firstBytes := firstBuf.Bytes()
|
||||||
|
chunkSize := len(firstBytes)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
firstPayload := firstBytes[:chunkSize]
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||||
|
if err != nil {
|
||||||
|
s.mu.Unlock()
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
|
||||||
|
firstBuf.Advance(int32(chunkSize))
|
||||||
|
if firstBuf.IsEmpty() {
|
||||||
|
firstBuf.Release()
|
||||||
|
mb = mb[1:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if len(mb) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.streamWriter.WriteMultiBuffer(mb)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
|
||||||
|
n := len(p)
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
chunkSize := len(p)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
firstPayload := p[:chunkSize]
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||||
|
if err != nil {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
p = p[chunkSize:]
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if len(p) == 0 {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := s.streamWriter.Write(p)
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) Close() error {
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
|
||||||
|
// and returns a StreamReader for subsequent stream chunks.
|
||||||
|
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
|
||||||
|
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := NewStreamReader(conn, aead)
|
||||||
|
|
||||||
|
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
_ = conn.SetReadDeadline(time.Time{})
|
||||||
|
|
||||||
|
if !saltFilter.Check(salt) {
|
||||||
|
return nil, nil, ErrSaltNotUnique
|
||||||
|
}
|
||||||
|
return reader, reqHeader, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
if c, ok := writer.(io.Closer); ok {
|
||||||
|
defer c.Close()
|
||||||
|
}
|
||||||
|
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user