diff --git a/mtproto/wrappers/abridged.go b/mtproto/wrappers/abridged.go deleted file mode 100644 index e4bfc13..0000000 --- a/mtproto/wrappers/abridged.go +++ /dev/null @@ -1,130 +0,0 @@ -package wrappers - -import ( - "bytes" - "io" - "net" - - "github.com/juju/errors" - - "github.com/9seconds/mtg/mtproto" - "github.com/9seconds/mtg/wrappers" -) - -type uint24 [3]byte - -const ( - abridgedSmallPacketLength = 0x7f - abridgedQuickAckLength = 0x80 - abridgedLargePacketLength = 16777216 // 256 ^ 3 -) - -type AbridgedReadWriteCloserWithAddr struct { - wrappers.BufferedReader - - conn wrappers.ReadWriteCloserWithAddr - opts *mtproto.ConnectionOpts -} - -func (a *AbridgedReadWriteCloserWithAddr) Read(p []byte) (int, error) { - return a.BufferedRead(p, func() error { - buf := &bytes.Buffer{} - buf.Grow(3) - - q := make([]byte, 1) - - if _, err := io.CopyN(buf, a.conn, 1); err != nil { - return errors.Annotate(err, "Cannot read message length") - } - msgLength := uint8(buf.Bytes()[0]) - q[0] = msgLength - buf.Reset() - - if msgLength >= abridgedQuickAckLength { - a.opts.ReadHacks.QuickAck = true - msgLength -= abridgedQuickAckLength - } - - msgLength32 := uint32(msgLength) - if msgLength == abridgedSmallPacketLength { - if _, err := io.CopyN(buf, a.conn, 3); err != nil { - return errors.Annotate(err, "Cannot read the correct message length") - } - number := uint24{} - copy(number[:], buf.Bytes()) - q = append(q, buf.Bytes()...) - msgLength32 = fromUint24(number) - } - msgLength32 *= 4 - - buf.Reset() - buf.Grow(int(msgLength32)) - - if _, err := io.CopyN(buf, a.conn, int64(msgLength32)); err != nil { - return errors.Annotate(err, "Cannot read message") - } - q = append(q, buf.Bytes()...) - a.Buffer.Write(buf.Bytes()) - - return nil - }) -} - -func (a *AbridgedReadWriteCloserWithAddr) Write(p []byte) (int, error) { - if len(p)%4 != 0 { - return 0, errors.Errorf("Incorrect packet length %d", len(p)) - } - if a.opts.WriteHacks.SimpleAck { - return a.conn.Write(reverseBytes(p)) - } - - packetLength := len(p) / 4 - switch { - case packetLength < abridgedSmallPacketLength: - newData := append([]byte{byte(packetLength)}, p...) - return a.conn.Write(newData) - - 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()) - } - - return 0, errors.Errorf("Packet is too big %d", len(p)) -} - -func (a *AbridgedReadWriteCloserWithAddr) Close() error { - return a.conn.Close() -} - -func (a *AbridgedReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return a.conn.LocalAddr() -} - -func (a *AbridgedReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return a.conn.RemoteAddr() -} - -func (a *AbridgedReadWriteCloserWithAddr) SocketID() string { - return a.conn.SocketID() -} - -func toUint24(number uint32) uint24 { - return uint24{byte(number), byte(number >> 8), byte(number >> 16)} -} - -func fromUint24(number uint24) uint32 { - return uint32(number[0]) + (uint32(number[1]) << 8) + (uint32(number[2]) << 16) -} - -func NewAbridgedRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto.ConnectionOpts) wrappers.ReadWriteCloserWithAddr { - return &AbridgedReadWriteCloserWithAddr{ - BufferedReader: wrappers.NewBufferedReader(), - conn: conn, - opts: connOpts, - } -} diff --git a/utils/read_current_data.go b/utils/read_current_data.go index c8669e2..d66802b 100644 --- a/utils/read_current_data.go +++ b/utils/read_current_data.go @@ -2,7 +2,7 @@ package utils import "io" -const readCurrentDataBufferSize = 1024 + 1 +const readCurrentDataBufferSize = 1024 + 1 // + 1 because telegram operates with blocks mod 4 func ReadCurrentData(src io.Reader) (rv []byte, err error) { buf := make([]byte, readCurrentDataBufferSize) diff --git a/utils/reverse_bytes.go b/utils/reverse_bytes.go new file mode 100644 index 0000000..447adbd --- /dev/null +++ b/utils/reverse_bytes.go @@ -0,0 +1,14 @@ +package utils + +func ReverseBytes(data []byte) []byte { + dataLen := len(data) + rv := make([]byte, dataLen) + + rv[dataLen/2] = data[dataLen/2] + for i := dataLen/2 - 1; i >= 0; i-- { + opp := dataLen - i - 1 + rv[i], rv[opp] = data[opp], data[i] + } + + return rv +} diff --git a/utils/uint24.go b/utils/uint24.go new file mode 100644 index 0000000..350f3d5 --- /dev/null +++ b/utils/uint24.go @@ -0,0 +1,11 @@ +package utils + +type Uint24 [3]byte + +func ToUint24(number uint32) Uint24 { + return Uint24{byte(number), byte(number >> 8), byte(number >> 16)} +} + +func FromUint24(number Uint24) uint32 { + return uint32(number[0]) + (uint32(number[1]) << 8) + (uint32(number[2]) << 16) +} diff --git a/wrappers/blockcipher.go b/wrappers/blockcipher.go new file mode 100644 index 0000000..55b90e1 --- /dev/null +++ b/wrappers/blockcipher.go @@ -0,0 +1,85 @@ +package wrappers + +import ( + "crypto/aes" + "crypto/cipher" + "net" + + "github.com/9seconds/mtg/utils" + "github.com/juju/errors" +) + +type WrapBlockCipher struct { + BufferedReader + + conn WrapStreamReadWriteCloser + encryptor cipher.BlockMode + decryptor cipher.BlockMode +} + +func (w *WrapBlockCipher) Read(p []byte) (int, error) { + return w.BufferedRead(p, func() error { + var buf []byte + + for len(buf) == 0 || len(buf)%aes.BlockSize != 0 { + rv, err := utils.ReadCurrentData(w.conn) + if err != nil { + return errors.Annotate(err, "Cannot read from socket") + } + buf = append(buf, rv...) + } + + w.decryptor.CryptBlocks(buf, buf) + w.Buffer.Write(buf) + + return nil + }) +} + +func (w *WrapBlockCipher) Write(p []byte) (int, error) { + if len(p)%aes.BlockSize > 0 { + return 0, errors.Errorf("Incorrect block size %d", len(p)) + } + + encrypted := make([]byte, len(p)) + w.encryptor.CryptBlocks(encrypted, p) + + return w.conn.Write(encrypted) +} + +func (w *WrapBlockCipher) LogDebug(msg string, data ...interface{}) { + w.conn.LogDebug(msg, data...) +} + +func (w *WrapBlockCipher) LogInfo(msg string, data ...interface{}) { + w.conn.LogInfo(msg, data...) +} + +func (w *WrapBlockCipher) LogWarn(msg string, data ...interface{}) { + w.conn.LogWarn(msg, data...) +} + +func (w *WrapBlockCipher) LogError(msg string, data ...interface{}) { + w.conn.LogError(msg, data...) +} + +func (w *WrapBlockCipher) LocalAddr() *net.TCPAddr { + return w.conn.LocalAddr() +} + +func (w *WrapBlockCipher) RemoteAddr() *net.TCPAddr { + return w.conn.RemoteAddr() +} + +func (w *WrapBlockCipher) Close() error { + return w.conn.Close() +} + +func NewWrapBlockCipher(conn WrapStreamReadWriteCloser, encryptor, decryptor cipher.BlockMode) WrapStreamReadWriteCloser { + return &WrapBlockCipher{ + BufferedReader: NewBufferedReader(), + conn: conn, + encryptor: encryptor, + decryptor: decryptor, + } +} diff --git a/wrappers/blockcipherrwc.go b/wrappers/blockcipherrwc.go deleted file mode 100644 index af5ccf0..0000000 --- a/wrappers/blockcipherrwc.go +++ /dev/null @@ -1,76 +0,0 @@ -package wrappers - -import ( - "crypto/aes" - "crypto/cipher" - "fmt" - "net" - - "github.com/juju/errors" - - "github.com/9seconds/mtg/utils" -) - -type BlockCipherReadWriteCloserWithAddr struct { - BufferedReader - - conn ReadWriteCloserWithAddr - encryptor cipher.BlockMode - decryptor cipher.BlockMode -} - -func (c *BlockCipherReadWriteCloserWithAddr) Read(p []byte) (int, error) { - return c.BufferedRead(p, func() error { - var buf []byte - - for len(buf) == 0 || len(buf)%aes.BlockSize != 0 { - rv, err := utils.ReadCurrentData(c.conn) - if err != nil { - return errors.Annotate(err, "Cannot read from socket") - } - buf = append(buf, rv...) - } - - c.decryptor.CryptBlocks(buf, buf) - c.Buffer.Write(buf) - - return nil - }) -} - -func (c *BlockCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) { - if len(p)%aes.BlockSize > 0 { - return 0, errors.Errorf("Incorrect block size %d", len(p)) - } - - encrypted := make([]byte, len(p)) - c.encryptor.CryptBlocks(encrypted, p) - - return c.conn.Write(encrypted) -} - -func (c *BlockCipherReadWriteCloserWithAddr) Close() error { - fmt.Println("BlockCipherReadWriteCloserWithAddr closes", "sockid", c.SocketID(), "bufsize", c.Buffer.Len()) - return c.conn.Close() -} - -func (c *BlockCipherReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return c.conn.LocalAddr() -} - -func (c *BlockCipherReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return c.conn.RemoteAddr() -} - -func (c *BlockCipherReadWriteCloserWithAddr) SocketID() string { - return c.conn.SocketID() -} - -func NewBlockCipherRWC(conn ReadWriteCloserWithAddr, encryptor, decryptor cipher.BlockMode) ReadWriteCloserWithAddr { - return &BlockCipherReadWriteCloserWithAddr{ - BufferedReader: NewBufferedReader(), - conn: conn, - encryptor: encryptor, - decryptor: decryptor, - } -} diff --git a/wrappers/buffer_pool.go b/wrappers/buffer_pool.go deleted file mode 100644 index ead700d..0000000 --- a/wrappers/buffer_pool.go +++ /dev/null @@ -1,27 +0,0 @@ -package wrappers - -import ( - "bytes" - "sync" -) - -var bufPool sync.Pool - -func getBuffer() *bytes.Buffer { - buf := bufPool.Get().(*bytes.Buffer) - buf.Reset() - - return buf -} - -func putBuffer(buf *bytes.Buffer) { - bufPool.Put(buf) -} - -func init() { - bufPool = sync.Pool{ - New: func() interface{} { - return &bytes.Buffer{} - }, - } -} diff --git a/wrappers/buffered_reader.go b/wrappers/buffered_reader.go index 09d3bca..379534a 100644 --- a/wrappers/buffered_reader.go +++ b/wrappers/buffered_reader.go @@ -1,19 +1,11 @@ package wrappers -import ( - "bytes" - - "github.com/juju/errors" -) +import "bytes" type BufferedReader struct { Buffer *bytes.Buffer } -var ( - BufferedReaderContinue = errors.New("Please continue reading") -) - func (b *BufferedReader) BufferedRead(p []byte, callback func() error) (int, error) { if b.Buffer.Len() > 0 { return b.flush(p) diff --git a/wrappers/conn.go b/wrappers/conn.go new file mode 100644 index 0000000..468820c --- /dev/null +++ b/wrappers/conn.go @@ -0,0 +1,115 @@ +package wrappers + +import ( + "net" + "time" + + "go.uber.org/zap" +) + +type ConnPurpose uint8 + +func (c ConnPurpose) String() string { + switch c { + case ConnPurposeClient: + return "client" + case ConnPurposeTelegram: + return "telegram" + } + + return "" +} + +const ( + ConnPurposeClient = iota + ConnPurposeTelegram +) + +const ( + connTimeoutRead = 5 * time.Minute + connTimeoutWrite = 5 * time.Minute +) + +type WrapConn struct { + purpose ConnPurpose + connID string + conn net.Conn + logger *zap.SugaredLogger + publicIPv4 net.IP + publicIPv6 net.IP +} + +func (w *WrapConn) Write(p []byte) (int, error) { + w.conn.SetWriteDeadline(time.Now().Add(connTimeoutWrite)) + n, err := w.conn.Write(p) + + w.logger.Debugw("Write to stream", "bytes", n, "error", err) + + return n, err +} + +func (w *WrapConn) Read(p []byte) (int, error) { + w.conn.SetReadDeadline(time.Now().Add(connTimeoutRead)) + n, err := w.conn.Read(p) + + w.logger.Debugw("Read from stream", "bytes", n, "error", err) + + return n, err +} + +func (w *WrapConn) Close() error { + defer w.LogDebug("Closed connection") + return w.conn.Close() +} + +func (w *WrapConn) LocalAddr() *net.TCPAddr { + addr := w.conn.LocalAddr().(*net.TCPAddr) + newAddr := *addr + + if w.RemoteAddr().IP.To4() != nil { + if w.publicIPv4 != nil { + newAddr.IP = w.publicIPv4 + } + } else if w.publicIPv6 != nil { + newAddr.IP = w.publicIPv6 + } + + return &newAddr +} + +func (w *WrapConn) RemoteAddr() *net.TCPAddr { + return w.conn.RemoteAddr().(*net.TCPAddr) +} + +func (w *WrapConn) LogDebug(msg string, data ...interface{}) { + w.logger.Debugw(msg, data...) +} + +func (w *WrapConn) LogInfo(msg string, data ...interface{}) { + w.logger.Infow(msg, data...) +} + +func (w *WrapConn) LogWarn(msg string, data ...interface{}) { + w.logger.Warnw(msg, data...) +} + +func (w *WrapConn) LogError(msg string, data ...interface{}) { + w.logger.Errorw(msg, data...) +} + +func NewConn(connID string, purpose ConnPurpose, conn net.Conn, publicIPv4, publicIPv6 net.IP) WrapStreamReadWriteCloser { + logger := zap.S().With( + "connection_id", connID, + "local_address", conn.LocalAddr(), + "remote_address", conn.RemoteAddr(), + ) + + return &WrapConn{ + logger: logger, + purpose: purpose, + connID: connID, + conn: conn, + publicIPv4: publicIPv4, + publicIPv6: publicIPv6, + } +} diff --git a/wrappers/ctx.go b/wrappers/ctx.go new file mode 100644 index 0000000..ee766f8 --- /dev/null +++ b/wrappers/ctx.go @@ -0,0 +1,76 @@ +package wrappers + +import ( + "context" + "net" + + "github.com/juju/errors" +) + +type WrapCtx struct { + cancel context.CancelFunc + conn WrapStreamReadWriteCloser + ctx context.Context +} + +func (w *WrapCtx) Read(p []byte) (int, error) { + select { + case <-w.ctx.Done(): + return 0, errors.Annotate(w.ctx.Err(), "Read is failed because of closed context") + default: + n, err := w.conn.Read(p) + if err != nil { + w.cancel() + } + return n, err + } +} + +func (w *WrapCtx) Write(p []byte) (int, error) { + select { + case <-w.ctx.Done(): + return 0, errors.Annotate(w.ctx.Err(), "Write is failed because of closed context") + default: + n, err := w.conn.Write(p) + if err != nil { + w.cancel() + } + return n, err + } +} + +func (w *WrapCtx) LogDebug(msg string, data ...interface{}) { + w.conn.LogDebug(msg, data...) +} + +func (w *WrapCtx) LogInfo(msg string, data ...interface{}) { + w.conn.LogInfo(msg, data...) +} + +func (w *WrapCtx) LogWarn(msg string, data ...interface{}) { + w.conn.LogWarn(msg, data...) +} + +func (w *WrapCtx) LogError(msg string, data ...interface{}) { + w.conn.LogError(msg, data...) +} + +func (w *WrapCtx) LocalAddr() *net.TCPAddr { + return w.conn.LocalAddr() +} + +func (w *WrapCtx) RemoteAddr() *net.TCPAddr { + return w.conn.RemoteAddr() +} + +func (w *WrapCtx) Close() error { + return w.conn.Close() +} + +func NewCtx(ctx context.Context, cancel context.CancelFunc, conn WrapStreamReadWriteCloser) WrapStreamReadWriteCloser { + return &WrapCtx{ + ctx: ctx, + cancel: cancel, + conn: conn, + } +} diff --git a/wrappers/ctxrwc.go b/wrappers/ctxrwc.go deleted file mode 100644 index fe6f515..0000000 --- a/wrappers/ctxrwc.go +++ /dev/null @@ -1,71 +0,0 @@ -package wrappers - -import ( - "context" - "net" - - "github.com/juju/errors" -) - -// CtxReadWriteCloser wraps underlying connection and does management of the -// context and its cancel function. -type CtxReadWriteCloserWithAddr struct { - ctx context.Context - conn ReadWriteCloserWithAddr - cancel context.CancelFunc -} - -// Read reads from connection -func (c *CtxReadWriteCloserWithAddr) Read(p []byte) (int, error) { - select { - case <-c.ctx.Done(): - return 0, errors.Annotate(c.ctx.Err(), "Read is failed because of closed context") - default: - n, err := c.conn.Read(p) - if err != nil { - c.cancel() - } - return n, err - } -} - -// Write writes into connection. -func (c *CtxReadWriteCloserWithAddr) Write(p []byte) (int, error) { - select { - case <-c.ctx.Done(): - return 0, errors.Annotate(c.ctx.Err(), "Write is failed because of closed context") - default: - n, err := c.conn.Write(p) - if err != nil { - c.cancel() - } - return n, err - } -} - -// Close closes underlying connection. -func (c *CtxReadWriteCloserWithAddr) Close() error { - return c.conn.Close() -} - -func (c *CtxReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return c.conn.LocalAddr() -} - -func (c *CtxReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return c.conn.RemoteAddr() -} - -func (c *CtxReadWriteCloserWithAddr) SocketID() string { - return c.conn.SocketID() -} - -// NewCtxRWC returns ReadWriteCloser which respects given context, -// cancellation etc. -func NewCtxRWC(ctx context.Context, cancel context.CancelFunc, conn ReadWriteCloserWithAddr) ReadWriteCloserWithAddr { - return &CtxReadWriteCloserWithAddr{ - conn: conn, - ctx: ctx, - cancel: cancel, - } -} diff --git a/wrappers/logrwc.go b/wrappers/logrwc.go deleted file mode 100644 index 74d3cbf..0000000 --- a/wrappers/logrwc.go +++ /dev/null @@ -1,59 +0,0 @@ -package wrappers - -import ( - "net" - - "go.uber.org/zap" -) - -// LogReadWriteCloser adds additional logging for reading/writing. All -// logging is performed for debug mode only. -type LogReadWriteCloserWithAddr struct { - conn ReadWriteCloserWithAddr - logger *zap.SugaredLogger - sockid string - name string -} - -// Read reads from connection -func (l *LogReadWriteCloserWithAddr) Read(p []byte) (n int, err error) { - n, err = l.conn.Read(p) - l.logger.Debugw("Finish reading", "name", l.name, "socketid", l.sockid, "nbytes", n, "error", err, "localAddr", l.LocalAddr()) - return -} - -// Write writes into connection. -func (l *LogReadWriteCloserWithAddr) Write(p []byte) (n int, err error) { - n, err = l.conn.Write(p) - l.logger.Debugw("Finish writing", "name", l.name, "socketid", l.sockid, "nbytes", n, "error", err, "localAddr", l.LocalAddr()) - return -} - -// Close closes underlying connection. -func (l *LogReadWriteCloserWithAddr) Close() error { - err := l.conn.Close() - l.logger.Debugw("Finish closing socket", "name", l.name, "socketid", l.sockid, "error", err) - return err -} - -func (l *LogReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return l.conn.LocalAddr() -} - -func (l *LogReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return l.conn.RemoteAddr() -} - -func (l *LogReadWriteCloserWithAddr) SocketID() string { - return l.sockid -} - -// NewLogRWC wraps ReadWriteCloser with logger calls. -func NewLogRWC(conn ReadWriteCloserWithAddr, logger *zap.SugaredLogger, sockid string, name string) ReadWriteCloserWithAddr { - return &LogReadWriteCloserWithAddr{ - conn: conn, - logger: logger, - sockid: sockid, - name: name, - } -} diff --git a/wrappers/mtproto_abridged.go b/wrappers/mtproto_abridged.go new file mode 100644 index 0000000..b9950c3 --- /dev/null +++ b/wrappers/mtproto_abridged.go @@ -0,0 +1,155 @@ +package wrappers + +import ( + "bytes" + "io" + "net" + + "github.com/juju/errors" + + "github.com/9seconds/mtg/mtproto" + "github.com/9seconds/mtg/utils" +) + +const ( + abridgedSmallPacketLength = 0x7f + abridgedQuickAckLength = 0x80 + abridgedLargePacketLength = 16777216 // 256 ^ 3 +) + +type MTProtoAbridged struct { + conn WrapStreamReadWriteCloser + opts *mtproto.ConnectionOpts + + readCounter uint32 + writeCounter uint32 +} + +func (m *MTProtoAbridged) Read() ([]byte, error) { + m.LogDebug("Read abridged packet", + "simple_ack", m.opts.WriteHacks.SimpleAck, + "quick_ack", m.opts.WriteHacks.QuickAck, + "counter", m.readCounter, + ) + + buf := &bytes.Buffer{} + buf.Grow(3) + + if _, err := io.CopyN(buf, m.conn, 1); err != nil { + return nil, errors.Annotate(err, "Cannot read message length") + } + msgLength := uint8(buf.Bytes()[0]) + buf.Reset() + + m.LogDebug("Abridged packet first byte", + "byte", msgLength, + "counter", m.readCounter, + ) + + if msgLength >= abridgedQuickAckLength { + m.opts.ReadHacks.QuickAck = true + msgLength -= abridgedQuickAckLength + } + + msgLength32 := uint32(msgLength) + if msgLength == abridgedSmallPacketLength { + if _, err := io.CopyN(buf, m.conn, 3); err != nil { + return nil, errors.Annotate(err, "Cannot read the correct message length") + } + number := utils.Uint24{} + copy(number[:], buf.Bytes()) + msgLength32 = utils.FromUint24(number) + } + msgLength32 *= 4 + + m.LogDebug("Abridged packet length", + "length", msgLength32, + "counter", m.readCounter, + ) + + buf.Reset() + buf.Grow(int(msgLength32)) + + if _, err := io.CopyN(buf, m.conn, int64(msgLength32)); err != nil { + return nil, errors.Annotate(err, "Cannot read message") + } + m.readCounter++ + + return buf.Bytes(), nil +} + +func (m *MTProtoAbridged) Write(p []byte) (int, error) { + m.LogDebug("Write abridged packet", + "length", len(p), + "simple_ack", m.opts.WriteHacks.SimpleAck, + "quick_ack", m.opts.WriteHacks.QuickAck, + "counter", m.writeCounter, + ) + + if len(p)%4 == 0 { + return 0, errors.Errorf("Incorrect packet length %d", len(p)) + } + + if m.opts.WriteHacks.SimpleAck { + return m.conn.Write(utils.ReverseBytes(p)) + } + + packetLength := len(p) / 4 + switch { + case packetLength < abridgedSmallPacketLength: + newData := append([]byte{byte(packetLength)}, p...) + + m.writeCounter++ + return m.conn.Write(newData) + + case packetLength < abridgedLargePacketLength: + length24 := utils.ToUint24(uint32(packetLength)) + + buf := &bytes.Buffer{} + buf.Grow(1 + 3 + len(p)) + + buf.WriteByte(byte(abridgedSmallPacketLength)) + buf.Write(length24[:]) + buf.Write(p) + + m.writeCounter++ + return m.conn.Write(buf.Bytes()) + } + + return 0, errors.Errorf("Packet is too big %d", len(p)) +} + +func (m *MTProtoAbridged) LogDebug(msg string, data ...interface{}) { + m.conn.LogDebug(msg, data...) +} + +func (m *MTProtoAbridged) LogInfo(msg string, data ...interface{}) { + m.conn.LogInfo(msg, data...) +} + +func (m *MTProtoAbridged) LogWarn(msg string, data ...interface{}) { + m.conn.LogWarn(msg, data...) +} + +func (m *MTProtoAbridged) LogError(msg string, data ...interface{}) { + m.conn.LogError(msg, data...) +} + +func (m *MTProtoAbridged) LocalAddr() *net.TCPAddr { + return m.conn.LocalAddr() +} + +func (m *MTProtoAbridged) RemoteAddr() *net.TCPAddr { + return m.conn.RemoteAddr() +} + +func (m *MTProtoAbridged) Close() error { + return m.conn.Close() +} + +func NewMTProtoAbridged(conn WrapStreamReadWriteCloser, opts *mtproto.ConnectionOpts) WrapPacketReadWriteCloser { + return &MTProtoAbridged{ + conn: conn, + opts: opts, + } +} diff --git a/wrappers/rwcaddr.go b/wrappers/rwcaddr.go deleted file mode 100644 index fee186c..0000000 --- a/wrappers/rwcaddr.go +++ /dev/null @@ -1,14 +0,0 @@ -package wrappers - -import ( - "io" - "net" -) - -type ReadWriteCloserWithAddr interface { - io.ReadWriteCloser - - LocalAddr() *net.TCPAddr - RemoteAddr() *net.TCPAddr - SocketID() string -} diff --git a/wrappers/streamcipher.go b/wrappers/streamcipher.go new file mode 100644 index 0000000..1ad764f --- /dev/null +++ b/wrappers/streamcipher.go @@ -0,0 +1,67 @@ +package wrappers + +import ( + "crypto/cipher" + "net" + + "github.com/juju/errors" +) + +type WrapStreamCipher struct { + encryptor cipher.Stream + decryptor cipher.Stream + conn WrapStreamReadWriteCloser +} + +func (w *WrapStreamCipher) Read(p []byte) (int, error) { + n, err := w.conn.Read(p) + if err != nil { + return 0, errors.Annotate(err, "Cannot read stream ciphered data") + } + w.decryptor.XORKeyStream(p, p[:n]) + + return n, nil +} + +func (w *WrapStreamCipher) Write(p []byte) (int, error) { + encrypted := make([]byte, len(p)) + w.encryptor.XORKeyStream(encrypted, p) + + return w.conn.Write(encrypted) +} + +func (w *WrapStreamCipher) LogDebug(msg string, data ...interface{}) { + w.conn.LogDebug(msg, data...) +} + +func (w *WrapStreamCipher) LogInfo(msg string, data ...interface{}) { + w.conn.LogInfo(msg, data...) +} + +func (w *WrapStreamCipher) LogWarn(msg string, data ...interface{}) { + w.conn.LogWarn(msg, data...) +} + +func (w *WrapStreamCipher) LogError(msg string, data ...interface{}) { + w.conn.LogError(msg, data...) +} + +func (w *WrapStreamCipher) LocalAddr() *net.TCPAddr { + return w.conn.LocalAddr() +} + +func (w *WrapStreamCipher) RemoteAddr() *net.TCPAddr { + return w.conn.RemoteAddr() +} + +func (w *WrapStreamCipher) Close() error { + return w.conn.Close() +} + +func NewStreamCipher(conn WrapStreamReadWriteCloser, encryptor, decryptor cipher.Stream) WrapStreamReadWriteCloser { + return &WrapStreamCipher{ + conn: conn, + encryptor: encryptor, + decryptor: decryptor, + } +} diff --git a/wrappers/streamcipherrwc.go b/wrappers/streamcipherrwc.go deleted file mode 100644 index f9ad598..0000000 --- a/wrappers/streamcipherrwc.go +++ /dev/null @@ -1,66 +0,0 @@ -package wrappers - -import ( - "crypto/cipher" - "net" -) - -// StreamCipherReadWriteCloser is a ReadWriteCloser which ciphers -// incoming and outgoing data with givem cipher.Stream instances. -type StreamCipherReadWriteCloserWithAddr struct { - encryptor cipher.Stream - decryptor cipher.Stream - conn ReadWriteCloserWithAddr -} - -// Read reads from connection -func (c *StreamCipherReadWriteCloserWithAddr) Read(p []byte) (n int, err error) { - n, err = c.conn.Read(p) - c.decryptor.XORKeyStream(p, p[:n]) - return -} - -// Write writes into connection. -func (c *StreamCipherReadWriteCloserWithAddr) Write(p []byte) (int, error) { - // This is to decrease an amount of allocations. Unfortunately, escape - // analysis in (at least Golang 1.10) is absolutely not perfect. For - // example, it understands that we want to have a slice locally, right? - // But since slice is effectively 2 ints + uintptr to [number]byte, the - // most heavyweight part is placed in heap. - buf := getBuffer() - defer putBuffer(buf) - buf.Grow(len(p)) - buf.Write(p) - - encrypted := buf.Bytes() - c.encryptor.XORKeyStream(encrypted, p) - - return c.conn.Write(encrypted) -} - -// Close closes underlying connection. -func (c *StreamCipherReadWriteCloserWithAddr) Close() error { - return c.conn.Close() -} - -func (c *StreamCipherReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return c.conn.LocalAddr() -} - -func (c *StreamCipherReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return c.conn.RemoteAddr() -} - -func (c *StreamCipherReadWriteCloserWithAddr) SocketID() string { - return c.conn.SocketID() -} - -// NewStreamCipherRWC returns wrapper which transparently -// encrypts/decrypts traffic with obfuscated2 protocol. -func NewStreamCipherRWC(conn ReadWriteCloserWithAddr, encryptor, decryptor cipher.Stream) ReadWriteCloserWithAddr { - return &StreamCipherReadWriteCloserWithAddr{ - conn: conn, - encryptor: encryptor, - decryptor: decryptor, - } -} diff --git a/wrappers/timeoutrwc.go b/wrappers/timeoutrwc.go deleted file mode 100644 index 3c48546..0000000 --- a/wrappers/timeoutrwc.go +++ /dev/null @@ -1,61 +0,0 @@ -package wrappers - -import ( - "net" - "time" - - "github.com/9seconds/mtg/config" -) - -type TimeoutReadWriteCloserWithAddr struct { - conn net.Conn - sock string - publicIPv4 net.IP - publicIPv6 net.IP -} - -func (t *TimeoutReadWriteCloserWithAddr) Read(p []byte) (int, error) { - t.conn.SetReadDeadline(time.Now().Add(config.TimeoutRead)) - return t.conn.Read(p) -} - -func (t *TimeoutReadWriteCloserWithAddr) Write(p []byte) (int, error) { - t.conn.SetWriteDeadline(time.Now().Add(config.TimeoutWrite)) - return t.conn.Write(p) -} - -func (t *TimeoutReadWriteCloserWithAddr) Close() error { - return t.conn.Close() -} - -func (t *TimeoutReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return t.conn.RemoteAddr().(*net.TCPAddr) -} - -func (t *TimeoutReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - addr := t.conn.LocalAddr().(*net.TCPAddr) - newAddr := *addr - - if t.RemoteAddr().IP.To4() != nil { - if t.publicIPv4 != nil { - newAddr.IP = t.publicIPv4 - } - } else if t.publicIPv6 != nil { - newAddr.IP = t.publicIPv6 - } - - return &newAddr -} - -func (t *TimeoutReadWriteCloserWithAddr) SocketID() string { - return t.sock -} - -func NewTimeoutRWC(conn net.Conn, sock string, ipv4, ipv6 net.IP) ReadWriteCloserWithAddr { - return &TimeoutReadWriteCloserWithAddr{ - conn: conn, - publicIPv4: ipv4, - publicIPv6: ipv6, - sock: sock, - } -} diff --git a/wrappers/trafficrwc.go b/wrappers/trafficrwc.go deleted file mode 100644 index 344a451..0000000 --- a/wrappers/trafficrwc.go +++ /dev/null @@ -1,51 +0,0 @@ -package wrappers - -import "net" - -// TrafficReadWriteCloser counts an amount of ingress/egress traffic by -// calling given callbacks. -type TrafficReadWriteCloserWithAddr struct { - conn ReadWriteCloserWithAddr - readCallback func(int) - writeCallback func(int) -} - -// Read reads from connection -func (t *TrafficReadWriteCloserWithAddr) Read(p []byte) (n int, err error) { - n, err = t.conn.Read(p) - t.readCallback(n) - return -} - -// Write writes into connection. -func (t *TrafficReadWriteCloserWithAddr) Write(p []byte) (n int, err error) { - n, err = t.conn.Write(p) - t.writeCallback(n) - return -} - -// Close closes underlying connection. -func (t *TrafficReadWriteCloserWithAddr) Close() error { - return t.conn.Close() -} - -func (t *TrafficReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return t.conn.LocalAddr() -} - -func (t *TrafficReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return t.conn.RemoteAddr() -} - -func (t *TrafficReadWriteCloserWithAddr) SocketID() string { - return t.conn.SocketID() -} - -// NewTrafficRWC wraps ReadWriteCloser to have read/write callbacks. -func NewTrafficRWC(conn ReadWriteCloserWithAddr, readCallback, writeCallback func(int)) ReadWriteCloserWithAddr { - return &TrafficReadWriteCloserWithAddr{ - conn: conn, - readCallback: readCallback, - writeCallback: writeCallback, - } -} diff --git a/wrappers/wrap.go b/wrappers/wrap.go new file mode 100644 index 0000000..d4bc7d4 --- /dev/null +++ b/wrappers/wrap.go @@ -0,0 +1,78 @@ +package wrappers + +import ( + "io" + "net" +) + +type Wrap interface { + LogDebug(msg string, data ...interface{}) + LogInfo(msg string, data ...interface{}) + LogWarn(msg string, data ...interface{}) + LogError(msg string, data ...interface{}) + + LocalAddr() *net.TCPAddr + RemoteAddr() *net.TCPAddr +} + +type WrapWriter interface { + io.Writer + Wrap +} + +type WrapWriteCloser interface { + io.Closer + WrapWriter +} + +type WrapStreamReader interface { + io.Reader + Wrap +} + +type WrapStreamReadCloser interface { + io.Closer + WrapStreamReader +} + +type WrapStreamReadWriter interface { + io.Writer + WrapStreamReader +} + +type WrapStreamWriteCloser interface { + io.Closer + io.Writer + Wrap +} + +type WrapStreamReadWriteCloser interface { + io.Closer + WrapStreamReadWriter +} + +type WrapPacketReader interface { + Read() ([]byte, error) + Wrap +} + +type WrapPacketReadWriter interface { + io.Writer + WrapPacketReader +} + +type WrapBlockReadCloser interface { + io.Closer + WrapPacketReader +} + +type WrapPacketWriteCloser interface { + io.Writer + io.Closer + Wrap +} + +type WrapPacketReadWriteCloser interface { + io.Closer + WrapPacketReadWriter +}