diff --git a/mtglib/internal/faketls/record/pools.go b/mtglib/internal/faketls/record/pools.go index 62f03e9..50b0fec 100644 --- a/mtglib/internal/faketls/record/pools.go +++ b/mtglib/internal/faketls/record/pools.go @@ -1,12 +1,22 @@ package record -import "sync" +import ( + "bytes" + "sync" +) -var recordPool = sync.Pool{ - New: func() interface{} { - return &Record{} - }, -} +var ( + recordPool = sync.Pool{ + New: func() interface{} { + return &Record{} + }, + } + bytesBufferPool = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) func AcquireRecord() *Record { return recordPool.Get().(*Record) @@ -16,3 +26,12 @@ func ReleaseRecord(r *Record) { r.Reset() recordPool.Put(r) } + +func acquireBytesBuffer() *bytes.Buffer { + return bytesBufferPool.Get().(*bytes.Buffer) +} + +func releaseBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + bytesBufferPool.Put(buf) +} diff --git a/mtglib/internal/faketls/record/record.go b/mtglib/internal/faketls/record/record.go index 31a1d9e..bce03b0 100644 --- a/mtglib/internal/faketls/record/record.go +++ b/mtglib/internal/faketls/record/record.go @@ -61,26 +61,22 @@ func (r *Record) Read(reader io.Reader) error { } func (r *Record) Dump(writer io.Writer) error { - buf := [2]byte{byte(r.Type), 0} + buf := acquireBytesBuffer() + defer releaseBytesBuffer(buf) - if _, err := writer.Write(buf[:1]); err != nil { - return fmt.Errorf("cannot dump type: %w", err) - } + bufSlice := [2]byte{byte(r.Type), 0} + buf.Write(bufSlice[:1]) - binary.BigEndian.PutUint16(buf[:], uint16(r.Version)) + binary.BigEndian.PutUint16(bufSlice[:], uint16(r.Version)) + buf.Write(bufSlice[:]) - if _, err := writer.Write(buf[:]); err != nil { - return fmt.Errorf("cannot dump version: %w", err) - } + binary.BigEndian.PutUint16(bufSlice[:], uint16(r.Payload.Len())) + buf.Write(bufSlice[:]) - binary.BigEndian.PutUint16(buf[:], uint16(r.Payload.Len())) + buf.Write(r.Payload.Bytes()) - if _, err := writer.Write(buf[:]); err != nil { - return fmt.Errorf("cannot dump payload length: %w", err) - } - - if _, err := writer.Write(r.Payload.Bytes()); err != nil { - return fmt.Errorf("cannot dump payload: %w", err) + if _, err := buf.WriteTo(writer); err != nil { + return fmt.Errorf("cannot dump record: %w", err) } return nil diff --git a/mtglib/internal/obfuscated2/client_handshake_test.go b/mtglib/internal/obfuscated2/client_handshake_test.go index 3edbdba..ef7402a 100644 --- a/mtglib/internal/obfuscated2/client_handshake_test.go +++ b/mtglib/internal/obfuscated2/client_handshake_test.go @@ -58,7 +58,7 @@ func (suite *ClientHandshakeTestSuite) TestOk() { copy(writeData, arr) }) - conn := &obfuscated2.Conn{ + conn := obfuscated2.Conn{ Conn: connMock, Encryptor: encryptor, Decryptor: decryptor, diff --git a/mtglib/internal/obfuscated2/conn.go b/mtglib/internal/obfuscated2/conn.go index 24b5a81..c3a4c05 100644 --- a/mtglib/internal/obfuscated2/conn.go +++ b/mtglib/internal/obfuscated2/conn.go @@ -10,11 +10,9 @@ type Conn struct { Encryptor cipher.Stream Decryptor cipher.Stream - - writeBuf []byte } -func (c *Conn) Read(p []byte) (int, error) { +func (c Conn) Read(p []byte) (int, error) { n, err := c.Conn.Read(p) if err != nil { return n, err // nolint: wrapcheck @@ -25,9 +23,16 @@ func (c *Conn) Read(p []byte) (int, error) { return n, nil } -func (c *Conn) Write(p []byte) (int, error) { - c.writeBuf = append(c.writeBuf[:0], p...) - c.Encryptor.XORKeyStream(c.writeBuf, p) +func (c Conn) Write(p []byte) (int, error) { + buf := acquireBytesBuffer() + defer releaseBytesBuffer(buf) - return c.Conn.Write(c.writeBuf) + buf.Write(p) + + payload := buf.Bytes() + c.Encryptor.XORKeyStream(payload, payload) + + n, err := buf.WriteTo(c.Conn) + + return int(n), err // nolint: wrapcheck } diff --git a/mtglib/internal/obfuscated2/pools.go b/mtglib/internal/obfuscated2/pools.go index ce3204f..811c540 100644 --- a/mtglib/internal/obfuscated2/pools.go +++ b/mtglib/internal/obfuscated2/pools.go @@ -1,16 +1,24 @@ package obfuscated2 import ( + "bytes" "crypto/sha256" "hash" "sync" ) -var sha256HasherPool = sync.Pool{ - New: func() interface{} { - return sha256.New() - }, -} +var ( + sha256HasherPool = sync.Pool{ + New: func() interface{} { + return sha256.New() + }, + } + bytesBufferPool = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) func acquireSha256Hasher() hash.Hash { return sha256HasherPool.Get().(hash.Hash) @@ -20,3 +28,12 @@ func releaseSha256Hasher(h hash.Hash) { h.Reset() sha256HasherPool.Put(h) } + +func acquireBytesBuffer() *bytes.Buffer { + return bytesBufferPool.Get().(*bytes.Buffer) +} + +func releaseBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + bytesBufferPool.Put(buf) +} diff --git a/mtglib/internal/obfuscated2/server_handshake_test.go b/mtglib/internal/obfuscated2/server_handshake_test.go index 3638369..9279295 100644 --- a/mtglib/internal/obfuscated2/server_handshake_test.go +++ b/mtglib/internal/obfuscated2/server_handshake_test.go @@ -17,7 +17,7 @@ type ServerHandshakeTestSuite struct { suite.Suite connMock *testlib.NetConnMock - proxyConn *obfuscated2.Conn + proxyConn obfuscated2.Conn encryptor cipher.Stream decryptor cipher.Stream } @@ -29,7 +29,7 @@ func (suite *ServerHandshakeTestSuite) SetupTest() { encryptor, decryptor, err := obfuscated2.ServerHandshake(buf) suite.NoError(err) - suite.proxyConn = &obfuscated2.Conn{ + suite.proxyConn = obfuscated2.Conn{ Conn: suite.connMock, Encryptor: encryptor, Decryptor: decryptor, diff --git a/mtglib/proxy.go b/mtglib/proxy.go index 50b038c..6baf4a2 100644 --- a/mtglib/proxy.go +++ b/mtglib/proxy.go @@ -192,7 +192,7 @@ func (p *Proxy) doObfuscated2Handshake(ctx *streamContext) error { ctx.dc = dc ctx.logger = ctx.logger.BindInt("dc", dc) - ctx.clientConn = &obfuscated2.Conn{ + ctx.clientConn = obfuscated2.Conn{ Conn: ctx.clientConn, Encryptor: encryptor, Decryptor: decryptor, @@ -214,7 +214,7 @@ func (p *Proxy) doTelegramCall(ctx *streamContext) error { return fmt.Errorf("cannot perform obfuscated2 handshake: %w", err) } - ctx.telegramConn = &obfuscated2.Conn{ + ctx.telegramConn = obfuscated2.Conn{ Conn: connTelegramTraffic{ Conn: conn, connID: ctx.connID,