diff --git a/proxy/shadowsocks_2022/inbound.go b/proxy/shadowsocks_2022/inbound.go index eaa637d30..ca07fa3ae 100644 --- a/proxy/shadowsocks_2022/inbound.go +++ b/proxy/shadowsocks_2022/inbound.go @@ -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 } } diff --git a/proxy/shadowsocks_2022/inbound_multi.go b/proxy/shadowsocks_2022/inbound_multi.go index 45df40e53..f46bf72a7 100644 --- a/proxy/shadowsocks_2022/inbound_multi.go +++ b/proxy/shadowsocks_2022/inbound_multi.go @@ -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 { diff --git a/proxy/shadowsocks_2022/outbound.go b/proxy/shadowsocks_2022/outbound.go index 71d21475d..ff1be9f06 100644 --- a/proxy/shadowsocks_2022/outbound.go +++ b/proxy/shadowsocks_2022/outbound.go @@ -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) } diff --git a/proxy/shadowsocks_2022/packet.go b/proxy/shadowsocks_2022/packet.go index f0d1daba7..fc906aeb8 100644 --- a/proxy/shadowsocks_2022/packet.go +++ b/proxy/shadowsocks_2022/packet.go @@ -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) diff --git a/proxy/shadowsocks_2022/shadowsocks_2022_test.go b/proxy/shadowsocks_2022/shadowsocks_2022_test.go index f6f930a27..bcec78d50 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022_test.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022_test.go @@ -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") + } +} diff --git a/proxy/shadowsocks_2022/stream.go b/proxy/shadowsocks_2022/stream.go index 3e21734e7..14afd66dc 100644 --- a/proxy/shadowsocks_2022/stream.go +++ b/proxy/shadowsocks_2022/stream.go @@ -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{