diff --git a/mtproto/rpc/rpc_proxy_request.go b/mtproto/rpc/rpc_proxy_request.go index 2ffa36f..feaecb3 100644 --- a/mtproto/rpc/rpc_proxy_request.go +++ b/mtproto/rpc/rpc_proxy_request.go @@ -29,10 +29,9 @@ type RPCProxyRequest struct { LocalIPPort [rpcProxyRequestIPPortLength]byte ADTag []byte Extras Extras - Message *bytes.Buffer } -func (r *RPCProxyRequest) Bytes() []byte { +func (r *RPCProxyRequest) Bytes(message []byte) []byte { buf := &bytes.Buffer{} flags := r.Flags @@ -40,8 +39,7 @@ func (r *RPCProxyRequest) Bytes() []byte { flags |= RPCProxyRequestFlagsQuickAck } - messageBytes := r.Message.Bytes() - if bytes.HasPrefix(messageBytes, rpcProxyRequestFlagsEncryptedPrefix[:]) { + if bytes.HasPrefix(message, rpcProxyRequestFlagsEncryptedPrefix[:]) { flags |= RPCProxyRequestFlagsEncrypted } @@ -58,9 +56,7 @@ func (r *RPCProxyRequest) Bytes() []byte { for i := 0; i < (buf.Len() % 4); i++ { buf.WriteByte(0x00) } - if r.Message != nil { - buf.Write(messageBytes) - } + buf.Write(message) return buf.Bytes() } diff --git a/mtproto/wrappers/frame.go b/mtproto/wrappers/frame.go index fc3e1d6..9a77c73 100644 --- a/mtproto/wrappers/frame.go +++ b/mtproto/wrappers/frame.go @@ -22,11 +22,11 @@ const ( var frameRWCPadding = [4]byte{0x04, 0x00, 0x00, 0x00} type FrameRWC struct { - conn wrappers.ReadWriteCloserWithAddr + wrappers.BufferedReader + conn wrappers.ReadWriteCloserWithAddr readSeqNo int32 writeSeqNo int32 - readBuf *bytes.Buffer } func (f *FrameRWC) Write(buf []byte) (int, error) { @@ -54,51 +54,49 @@ func (f *FrameRWC) Write(buf []byte) (int, error) { } func (f *FrameRWC) Read(p []byte) (int, error) { - if f.readBuf.Len() > 0 { - return f.flush(p) - } + return f.BufferedRead(p, func() error { + buf := &bytes.Buffer{} + for { + buf.Reset() + if _, err := io.CopyN(buf, f.conn, 4); err != nil { + return errors.Annotate(err, "Cannot read frame padding") + } + if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) { + break + } + } + + messageLength := binary.LittleEndian.Uint32(buf.Bytes()) + if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength { + return errors.Errorf("Incorrect frame message length %d", messageLength) + } + sum := crc32.NewIEEE() + sum.Write(buf.Bytes()) - buf := &bytes.Buffer{} - for { buf.Reset() - if _, err := io.CopyN(buf, f.conn, 4); err != nil { - return 0, errors.Annotate(err, "Cannot read frame padding") + buf.Grow(int(messageLength) - 4) // -4 because we already read the first number + if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil { + return errors.Annotate(err, "Cannot read the message frame") } - if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) { - break + sum.Write(buf.Bytes()) + + var seqNo int32 + binary.Read(buf, binary.LittleEndian, seqNo) + if seqNo != f.readSeqNo { + return errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo) } - } + f.readSeqNo++ - messageLength := binary.LittleEndian.Uint32(buf.Bytes()) - if messageLength%4 != 0 || messageLength < frameRWCMinMessageLength || messageLength > frameRWCMaxMessageLength { - return 0, errors.Errorf("Incorrect frame message length %d", messageLength) - } - sum := crc32.NewIEEE() - sum.Write(buf.Bytes()) + data := buf.Bytes()[:int(messageLength)-4-4-4] + checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:]) + if checksum != sum.Sum32() { + return errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum) - buf.Reset() - buf.Grow(int(messageLength) - 4) // -4 because we already read the first number - if _, err := io.CopyN(buf, f.conn, int64(messageLength)-4); err != nil { - return 0, errors.Annotate(err, "Cannot read the message frame") - } - sum.Write(buf.Bytes()) + } + f.Buffer.Write(data) - var seqNo int32 - binary.Read(buf, binary.LittleEndian, seqNo) - if seqNo != f.readSeqNo { - return 0, errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, f.readSeqNo) - } - f.readSeqNo++ - - data := buf.Bytes()[:int(messageLength)-4-4-4] - checksum := binary.LittleEndian.Uint32(buf.Bytes()[int(messageLength)-4-4-4:]) - if checksum != sum.Sum32() { - return 0, errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum) - - } - f.readBuf.Write(data) - - return f.flush(p) + return nil + }) } func (f *FrameRWC) Close() error { @@ -109,28 +107,11 @@ func (f *FrameRWC) Addr() *net.TCPAddr { return f.conn.Addr() } -func (f *FrameRWC) flush(p []byte) (int, error) { - sizeToRead := len(p) - if f.readBuf.Len() < sizeToRead { - sizeToRead = f.readBuf.Len() - } - - data := f.readBuf.Bytes() - copy(p, data[:sizeToRead]) - if sizeToRead == f.readBuf.Len() { - f.readBuf.Reset() - } else { - f.readBuf = bytes.NewBuffer(data[sizeToRead:]) - } - - return sizeToRead, nil -} - func NewFrameRWC(conn wrappers.ReadWriteCloserWithAddr, seqNo int32) wrappers.ReadWriteCloserWithAddr { return &FrameRWC{ - conn: conn, - readSeqNo: seqNo, - writeSeqNo: seqNo, - readBuf: &bytes.Buffer{}, + BufferedReader: wrappers.NewBufferedReader(), + conn: conn, + readSeqNo: seqNo, + writeSeqNo: seqNo, } } diff --git a/wrappers/blockcipherrwc.go b/wrappers/blockcipherrwc.go index ad2c9d1..4d11d5d 100644 --- a/wrappers/blockcipherrwc.go +++ b/wrappers/blockcipherrwc.go @@ -1,7 +1,6 @@ package wrappers import ( - "bytes" "crypto/aes" "crypto/cipher" "net" @@ -10,7 +9,7 @@ import ( ) type BlockCipherReadWriteCloserWithAddr struct { - buf *bytes.Buffer + BufferedReader conn ReadWriteCloserWithAddr encryptor cipher.BlockMode @@ -18,20 +17,19 @@ type BlockCipherReadWriteCloserWithAddr struct { } func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) { - if c.buf.Len() > 0 { - return c.flush(p) - } - - for c.buf.Len() == 0 || c.buf.Len()%aes.BlockSize != 0 { - n, err := c.conn.Read(p) - if err != nil { - return 0, errors.Annotate(err, "Cannot read from socket") + return c.BufferedRead(p, func() error { + bufferLength := c.Buffer.Len() + for bufferLength%aes.BlockSize != 0 || bufferLength == 0 { + n, err := c.conn.Read(p) + if err != nil { + return errors.Annotate(err, "Cannot read from socket") + } + c.Buffer.Write(p[:n]) } - c.buf.Write(p[:n]) - } - c.decryptor.CryptBlocks(c.buf.Bytes(), c.buf.Bytes()) + c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.Bytes()) - return c.flush(p) + return nil + }) } func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) { @@ -39,19 +37,13 @@ func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) { return 0, errors.Errorf("Incorrect block size %d", len(p)) } - buf := getBuffer() - defer putBuffer(buf) - buf.Grow(len(p)) - buf.Write(p) - - encrypted := buf.Bytes() + encrypted := make([]byte, len(p)) c.encryptor.CryptBlocks(encrypted, p) return c.conn.Write(encrypted) } func (c *BlockCipherReadWriteCloserWithAddr) Close() error { - defer putBuffer(c.buf) return c.conn.Close() } @@ -59,32 +51,11 @@ func (c *BlockCipherReadWriteCloserWithAddr) Addr() *net.TCPAddr { return c.conn.Addr() } -func (c *BlockCipherReadWriteCloserWithAddr) flush(p []byte) (int, error) { - sizeToRead := len(p) - if c.buf.Len() < sizeToRead { - sizeToRead = c.buf.Len() - } - - data := c.buf.Bytes() - copy(p, data[:sizeToRead]) - if sizeToRead == c.buf.Len() { - c.buf.Reset() - } else { - newBuf := getBuffer() - newBuf.Write(data[sizeToRead:]) - - putBuffer(c.buf) - c.buf = newBuf - } - - return sizeToRead, nil -} - func NewBlockCipherRWC(conn ReadWriteCloserWithAddr, encryptor, decryptor cipher.BlockMode) ReadWriteCloserWithAddr { return &BlockCipherReadWriteCloserWithAddr{ - buf: getBuffer(), - conn: conn, - encryptor: encryptor, - decryptor: decryptor, + BufferedReader: NewBufferedReader(), + conn: conn, + encryptor: encryptor, + decryptor: decryptor, } } diff --git a/wrappers/buffered_reader.go b/wrappers/buffered_reader.go new file mode 100644 index 0000000..a824805 --- /dev/null +++ b/wrappers/buffered_reader.go @@ -0,0 +1,38 @@ +package wrappers + +import "bytes" + +type BufferedReader struct { + Buffer *bytes.Buffer +} + +func (b *BufferedReader) BufferedRead(p []byte, callback func() error) (int, error) { + if b.Buffer.Len() > 0 { + return b.flush(p) + } + if err := callback(); err != nil { + return 0, err + } + return b.flush(p) +} + +func (b *BufferedReader) flush(p []byte) (int, error) { + sizeToRead := len(p) + if b.Buffer.Len() < sizeToRead { + sizeToRead = b.Buffer.Len() + } + + data := b.Buffer.Bytes() + copy(p, data[:sizeToRead]) + if sizeToRead == b.Buffer.Len() { + b.Buffer.Reset() + } else { + b.Buffer = bytes.NewBuffer(data[sizeToRead:]) + } + + return sizeToRead, nil +} + +func NewBufferedReader() BufferedReader { + return BufferedReader{Buffer: &bytes.Buffer{}} +}