From a7d46727baf50d3aa9bca708ad05473fc6f7eedb Mon Sep 17 00:00:00 2001 From: 9seconds Date: Sat, 7 Jul 2018 17:40:21 +0300 Subject: [PATCH] Move intermediate --- mtproto/wrappers/intermediate.go | 90 ----------------------- wrappers/mtproto_abridged.go | 21 ++++-- wrappers/mtproto_intermediate.go | 119 +++++++++++++++++++++++++++++++ 3 files changed, 134 insertions(+), 96 deletions(-) delete mode 100644 mtproto/wrappers/intermediate.go create mode 100644 wrappers/mtproto_intermediate.go diff --git a/mtproto/wrappers/intermediate.go b/mtproto/wrappers/intermediate.go deleted file mode 100644 index 0468d47..0000000 --- a/mtproto/wrappers/intermediate.go +++ /dev/null @@ -1,90 +0,0 @@ -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 { - buf := &bytes.Buffer{} - buf.Grow(4) - - if _, err := io.CopyN(buf, i.conn, 4); err != nil { - return errors.Annotate(err, "Cannot read message length") - } - length := binary.LittleEndian.Uint32(buf.Bytes()) - buf.Reset() - buf.Grow(int(length)) - - if length > intermediateQuickAckLength { - i.opts.ReadHacks.QuickAck = true - length -= intermediateQuickAckLength - } - - 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.WriteHacks.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 (i *IntermediateReadWriteCloserWithAddr) SocketID() string { - return i.conn.SocketID() -} - -func NewIntermediateRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto.ConnectionOpts) wrappers.ReadWriteCloserWithAddr { - return &IntermediateReadWriteCloserWithAddr{ - BufferedReader: wrappers.NewBufferedReader(), - conn: conn, - opts: connOpts, - } -} diff --git a/wrappers/mtproto_abridged.go b/wrappers/mtproto_abridged.go index 1aa7f8c..f1c1a5c 100644 --- a/wrappers/mtproto_abridged.go +++ b/wrappers/mtproto_abridged.go @@ -26,9 +26,9 @@ type MTProtoAbridged struct { } func (m *MTProtoAbridged) Read() ([]byte, error) { - m.LogDebug("Read abridged packet", - "simple_ack", m.opts.WriteHacks.SimpleAck, - "quick_ack", m.opts.WriteHacks.QuickAck, + m.LogDebug("Read packet", + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, "counter", m.readCounter, ) @@ -41,9 +41,11 @@ func (m *MTProtoAbridged) Read() ([]byte, error) { msgLength := uint8(buf.Bytes()[0]) buf.Reset() - m.LogDebug("Abridged packet first byte", + m.LogDebug("Packet first byte", "byte", msgLength, "counter", m.readCounter, + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, ) if msgLength >= mtprotoAbridgedQuickAckLength { @@ -62,8 +64,10 @@ func (m *MTProtoAbridged) Read() ([]byte, error) { } msgLength32 *= 4 - m.LogDebug("Abridged packet length", + m.LogDebug("Packet length", "length", msgLength32, + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, "counter", m.readCounter, ) @@ -79,7 +83,7 @@ func (m *MTProtoAbridged) Read() ([]byte, error) { } func (m *MTProtoAbridged) Write(p []byte) (int, error) { - m.LogDebug("Write abridged packet", + m.LogDebug("Write packet", "length", len(p), "simple_ack", m.opts.WriteHacks.SimpleAck, "quick_ack", m.opts.WriteHacks.QuickAck, @@ -91,6 +95,7 @@ func (m *MTProtoAbridged) Write(p []byte) (int, error) { } if m.opts.WriteHacks.SimpleAck { + m.writeCounter++ return m.conn.Write(utils.ReverseBytes(p)) } @@ -120,18 +125,22 @@ func (m *MTProtoAbridged) Write(p []byte) (int, error) { } func (m *MTProtoAbridged) LogDebug(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "abridged"}) m.conn.LogDebug(msg, data...) } func (m *MTProtoAbridged) LogInfo(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "abridged"}) m.conn.LogInfo(msg, data...) } func (m *MTProtoAbridged) LogWarn(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "abridged"}) m.conn.LogWarn(msg, data...) } func (m *MTProtoAbridged) LogError(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "abridged"}) m.conn.LogError(msg, data...) } diff --git a/wrappers/mtproto_intermediate.go b/wrappers/mtproto_intermediate.go new file mode 100644 index 0000000..2605a71 --- /dev/null +++ b/wrappers/mtproto_intermediate.go @@ -0,0 +1,119 @@ +package wrappers + +import ( + "bytes" + "encoding/binary" + "io" + "net" + + "github.com/9seconds/mtg/mtproto" + "github.com/juju/errors" +) + +const mtprotoIntermediateQuickAckLength = 0x80000000 + +type MTProtoIntermediate struct { + conn WrapStreamReadWriteCloser + opts *mtproto.ConnectionOpts + + readCounter uint32 + writeCounter uint32 +} + +func (m *MTProtoIntermediate) Read() ([]byte, error) { + m.LogDebug("Read packet", + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, + "counter", m.readCounter, + ) + + buf := &bytes.Buffer{} + buf.Grow(4) + + if _, err := io.CopyN(buf, m.conn, 4); err != nil { + return nil, errors.Annotate(err, "Cannot read message length") + } + length := binary.LittleEndian.Uint32(buf.Bytes()) + + m.LogDebug("Packet message length", + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, + "counter", m.readCounter, + "length", length, + ) + + if length > mtprotoIntermediateQuickAckLength { + m.opts.ReadHacks.QuickAck = true + length -= mtprotoIntermediateQuickAckLength + } + + buf.Reset() + buf.Grow(int(length)) + if _, err := io.CopyN(buf, m.conn, int64(length)); err != nil { + return nil, errors.Annotate(err, "Cannot read the message") + } + + if length%4 != 0 { + length -= length % 4 + } + m.readCounter++ + + return buf.Bytes()[:length], nil +} + +func (m *MTProtoIntermediate) Write(p []byte) (int, error) { + m.LogDebug("Write packet", + "simple_ack", m.opts.WriteHacks.SimpleAck, + "quick_ack", m.opts.WriteHacks.QuickAck, + "counter", m.writeCounter, + ) + m.writeCounter++ + + if m.opts.ReadHacks.SimpleAck { + return m.conn.Write(p) + } + + var length [4]byte + binary.LittleEndian.PutUint32(length[:], uint32(len(p))) + + return m.conn.Write(append(length[:], p...)) +} + +func (m *MTProtoIntermediate) LogDebug(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "intermediate"}) + m.conn.LogDebug(msg, data...) +} + +func (m *MTProtoIntermediate) LogInfo(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "intermediate"}) + m.conn.LogInfo(msg, data...) +} + +func (m *MTProtoIntermediate) LogWarn(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "intermediate"}) + m.conn.LogWarn(msg, data...) +} + +func (m *MTProtoIntermediate) LogError(msg string, data ...interface{}) { + data = append(data, []interface{}{"type", "intermediate"}) + m.conn.LogError(msg, data...) +} + +func (m *MTProtoIntermediate) LocalAddr() *net.TCPAddr { + return m.conn.LocalAddr() +} + +func (m *MTProtoIntermediate) RemoteAddr() *net.TCPAddr { + return m.conn.RemoteAddr() +} + +func (m *MTProtoIntermediate) Close() error { + return m.conn.Close() +} + +func NewMTProtoIntermediate(conn WrapStreamReadWriteCloser, opts *mtproto.ConnectionOpts) WrapPacketReadWriteCloser { + return &MTProtoIntermediate{ + conn: conn, + opts: opts, + } +}