Fix ss2022 big buffer && udp mult user

This commit is contained in:
Fangliding
2026-09-29 18:47:24 +08:00
parent e5e85ca9da
commit 7dc35bac94
6 changed files with 157 additions and 54 deletions
+2 -3
View File
@@ -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
}
}
+6 -11
View File
@@ -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 {
+1 -1
View File
@@ -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)
}
+56 -12
View File
@@ -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")
}
}
+28 -26
View File
@@ -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{