mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-30 10:57:57 +03:00
Fix ss2022 big buffer && udp mult user
This commit is contained in:
@@ -146,9 +146,8 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
||||
}
|
||||
|
||||
if len(reqHeader.EarlyData) > 0 {
|
||||
earlyBuf := buf.New()
|
||||
earlyBuf.Write(reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -283,9 +283,8 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
||||
}
|
||||
|
||||
if len(reqHeader.EarlyData) > 0 {
|
||||
earlyBuf := buf.New()
|
||||
earlyBuf.Write(reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -361,15 +360,11 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
|
||||
} 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
|
||||
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
|
||||
i.udpMasterCipher.Decrypt(decryptedHash[:], packetBytes[16:32])
|
||||
for k := 0; k < 16; k++ {
|
||||
decryptedHash[k] ^= rawHeader[k]
|
||||
}
|
||||
|
||||
user, ok := i.usersByHash.Load(decryptedHash)
|
||||
if !ok || user == nil {
|
||||
|
||||
@@ -47,7 +47,7 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
||||
}
|
||||
|
||||
finalPSK := pskList[len(pskList)-1]
|
||||
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
|
||||
udpCodec, err := NewUDPPacketCodec(method, pskList)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to create udp packet codec").Base(err)
|
||||
}
|
||||
|
||||
@@ -17,8 +17,10 @@ import (
|
||||
|
||||
type UDPCodec struct {
|
||||
method *CipherMethod
|
||||
pskList [][]byte
|
||||
psk []byte
|
||||
blockCipher cipher.Block
|
||||
blockCiphers []cipher.Block
|
||||
chachaCipher cipher.AEAD
|
||||
clientBodyCipher cipher.AEAD
|
||||
clientSessionID uint64
|
||||
@@ -48,11 +50,23 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||
c, err := newUDPCodec(method, psk)
|
||||
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
|
||||
finalPSK := pskList[len(pskList)-1]
|
||||
c, err := newUDPCodec(method, finalPSK)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.pskList = pskList
|
||||
if len(pskList) > 1 && !method.IsChaCha {
|
||||
c.blockCiphers = make([]cipher.Block, len(pskList))
|
||||
for i, psk := range pskList {
|
||||
c.blockCiphers[i], err = method.NewBlock(psk)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var sessID [8]byte
|
||||
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||
return nil, err
|
||||
@@ -60,7 +74,7 @@ func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
|
||||
|
||||
if !method.IsChaCha {
|
||||
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
|
||||
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], method.KeySaltLength)
|
||||
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -130,21 +144,50 @@ func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*bu
|
||||
}
|
||||
|
||||
// AES mode:
|
||||
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
|
||||
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||
var sessBytes [8]byte
|
||||
binary.BigEndian.PutUint64(sessBytes[:], sessID)
|
||||
|
||||
var rawHeader [16]byte
|
||||
copy(rawHeader[:8], sessBytes[:])
|
||||
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||
|
||||
eihCount := 0
|
||||
if len(c.pskList) > 1 {
|
||||
eihCount = len(c.pskList) - 1
|
||||
}
|
||||
|
||||
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||
if totalLen > buf.Size {
|
||||
return nil, ErrPacketTooLarge
|
||||
}
|
||||
|
||||
outBuf := buf.New()
|
||||
|
||||
var rawHeader [16]byte
|
||||
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
|
||||
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||
if len(c.pskList) > 1 {
|
||||
// Multi-user / Relay mode:
|
||||
// 1. Header (16B) encrypted with first hop's block cipher
|
||||
var encryptedHeader [16]byte
|
||||
c.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
outBuf.Write(encryptedHeader[:])
|
||||
|
||||
var encryptedHeader [16]byte
|
||||
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
outBuf.Write(encryptedHeader[:])
|
||||
// 2. Multi-hop EIHs for intermediate hops
|
||||
for i := 0; i < len(c.pskList)-1; i++ {
|
||||
nextPSK := c.pskList[i+1]
|
||||
pskHash := DeriveUserPSKHash(nextPSK)
|
||||
var eihPlain [16]byte
|
||||
for k := 0; k < 16; k++ {
|
||||
eihPlain[k] = pskHash[k] ^ rawHeader[k]
|
||||
}
|
||||
var encryptedEIH [16]byte
|
||||
c.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
|
||||
outBuf.Write(encryptedEIH[:])
|
||||
}
|
||||
} else {
|
||||
// Single-user mode:
|
||||
var encryptedHeader [16]byte
|
||||
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
outBuf.Write(encryptedHeader[:])
|
||||
}
|
||||
|
||||
bodyAead := c.clientBodyCipher
|
||||
|
||||
@@ -163,7 +206,8 @@ func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*bu
|
||||
}
|
||||
outBuf.Write(payload)
|
||||
|
||||
plainBytes := outBuf.Bytes()[16:]
|
||||
headerOffset := 16 + eihCount*16
|
||||
plainBytes := outBuf.Bytes()[headerOffset:]
|
||||
bodyNonce := rawHeader[4:16]
|
||||
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||
|
||||
@@ -272,7 +272,7 @@ func TestUDPCodec(t *testing.T) {
|
||||
psk := make([]byte, method.KeySaltLength)
|
||||
_, _ = rand.Read(psk)
|
||||
|
||||
clientCodec, err := NewUDPPacketCodec(method, psk)
|
||||
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
|
||||
common.Must(err)
|
||||
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
|
||||
common.Must(err)
|
||||
@@ -360,3 +360,66 @@ func TestMultiUserManager(t *testing.T) {
|
||||
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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -119,8 +119,16 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
|
||||
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
defer buf.ReleaseMulti(mb)
|
||||
for _, b := range mb {
|
||||
if err := w.WriteChunk(b.Bytes()); err != nil {
|
||||
return err
|
||||
p := b.Bytes()
|
||||
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
|
||||
@@ -168,7 +176,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||
if payloadLen == 0 {
|
||||
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||
return 0, ErrInvalidRequest
|
||||
}
|
||||
|
||||
@@ -194,11 +202,10 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
||||
|
||||
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
if r.cached > 0 {
|
||||
b := buf.New()
|
||||
b.Write(r.buffer[r.offset : r.offset+r.cached])
|
||||
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
|
||||
r.cached = 0
|
||||
r.offset = 0
|
||||
return buf.MultiBuffer{b}, nil
|
||||
return mb, nil
|
||||
}
|
||||
|
||||
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||
@@ -212,7 +219,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||
if payloadLen == 0 {
|
||||
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||
return nil, ErrInvalidRequest
|
||||
}
|
||||
|
||||
@@ -227,9 +234,8 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
}
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
b := buf.New()
|
||||
b.Write(decryptedPayload)
|
||||
return buf.MultiBuffer{b}, nil
|
||||
mb := buf.MergeBytes(nil, decryptedPayload)
|
||||
return mb, nil
|
||||
}
|
||||
|
||||
type ClientRequestHeader struct {
|
||||
@@ -282,31 +288,27 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
||||
}
|
||||
IncreaseNonce(reader.Nonce())
|
||||
|
||||
b := buf.New()
|
||||
b.Write(plainVar)
|
||||
defer b.Release()
|
||||
|
||||
dest, err := ReadAddressPort(b)
|
||||
dest, addrLen, err := parseAddressPort(plainVar)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dest.Network = net.Network_TCP
|
||||
|
||||
var padLenBytes [2]byte
|
||||
if _, err := b.Read(padLenBytes[:]); err != nil {
|
||||
return nil, err
|
||||
offset := addrLen
|
||||
if len(plainVar) < offset+2 {
|
||||
return nil, ErrPacketTooShort
|
||||
}
|
||||
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
|
||||
if int(b.Len()) < paddingLen {
|
||||
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
|
||||
offset += 2
|
||||
|
||||
if len(plainVar) < offset+paddingLen {
|
||||
return nil, ErrNoPadding
|
||||
}
|
||||
if paddingLen > 0 {
|
||||
b.Advance(int32(paddingLen))
|
||||
}
|
||||
offset += paddingLen
|
||||
|
||||
var earlyData []byte
|
||||
if b.Len() > 0 {
|
||||
earlyData = make([]byte, b.Len())
|
||||
copy(earlyData, b.Bytes())
|
||||
if len(plainVar) > offset {
|
||||
earlyData = plainVar[offset:]
|
||||
}
|
||||
|
||||
return &ClientRequestHeader{
|
||||
|
||||
Reference in New Issue
Block a user