From 5413a5b0904ed032a6722f33ed54293943194a29 Mon Sep 17 00:00:00 2001 From: 9seconds Date: Thu, 5 Jul 2018 20:46:44 +0300 Subject: [PATCH] Fix infinite hangs --- mtproto/wrappers/abridged.go | 6 +++--- mtproto/wrappers/frame.go | 2 +- mtproto/wrappers/intermediate.go | 3 +++ mtproto/wrappers/proxy_request.go | 11 +++++------ wrappers/blockcipherrwc.go | 11 ++++++----- 5 files changed, 18 insertions(+), 15 deletions(-) diff --git a/mtproto/wrappers/abridged.go b/mtproto/wrappers/abridged.go index 8d38a97..9c22785 100644 --- a/mtproto/wrappers/abridged.go +++ b/mtproto/wrappers/abridged.go @@ -28,6 +28,9 @@ type AbridgedReadWriteCloserWithAddr struct { func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) { return a.BufferedRead(p, func() error { + a.opts.QuickAck = false + a.opts.SimpleAck = false + buf := &bytes.Buffer{} buf.Grow(3) @@ -37,7 +40,6 @@ func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) { msgLength := uint8(buf.Bytes()[0]) buf.Reset() - a.opts.QuickAck = false if msgLength >= abridgedQuickAckLength { a.opts.QuickAck = true msgLength -= 0x80 @@ -78,13 +80,11 @@ func (a *AbridgedReadWriteCloserWithAddr) Write(p []byte) (int, error) { case packetLength < abridgedLargePacketLength: length24 := toUint24(uint32(packetLength)) - buf := &bytes.Buffer{} buf.Grow(1 + 3 + len(p)) buf.WriteByte(byte(abridgedSmallPacketLength)) buf.Write(length24[:]) buf.Write(p) - return a.conn.Write(buf.Bytes()) default: diff --git a/mtproto/wrappers/frame.go b/mtproto/wrappers/frame.go index 907dbde..2c5cf8a 100644 --- a/mtproto/wrappers/frame.go +++ b/mtproto/wrappers/frame.go @@ -66,7 +66,7 @@ func (f *FrameRWC) Read(p []byte) (int, error) { if _, err := io.CopyN(writer, f.conn, 4); err != nil { return errors.Annotate(err, "Cannot read frame padding") } - if !bytes.Equal(buf.Bytes(), frameRWCPadding[:]) { + if !bytes.Equal(buf.Bytes(), frameRWCPadding) { break } } diff --git a/mtproto/wrappers/intermediate.go b/mtproto/wrappers/intermediate.go index 4e35223..d80eabe 100644 --- a/mtproto/wrappers/intermediate.go +++ b/mtproto/wrappers/intermediate.go @@ -23,6 +23,9 @@ type IntermediateReadWriteCloserWithAddr struct { func (i *IntermediateReadWriteCloserWithAddr) Read(p []byte) (int, error) { return i.BufferedRead(p, func() error { + i.opts.QuickAck = false + i.opts.SimpleAck = false + buf := &bytes.Buffer{} buf.Grow(4) diff --git a/mtproto/wrappers/proxy_request.go b/mtproto/wrappers/proxy_request.go index 1d2ba67..7d0adbe 100644 --- a/mtproto/wrappers/proxy_request.go +++ b/mtproto/wrappers/proxy_request.go @@ -29,15 +29,16 @@ func (p *ProxyRequestReadWriteCloserWithAddr) Read(buf []byte) (int, error) { return errors.Annotate(err, "Cannot read RPC tag") } - if bytes.Equal(ansBuf.Bytes(), rpc.TagCloseExt) { + switch { + case bytes.Equal(ansBuf.Bytes(), rpc.TagCloseExt): return p.readCloseExt() - } else if bytes.Equal(ansBuf.Bytes(), rpc.TagProxyAns) { + case bytes.Equal(ansBuf.Bytes(), rpc.TagProxyAns): return p.readProxyAns(buf) - } else if bytes.Equal(ansBuf.Bytes(), rpc.TagSimpleAck) { + case bytes.Equal(ansBuf.Bytes(), rpc.TagSimpleAck): return p.readSimpleAck() } - return nil + return errors.Errorf("Unknown RPC answer %s", ansBuf.Bytes()) }) } @@ -80,8 +81,6 @@ func (p *ProxyRequestReadWriteCloserWithAddr) Write(raw []byte) (int, error) { if _, err := p.conn.Write(p.req.Bytes(raw)); err != nil { return 0, err } - p.req.Options.SimpleAck = false - p.req.Options.QuickAck = false return len(raw), nil } diff --git a/wrappers/blockcipherrwc.go b/wrappers/blockcipherrwc.go index 9db024d..8d2a06c 100644 --- a/wrappers/blockcipherrwc.go +++ b/wrappers/blockcipherrwc.go @@ -1,6 +1,7 @@ package wrappers import ( + "bytes" "crypto/aes" "crypto/cipher" "net" @@ -18,16 +19,16 @@ type BlockCipherReadWriteCloserWithAddr struct { func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) { return c.BufferedRead(p, func() error { - bufferLength := c.Buffer.Len() - for bufferLength%aes.BlockSize != 0 || bufferLength == 0 { + buf := &bytes.Buffer{} + for buf.Len()%aes.BlockSize != 0 || buf.Len() == 0 { n, err := c.conn.Read(p) if err != nil { return errors.Annotate(err, "Cannot read from socket") } - c.Buffer.Write(p[:n]) - bufferLength = c.Buffer.Len() + buf.Write(p[:n]) } - c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.Bytes()) + c.decryptor.CryptBlocks(buf.Bytes(), buf.Bytes()) + c.Buffer.Write(buf.Bytes()) return nil })