From d636010377261d962912038b967523ba3b0f7652 Mon Sep 17 00:00:00 2001 From: 9seconds Date: Thu, 5 Jul 2018 09:49:15 +0300 Subject: [PATCH] Add intermediate rwc --- mtproto/wrappers/abridged.go | 19 +++---- mtproto/wrappers/intermediate.go | 83 +++++++++++++++++++++++++++++++ mtproto/wrappers/proxy_request.go | 5 +- 3 files changed, 96 insertions(+), 11 deletions(-) create mode 100644 mtproto/wrappers/intermediate.go diff --git a/mtproto/wrappers/abridged.go b/mtproto/wrappers/abridged.go index db5a99e..66b440d 100644 --- a/mtproto/wrappers/abridged.go +++ b/mtproto/wrappers/abridged.go @@ -20,14 +20,14 @@ const ( abridgedLargePacketLength = 16777216 // 256 ^ 3 ) -type AbridgedReadWriteCloserAddr struct { +type AbridgedReadWriteCloserWithAddr struct { wrappers.BufferedReader conn wrappers.ReadWriteCloserWithAddr opts *mtproto.ConnectionOpts } -func (a *AbridgedReadWriteCloserAddr) Read(p []byte) (int, error) { +func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) { return a.BufferedRead(p, func() error { var msgLength uint8 if err := binary.Read(a.conn, binary.LittleEndian, &msgLength); err != nil { @@ -62,7 +62,7 @@ func (a *AbridgedReadWriteCloserAddr) Read(p []byte) (int, error) { }) } -func (a *AbridgedReadWriteCloserAddr) Write(p []byte) (int, error) { +func (a *AbridgedReadWriteCloserWithAddr) Write(p []byte) (int, error) { if len(p)%4 != 0 { return 0, errors.Errorf("Incorrect packet length %d", len(p)) } @@ -92,15 +92,15 @@ func (a *AbridgedReadWriteCloserAddr) Write(p []byte) (int, error) { } } -func (a *AbridgedReadWriteCloserAddr) Close() error { +func (a *AbridgedReadWriteCloserWithAddr) Close() error { return a.conn.Close() } -func (a *AbridgedReadWriteCloserAddr) LocalAddr() *net.TCPAddr { +func (a *AbridgedReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { return a.conn.LocalAddr() } -func (a *AbridgedReadWriteCloserAddr) RemoteAddr() *net.TCPAddr { +func (a *AbridgedReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { return a.conn.RemoteAddr() } @@ -113,8 +113,9 @@ func fromUint24(number uint24) uint32 { } func NewAbridgedRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto.ConnectionOpts) wrappers.ReadWriteCloserWithAddr { - return &AbridgedReadWriteCloserAddr{ - conn: conn, - opts: connOpts, + return &AbridgedReadWriteCloserWithAddr{ + BufferedReader: wrappers.NewBufferedReader(), + conn: conn, + opts: connOpts, } } diff --git a/mtproto/wrappers/intermediate.go b/mtproto/wrappers/intermediate.go new file mode 100644 index 0000000..1c2a6ab --- /dev/null +++ b/mtproto/wrappers/intermediate.go @@ -0,0 +1,83 @@ +package wrappers + +import ( + "bytes" + "encoding/binary" + "io" + "net" + + "github.com/juju/errors" + + "github.com/9seconds/mtg/mtproto" + "github.com/9seconds/mtg/wrappers" +) + +const intermediateQuickAckLength = 0x80000000 + +type IntermediateReadWriteCloserWithAddr struct { + wrappers.BufferedReader + + conn wrappers.ReadWriteCloserWithAddr + opts *mtproto.ConnectionOpts +} + +func (i *IntermediateReadWriteCloserWithAddr) Read(p []byte) (int, error) { + return i.BufferedRead(p, func() error { + var length uint32 + if err := binary.Read(i.conn, binary.LittleEndian, &length); err != nil { + return errors.Annotate(err, "Cannot read message length") + } + + if length > intermediateQuickAckLength { + i.opts.QuickAck = true + length -= intermediateQuickAckLength + } + + buf := &bytes.Buffer{} + buf.Grow(int(length)) + if _, err := io.CopyN(buf, i.conn, int64(length)); err != nil { + return errors.Annotate(err, "Cannot read the message") + } + + if length%4 != 0 { + length -= length % 4 + i.Buffer.Write(buf.Bytes()[:length]) + return nil + } + + i.Buffer.Write(buf.Bytes()) + + return nil + }) +} + +func (i *IntermediateReadWriteCloserWithAddr) Write(p []byte) (int, error) { + if i.opts.SimpleAck { + return i.conn.Write(p) + } + + var length [4]byte + binary.LittleEndian.PutUint32(length[:], uint32(len(p))) + + return i.conn.Write(append(length[:], p...)) +} + +func (i *IntermediateReadWriteCloserWithAddr) Close() error { + return i.conn.Close() +} + +func (i *IntermediateReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { + return i.conn.LocalAddr() +} + +func (i *IntermediateReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { + return i.conn.RemoteAddr() +} + +func NewIntermediateRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto.ConnectionOpts) wrappers.ReadWriteCloserWithAddr { + return &IntermediateReadWriteCloserWithAddr{ + BufferedReader: wrappers.NewBufferedReader(), + conn: conn, + opts: connOpts, + } +} diff --git a/mtproto/wrappers/proxy_request.go b/mtproto/wrappers/proxy_request.go index 5b1b269..657f120 100644 --- a/mtproto/wrappers/proxy_request.go +++ b/mtproto/wrappers/proxy_request.go @@ -105,7 +105,8 @@ func NewProxyRequestRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto } return &ProxyRequestReadWriteCloserWithAddr{ - conn: conn, - req: req, + BufferedReader: wrappers.NewBufferedReader(), + conn: conn, + req: req, }, nil }