diff --git a/faketls/client_protocol.go b/faketls/client_protocol.go index 1f6a207..135c422 100644 --- a/faketls/client_protocol.go +++ b/faketls/client_protocol.go @@ -63,7 +63,12 @@ func (c *ClientProtocol) tlsHandshake(conn io.ReadWriter) error { return fmt.Errorf("cannot read initial record: %w", err) } - clientHello, err := tlstypes.ParseClientHello(helloRecord.Data.Bytes()) + buf := acquireBytesBuffer() + defer releaseBytesBuffer(buf) + + helloRecord.Data.WriteBytes(buf) + + clientHello, err := tlstypes.ParseClientHello(buf.Bytes()) if err != nil { return fmt.Errorf("cannot parse client hello: %w", err) } diff --git a/faketls/cloak.go b/faketls/cloak.go index c33d06f..d9a2cc2 100644 --- a/faketls/cloak.go +++ b/faketls/cloak.go @@ -28,15 +28,9 @@ func cloak(one, another io.ReadWriteCloser) { wg.Add(2) - go func() { - defer wg.Done() - io.Copy(one, another) // nolint: errcheck - }() + go cloakPipe(one, another, wg) - go func() { - defer wg.Done() - io.Copy(another, one) // nolint: errcheck - }() + go cloakPipe(another, one, wg) go func() { wg.Wait() @@ -69,3 +63,12 @@ func cloak(one, another io.ReadWriteCloser) { <-ctx.Done() } + +func cloakPipe(one io.Writer, another io.Reader, wg *sync.WaitGroup) { + defer wg.Done() + + buf := acquireCloakBuffer() + defer releaseCloakBuffer(buf) + + io.CopyBuffer(one, another, *buf) // nolint: errcheck +} diff --git a/faketls/pools.go b/faketls/pools.go new file mode 100644 index 0000000..ca7b6c4 --- /dev/null +++ b/faketls/pools.go @@ -0,0 +1,39 @@ +package faketls + +import ( + "bytes" + "sync" +) + +const cloakBufferSize = 1024 + +var ( + poolBytesBuffer = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } + poolCloakBuffer = sync.Pool{ + New: func() interface{} { + rv := make([]byte, cloakBufferSize) + return &rv + }, + } +) + +func acquireBytesBuffer() *bytes.Buffer { + return poolBytesBuffer.Get().(*bytes.Buffer) +} + +func acquireCloakBuffer() *[]byte { + return poolCloakBuffer.Get().(*[]byte) +} + +func releaseBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + poolBytesBuffer.Put(buf) +} + +func releaseCloakBuffer(buf *[]byte) { + poolCloakBuffer.Put(buf) +} diff --git a/tlstypes/client_hello.go b/tlstypes/client_hello.go index dbdcde4..ffb6e87 100644 --- a/tlstypes/client_hello.go +++ b/tlstypes/client_hello.go @@ -25,7 +25,7 @@ func (c ClientHello) Digest() []byte { } mac := hmac.New(sha256.New, config.C.Secret) - mac.Write(rec.Bytes()) // nolint: errcheck + rec.WriteBytes(mac) computedDigest := mac.Sum(nil) for i := range computedDigest { diff --git a/tlstypes/consts.go b/tlstypes/consts.go index 72e6935..6e4156a 100644 --- a/tlstypes/consts.go +++ b/tlstypes/consts.go @@ -1,5 +1,7 @@ package tlstypes +import "io" + type RecordType uint8 const ( @@ -69,11 +71,16 @@ var ( ) type Byter interface { - Bytes() []byte + WriteBytes(io.Writer) + Len() int } type RawBytes []byte -func (r RawBytes) Bytes() []byte { - return []byte(r) +func (r RawBytes) WriteBytes(writer io.Writer) { + writer.Write(r) // nolint: errcheck +} + +func (r RawBytes) Len() int { + return len(r) } diff --git a/tlstypes/handshake.go b/tlstypes/handshake.go index ec0accf..f70fc92 100644 --- a/tlstypes/handshake.go +++ b/tlstypes/handshake.go @@ -1,7 +1,7 @@ package tlstypes import ( - "bytes" + "io" "github.com/9seconds/mtg/utils" ) @@ -14,24 +14,31 @@ type Handshake struct { Tail Byter } -func (h *Handshake) Bytes() []byte { - buf := bytes.Buffer{} - packetBuf := bytes.Buffer{} +func (h *Handshake) WriteBytes(writer io.Writer) { + packetBuf := acquireBytesBuffer() + defer releaseBytesBuffer(packetBuf) - buf.WriteByte(byte(h.Type)) + writer.Write([]byte{byte(h.Type)}) // nolint: errcheck packetBuf.Write(h.Version.Bytes()) packetBuf.Write(h.Random[:]) packetBuf.WriteByte(byte(len(h.SessionID))) packetBuf.Write(h.SessionID) - packetBuf.Write(h.Tail.Bytes()) + h.Tail.WriteBytes(packetBuf) sizeUint24 := utils.ToUint24(uint32(packetBuf.Len())) sizeUint24Bytes := sizeUint24[:] sizeUint24Bytes[0], sizeUint24Bytes[2] = sizeUint24Bytes[2], sizeUint24Bytes[0] - buf.Write(sizeUint24Bytes) - packetBuf.WriteTo(&buf) // nolint: errcheck + writer.Write(sizeUint24Bytes) // nolint: errcheck + packetBuf.WriteTo(writer) // nolint: errcheck +} - return buf.Bytes() +func (h *Handshake) Len() int { + buf := acquireBytesBuffer() + defer releaseBytesBuffer(buf) + + h.WriteBytes(buf) + + return buf.Len() } diff --git a/tlstypes/pools.go b/tlstypes/pools.go new file mode 100644 index 0000000..f51aae7 --- /dev/null +++ b/tlstypes/pools.go @@ -0,0 +1,23 @@ +package tlstypes + +import ( + "bytes" + "sync" +) + +var ( + poolBytesBuffer = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) + +func acquireBytesBuffer() *bytes.Buffer { + return poolBytesBuffer.Get().(*bytes.Buffer) +} + +func releaseBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + poolBytesBuffer.Put(buf) +} diff --git a/tlstypes/record.go b/tlstypes/record.go index d6a71dd..5dae77e 100644 --- a/tlstypes/record.go +++ b/tlstypes/record.go @@ -15,16 +15,15 @@ type Record struct { Data Byter } -func (r Record) Bytes() []byte { - buf := bytes.Buffer{} - data := r.Data.Bytes() +func (r Record) WriteBytes(writer io.Writer) { + writer.Write([]byte{byte(r.Type)}) // nolint: errcheck + writer.Write(r.Version.Bytes()) // nolint: errcheck + binary.Write(writer, binary.BigEndian, uint16(r.Data.Len())) // nolint: errcheck + r.Data.WriteBytes(writer) +} - buf.WriteByte(byte(r.Type)) - buf.Write(r.Version.Bytes()) - binary.Write(&buf, binary.BigEndian, uint16(len(data))) // nolint: errcheck - buf.Write(data) - - return buf.Bytes() +func (r Record) Len() int { + return 1 + 2 + 2 + r.Data.Len() } func ReadRecord(reader io.Reader) (Record, error) { diff --git a/tlstypes/server_hello.go b/tlstypes/server_hello.go index 3c9249c..a92325c 100644 --- a/tlstypes/server_hello.go +++ b/tlstypes/server_hello.go @@ -20,20 +20,22 @@ type ServerHello struct { } func (s ServerHello) WelcomePacket() []byte { + buf := &bytes.Buffer{} + s.Random = [32]byte{} rec := Record{ Type: RecordTypeHandshake, Version: Version12, Data: &s, } - buf := bytes.NewBuffer(rec.Bytes()) + rec.WriteBytes(buf) recChangeCipher := Record{ Type: RecordTypeChangeCipherSpec, Version: Version12, Data: RawBytes([]byte{0x01}), } - buf.Write(recChangeCipher.Bytes()) + recChangeCipher.WriteBytes(buf) hostCert := make([]byte, 1024+mrand.Intn(3092)) rand.Read(hostCert) // nolint: errcheck @@ -43,7 +45,8 @@ func (s ServerHello) WelcomePacket() []byte { Version: Version12, Data: RawBytes(hostCert), } - buf.Write(recData.Bytes()) + recData.WriteBytes(buf) + packet := buf.Bytes() mac := hmac.New(sha256.New, config.C.Secret) diff --git a/wrappers/packet/mtproto_frame.go b/wrappers/packet/mtproto_frame.go index 6a1cd1a..6190d39 100644 --- a/wrappers/packet/mtproto_frame.go +++ b/wrappers/packet/mtproto_frame.go @@ -42,7 +42,9 @@ type wrapperMtprotoFrame struct { } func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funlen - buf := &bytes.Buffer{} + buf := acquireMtprotoFrameBytesBuffer() + defer releaseMtprotoFrameBytesBuffer(buf) + sum := crc32.NewIEEE() writer := io.MultiWriter(buf, sum) @@ -71,7 +73,6 @@ func (w *wrapperMtprotoFrame) Read() (conntypes.Packet, error) { // nolint: funl } buf.Reset() - buf.Grow(int(messageLength) - 4 - 4) if _, err := io.CopyN(writer, w.parent, int64(messageLength)-4-4); err != nil { return nil, fmt.Errorf("cannot read the message frame: %w", err) @@ -113,8 +114,8 @@ func (w *wrapperMtprotoFrame) Write(p conntypes.Packet) error { messageLength := 4 + 4 + len(p) + 4 paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize - buf := &bytes.Buffer{} - buf.Grow(messageLength + paddingLength) + buf := acquireMtprotoFrameBytesBuffer() + defer releaseMtprotoFrameBytesBuffer(buf) binary.Write(buf, binary.LittleEndian, uint32(messageLength)) // nolint: errcheck binary.Write(buf, binary.LittleEndian, w.writeSeqNo) // nolint: errcheck diff --git a/wrappers/packet/pools.go b/wrappers/packet/pools.go new file mode 100644 index 0000000..c27e30e --- /dev/null +++ b/wrappers/packet/pools.go @@ -0,0 +1,23 @@ +package packet + +import ( + "bytes" + "sync" +) + +var ( + poolMtprotoFrameBytesBuffer = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) + +func acquireMtprotoFrameBytesBuffer() *bytes.Buffer { + return poolMtprotoFrameBytesBuffer.Get().(*bytes.Buffer) +} + +func releaseMtprotoFrameBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + poolMtprotoFrameBytesBuffer.Put(buf) +} diff --git a/wrappers/packetack/client_abridged.go b/wrappers/packetack/client_abridged.go index 1b8aa9a..556a4f2 100644 --- a/wrappers/packetack/client_abridged.go +++ b/wrappers/packetack/client_abridged.go @@ -88,7 +88,9 @@ func (w *wrapperClientAbridged) Write(packet conntypes.Packet, acks *conntypes.C return nil case packetLength < clientAbridgedLargePacketLength: length24 := utils.ToUint24(uint32(packetLength)) - buf := bytes.Buffer{} + + buf := acquireClientBytesBuffer() + defer releaseClientBytesBuffer(buf) buf.WriteByte(byte(clientAbridgedSmallPacketLength)) buf.Write(length24[:]) diff --git a/wrappers/packetack/client_intermediate_secure.go b/wrappers/packetack/client_intermediate_secure.go index 153e779..9c3c448 100644 --- a/wrappers/packetack/client_intermediate_secure.go +++ b/wrappers/packetack/client_intermediate_secure.go @@ -1,7 +1,6 @@ package packetack import ( - "bytes" "encoding/binary" "fmt" "math/rand" @@ -35,11 +34,13 @@ func (w *wrapperClientIntermediateSecure) Write(packet conntypes.Packet, acks *c return nil } - buf := bytes.Buffer{} + buf := acquireClientBytesBuffer() + defer releaseClientBytesBuffer(buf) + paddingLength := rand.Intn(4) buf.Grow(4 + len(packet) + paddingLength) - binary.Write(&buf, binary.LittleEndian, uint32(len(packet)+paddingLength)) // nolint: errcheck + binary.Write(buf, binary.LittleEndian, uint32(len(packet)+paddingLength)) // nolint: errcheck buf.Write(packet) buf.Write(make([]byte, paddingLength)) diff --git a/wrappers/packetack/pools.go b/wrappers/packetack/pools.go new file mode 100644 index 0000000..30a0d7d --- /dev/null +++ b/wrappers/packetack/pools.go @@ -0,0 +1,23 @@ +package packetack + +import ( + "bytes" + "sync" +) + +var ( + poolClientBytesBuffer = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) + +func acquireClientBytesBuffer() *bytes.Buffer { + return poolClientBytesBuffer.Get().(*bytes.Buffer) +} + +func releaseClientBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + poolClientBytesBuffer.Put(buf) +} diff --git a/wrappers/packetack/proxy.go b/wrappers/packetack/proxy.go index 62ac7d2..bd98edf 100644 --- a/wrappers/packetack/proxy.go +++ b/wrappers/packetack/proxy.go @@ -23,8 +23,8 @@ type wrapperProxy struct { func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.ConnectionAcks) error { buf := bytes.Buffer{} - flags := w.flags + if acks.Quick { flags |= rpc.ProxyRequestFlagsQuickAck } @@ -43,6 +43,7 @@ func (w *wrapperProxy) Write(packet conntypes.Packet, acks *conntypes.Connection buf.WriteByte(byte(len(config.C.AdTag))) buf.Write(config.C.AdTag) buf.Write(make([]byte, (4-buf.Len()%4)%4)) + buf.Grow(len(packet)) buf.Write(packet) return w.proxy.Write(buf.Bytes()) diff --git a/wrappers/stream/faketls.go b/wrappers/stream/faketls.go index 3db845b..d074dcb 100644 --- a/wrappers/stream/faketls.go +++ b/wrappers/stream/faketls.go @@ -1,6 +1,7 @@ package stream import ( + "bytes" "errors" "fmt" "net" @@ -39,13 +40,19 @@ func (w *wrapperFakeTLS) WriteTimeout(p []byte, timeout time.Duration) (int, err func (w *wrapperFakeTLS) write(p []byte, writeFunc func([]byte) (int, error)) (int, error) { sum := 0 + buf := acquireBytesBuffer() + defer releaseBytesBuffer(buf) + for _, v := range tlstypes.MakeRecords(p) { - _, err := writeFunc(v.Bytes()) + buf.Reset() + v.WriteBytes(buf) + + _, err := writeFunc(buf.Bytes()) if err != nil { return sum, err } - sum += len(v.Data.Bytes()) + sum += v.Data.Len() } return sum, nil @@ -86,7 +93,10 @@ func NewFakeTLS(socket conntypes.StreamReadWriteCloser) conntypes.StreamReadWrit switch rec.Type { case tlstypes.RecordTypeChangeCipherSpec: case tlstypes.RecordTypeApplicationData: - return rec.Data.Bytes(), nil + buf := &bytes.Buffer{} + rec.Data.WriteBytes(buf) + + return buf.Bytes(), nil default: return nil, fmt.Errorf("unsupported record type %v", rec.Type) } diff --git a/wrappers/stream/mtproto_cipher.go b/wrappers/stream/mtproto_cipher.go index a46b528..fd99169 100644 --- a/wrappers/stream/mtproto_cipher.go +++ b/wrappers/stream/mtproto_cipher.go @@ -1,7 +1,6 @@ package stream import ( - "bytes" "crypto/aes" "crypto/cipher" "crypto/md5" // nolint: gosec @@ -54,7 +53,9 @@ func mtprotoDeriveKeys(purpose mtprotoCipherPurpose, resp *rpc.NonceResponse, client, remote *net.TCPAddr, secret []byte) ([]byte, []byte) { - message := bytes.Buffer{} + message := acquireBytesBuffer() + defer releaseBytesBuffer(message) + message.Write(resp.Nonce) // nolint: gosec message.Write(req.Nonce) // nolint: gosec message.Write(req.CryptoTS) // nolint: gosec diff --git a/wrappers/stream/obfuscated2.go b/wrappers/stream/obfuscated2.go index 603319f..e005c84 100644 --- a/wrappers/stream/obfuscated2.go +++ b/wrappers/stream/obfuscated2.go @@ -1,11 +1,9 @@ package stream import ( - "bytes" "crypto/cipher" "fmt" "net" - "sync" "time" "go.uber.org/zap" @@ -13,23 +11,6 @@ import ( "github.com/9seconds/mtg/conntypes" ) -var ( - poolWrapperObfuscated2WritePool = sync.Pool{ - New: func() interface{} { - return &bytes.Buffer{} - }, - } -) - -func poolWrapperObfuscated2WritePoolAcquire() *bytes.Buffer { - return poolWrapperObfuscated2WritePool.Get().(*bytes.Buffer) -} - -func poolWrapperObfuscated2WritePoolRelease(buf *bytes.Buffer) { - buf.Reset() - poolWrapperObfuscated2WritePool.Put(buf) -} - type wrapperObfuscated2 struct { encryptor cipher.Stream decryptor cipher.Stream @@ -59,8 +40,8 @@ func (w *wrapperObfuscated2) Read(p []byte) (int, error) { } func (w *wrapperObfuscated2) WriteTimeout(p []byte, timeout time.Duration) (int, error) { - buffer := poolWrapperObfuscated2WritePoolAcquire() - defer poolWrapperObfuscated2WritePoolRelease(buffer) + buffer := acquireBytesBuffer() + defer releaseBytesBuffer(buffer) buffer.Write(p) @@ -72,8 +53,8 @@ func (w *wrapperObfuscated2) WriteTimeout(p []byte, timeout time.Duration) (int, } func (w *wrapperObfuscated2) Write(p []byte) (int, error) { - buffer := poolWrapperObfuscated2WritePoolAcquire() - defer poolWrapperObfuscated2WritePoolRelease(buffer) + buffer := acquireBytesBuffer() + defer releaseBytesBuffer(buffer) buffer.Write(p) diff --git a/wrappers/stream/pools.go b/wrappers/stream/pools.go new file mode 100644 index 0000000..3fd9907 --- /dev/null +++ b/wrappers/stream/pools.go @@ -0,0 +1,23 @@ +package stream + +import ( + "bytes" + "sync" +) + +var ( + poolBytesBuffer = sync.Pool{ + New: func() interface{} { + return &bytes.Buffer{} + }, + } +) + +func acquireBytesBuffer() *bytes.Buffer { + return poolBytesBuffer.Get().(*bytes.Buffer) +} + +func releaseBytesBuffer(buf *bytes.Buffer) { + buf.Reset() + poolBytesBuffer.Put(buf) +}