diff --git a/client/client.go b/client/client.go index cb6b291..648fec2 100644 --- a/client/client.go +++ b/client/client.go @@ -8,5 +8,4 @@ import ( "github.com/9seconds/mtg/wrappers" ) -// Init has to initialize client connection based on given config. -type Init func(net.Conn, *config.Config) (*mtproto.ConnectionOpts, wrappers.ReadWriteCloserWithAddr, error) +type Init func(net.Conn, string, *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) diff --git a/client/direct.go b/client/direct.go index 98b518d..9f0df10 100644 --- a/client/direct.go +++ b/client/direct.go @@ -14,28 +14,29 @@ import ( const handshakeTimeout = 10 * time.Second -// DirectInit initializes client to access Telegram bypassing middleproxies. -func DirectInit(conn net.Conn, conf *config.Config) (*mtproto.ConnectionOpts, wrappers.ReadWriteCloserWithAddr, error) { - if err := config.SetSocketOptions(conn); err != nil { +func DirectInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) { + if err := config.SetSocketOptions(socket); err != nil { return nil, nil, errors.Annotate(err, "Cannot set socket options") } - conn.SetReadDeadline(time.Now().Add(handshakeTimeout)) // nolint: errcheck - frame, err := obfuscated2.ExtractFrame(conn) - conn.SetReadDeadline(time.Time{}) // nolint: errcheck + socket.SetReadDeadline(time.Now().Add(handshakeTimeout)) + frame, err := obfuscated2.ExtractFrame(socket) if err != nil { return nil, nil, errors.Annotate(err, "Cannot extract frame") } + socket.SetReadDeadline(time.Time{}) + conn := wrappers.NewConn(socket, connID, wrappers.ConnPurposeClient, conf.PublicIPv4, conf.PublicIPv6) obfs2, connOpts, err := obfuscated2.ParseObfuscated2ClientFrame(conf.Secret, frame) if err != nil { return nil, nil, errors.Annotate(err, "Cannot parse obfuscated frame") } connOpts.ConnectionProto = mtproto.ConnectionProtocolAny - connOpts.ClientAddr = conn.RemoteAddr().(*net.TCPAddr) + connOpts.ClientAddr = conn.RemoteAddr() - socket := wrappers.NewTimeoutRWC(conn, conf.PublicIPv4, conf.PublicIPv6) - socket = wrappers.NewStreamCipherRWC(socket, obfs2.Encryptor, obfs2.Decryptor) + conn = wrappers.NewStreamCipher(conn, obfs2.Encryptor, obfs2.Decryptor) - return connOpts, socket, nil + conn.Logger().Infow("Client connection initialized") + + return conn, connOpts, nil } diff --git a/client/middle.go b/client/middle.go index d5dc4d3..b85deef 100644 --- a/client/middle.go +++ b/client/middle.go @@ -5,21 +5,25 @@ import ( "github.com/9seconds/mtg/config" "github.com/9seconds/mtg/mtproto" - mtwrappers "github.com/9seconds/mtg/mtproto/wrappers" "github.com/9seconds/mtg/wrappers" ) -func MiddleInit(conn net.Conn, conf *config.Config) (*mtproto.ConnectionOpts, wrappers.ReadWriteCloserWithAddr, error) { - opts, newConn, err := DirectInit(conn, conf) +func MiddleInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) { + conn, opts, err := DirectInit(socket, connID, conf) if err != nil { return nil, nil, err } + connStream := conn.(wrappers.StreamReadWriteCloser) - if opts.ConnectionType == mtproto.ConnectionTypeAbridged { - newConn = mtwrappers.NewAbridgedRWC(newConn, opts) - } else { - newConn = mtwrappers.NewIntermediateRWC(newConn, opts) + newConn := wrappers.NewMTProtoAbridged(connStream, opts) + if opts.ConnectionType != mtproto.ConnectionTypeAbridged { + newConn = wrappers.NewMTProtoIntermediate(connStream, opts) } - return opts, newConn, nil + opts.ConnectionProto = mtproto.ConnectionProtocolIPv4 + if socket.LocalAddr().(*net.TCPAddr).IP.To4() == nil { + opts.ConnectionProto = mtproto.ConnectionProtocolIPv6 + } + + return newConn, opts, err } diff --git a/config/config.go b/config/config.go index b5cfe9e..59d06cd 100644 --- a/config/config.go +++ b/config/config.go @@ -5,7 +5,6 @@ import ( "fmt" "net" "strconv" - "time" "github.com/juju/errors" ) @@ -14,10 +13,6 @@ import ( const ( BufferWriteSize = 32 * 1024 BufferReadSize = 32 * 1024 - BufferSizeCopy = 32 * 1024 - - TimeoutRead = time.Minute - TimeoutWrite = time.Minute ) // Config represents common configuration of mtg. @@ -63,6 +58,10 @@ func (c *Config) StatAddr() string { return getAddr(c.StatsIP, c.StatsPort) } +func (c *Config) UseMiddleProxy() bool { + return len(c.AdTag) > 0 +} + // GetURLs returns configured IPURLs instance with links to this server. func (c *Config) GetURLs() IPURLs { urls := IPURLs{} diff --git a/main.go b/main.go index 14cf85a..c0ffba1 100644 --- a/main.go +++ b/main.go @@ -111,16 +111,21 @@ func main() { zapcore.NewJSONEncoder(encoderCfg), zapcore.Lock(os.Stderr), atom, - )).Sugar() + )) + zap.ReplaceGlobals(logger) + defer logger.Sync() - stat := proxy.NewStats(conf) - go stat.Serve() + if conf.UseMiddleProxy() { + zap.S().Infow("Use middle proxy connection to Telegram") + } else { + zap.S().Infow("Use direct connection to Telegram") + } - srv := proxy.NewServer(conf, logger, stat) printURLs(conf.GetURLs()) - if err := srv.Serve(); err != nil { - logger.Fatal(err.Error()) + server := proxy.NewProxy(conf) + if err := server.Serve(); err != nil { + zap.S().Fatalw("Server stopped", "error", err) } } diff --git a/mtproto/connection_options.go b/mtproto/connection_options.go index 585e0f7..0285488 100644 --- a/mtproto/connection_options.go +++ b/mtproto/connection_options.go @@ -13,14 +13,19 @@ type ConnectionType uint8 type ConnectionProtocol uint8 +type Hacks struct { + SimpleAck bool + QuickAck bool +} + // ConnectionOpts presents an options, metadata on connection requested // by the user on handshake. type ConnectionOpts struct { DC int16 ConnectionType ConnectionType ConnectionProto ConnectionProtocol - QuickAck bool - SimpleAck bool + ReadHacks Hacks + WriteHacks Hacks ClientAddr *net.TCPAddr } diff --git a/mtproto/rpc/handshake_request.go b/mtproto/rpc/handshake_request.go new file mode 100644 index 0000000..924859d --- /dev/null +++ b/mtproto/rpc/handshake_request.go @@ -0,0 +1,22 @@ +package rpc + +import "bytes" + +type HandshakeRequest struct { +} + +func (r *HandshakeRequest) Bytes() []byte { + buf := &bytes.Buffer{} + buf.Grow(len(TagHandshake) + len(HandshakeFlags) + len(HandshakeSenderPID) + len(HandshakePeerPID)) + + buf.Write(TagHandshake) + buf.Write(HandshakeFlags) + buf.Write(HandshakeSenderPID) + buf.Write(HandshakePeerPID) + + return buf.Bytes() +} + +func NewHandshakeRequest() *HandshakeRequest { + return &HandshakeRequest{} +} diff --git a/mtproto/rpc/rpc_handshake_response.go b/mtproto/rpc/handshake_response.go similarity index 62% rename from mtproto/rpc/rpc_handshake_response.go rename to mtproto/rpc/handshake_response.go index af74469..40534d4 100644 --- a/mtproto/rpc/rpc_handshake_response.go +++ b/mtproto/rpc/handshake_response.go @@ -6,14 +6,14 @@ import ( "github.com/juju/errors" ) -type RPCHandshakeResponse struct { +type HandshakeResponse struct { Type []byte Flags []byte SenderPID []byte PeerPID []byte } -func (r *RPCHandshakeResponse) Bytes() []byte { +func (r *HandshakeResponse) Bytes() []byte { buf := &bytes.Buffer{} buf.Write(r.Type[:]) @@ -24,23 +24,23 @@ func (r *RPCHandshakeResponse) Bytes() []byte { return buf.Bytes() } -func (r *RPCHandshakeResponse) Valid(req *RPCHandshakeRequest) error { - if !bytes.Equal(r.Type, RPCTagHandshake) { +func (r *HandshakeResponse) Valid(req *HandshakeRequest) error { + if !bytes.Equal(r.Type, TagHandshake) { return errors.New("Unexpected handshake tag") } - if !bytes.Equal(r.PeerPID, RPCHandshakeSenderPID) { + if !bytes.Equal(r.PeerPID, HandshakeSenderPID) { return errors.New("Incorrect sender PID") } return nil } -func NewRPCHandshakeResponse(data []byte) (*RPCHandshakeResponse, error) { +func NewHandshakeResponse(data []byte) (*HandshakeResponse, error) { if len(data) != 32 { return nil, errors.New("Incorrect handshake response length") } - return &RPCHandshakeResponse{ + return &HandshakeResponse{ Type: data[:4], Flags: data[4:8], SenderPID: data[8:20], diff --git a/mtproto/rpc/rpc_nonce_request.go b/mtproto/rpc/nonce_request.go similarity index 77% rename from mtproto/rpc/rpc_nonce_request.go rename to mtproto/rpc/nonce_request.go index 55f5543..e4fb8a0 100644 --- a/mtproto/rpc/rpc_nonce_request.go +++ b/mtproto/rpc/nonce_request.go @@ -9,25 +9,25 @@ import ( "github.com/juju/errors" ) -type RPCNonceRequest struct { +type NonceRequest struct { KeySelector []byte CryptoTS []byte Nonce []byte } -func (r *RPCNonceRequest) Bytes() []byte { +func (r *NonceRequest) Bytes() []byte { buf := &bytes.Buffer{} - buf.Write(RPCTagNonce) + buf.Write(TagNonce) buf.Write(r.KeySelector) - buf.Write(RPCNonceCryptoAES) + buf.Write(NonceCryptoAES) buf.Write(r.CryptoTS) buf.Write(r.Nonce) return buf.Bytes() } -func NewRPCNonceRequest(proxySecret []byte) (*RPCNonceRequest, error) { +func NewNonceRequest(proxySecret []byte) (*NonceRequest, error) { nonce := make([]byte, 16) keySelector := make([]byte, 4) cryptoTS := make([]byte, 4) @@ -40,7 +40,7 @@ func NewRPCNonceRequest(proxySecret []byte) (*RPCNonceRequest, error) { timestamp := time.Now().Truncate(time.Second).Unix() % 4294967296 // 256 ^ 4 - do not know how to name binary.LittleEndian.PutUint32(cryptoTS, uint32(timestamp)) - return &RPCNonceRequest{ + return &NonceRequest{ KeySelector: keySelector, CryptoTS: cryptoTS, Nonce: nonce, diff --git a/mtproto/rpc/rpc_nonce_response.go b/mtproto/rpc/nonce_response.go similarity index 55% rename from mtproto/rpc/rpc_nonce_response.go rename to mtproto/rpc/nonce_response.go index 4f753b6..85c9383 100644 --- a/mtproto/rpc/rpc_nonce_response.go +++ b/mtproto/rpc/nonce_response.go @@ -6,17 +6,17 @@ import ( "github.com/juju/errors" ) -type RPCNonceResponse struct { - RPCNonceRequest +type NonceResponse struct { + NonceRequest - RPCType []byte - Crypto []byte + Type []byte + Crypto []byte } -func (r *RPCNonceResponse) Bytes() []byte { +func (r *NonceResponse) Bytes() []byte { buf := &bytes.Buffer{} - buf.Write(r.RPCType) + buf.Write(r.Type) buf.Write(r.KeySelector) buf.Write(r.Crypto) buf.Write(r.CryptoTS) @@ -25,11 +25,11 @@ func (r *RPCNonceResponse) Bytes() []byte { return buf.Bytes() } -func (r *RPCNonceResponse) Valid(req *RPCNonceRequest) error { - if !bytes.Equal(r.RPCType, RPCTagNonce) { +func (r *NonceResponse) Valid(req *NonceRequest) error { + if !bytes.Equal(r.Type, TagNonce) { return errors.New("Unexpected RPC type") } - if !bytes.Equal(r.Crypto, RPCNonceCryptoAES) { + if !bytes.Equal(r.Crypto, NonceCryptoAES) { return errors.New("Unexpected crypto type") } if !bytes.Equal(r.KeySelector, req.KeySelector) { @@ -39,18 +39,18 @@ func (r *RPCNonceResponse) Valid(req *RPCNonceRequest) error { return nil } -func NewRPCNonceResponse(data []byte) (*RPCNonceResponse, error) { +func NewNonceResponse(data []byte) (*NonceResponse, error) { if len(data) != 32 { return nil, errors.New("Unexpected message length") } - return &RPCNonceResponse{ - RPCNonceRequest: RPCNonceRequest{ + return &NonceResponse{ + NonceRequest: NonceRequest{ KeySelector: data[4:8], CryptoTS: data[12:16], Nonce: data[16:], }, - RPCType: data[:4], - Crypto: data[8:12], + Type: data[:4], + Crypto: data[8:12], }, nil } diff --git a/mtproto/rpc/proxy_flags.go b/mtproto/rpc/proxy_flags.go new file mode 100644 index 0000000..5599537 --- /dev/null +++ b/mtproto/rpc/proxy_flags.go @@ -0,0 +1,55 @@ +package rpc + +import ( + "encoding/binary" + "strings" +) + +type proxyRequestFlags uint32 + +const ( + proxyRequestFlagsHasAdTag proxyRequestFlags = 0x8 + proxyRequestFlagsEncrypted = 0x2 + proxyRequestFlagsMagic = 0x1000 + proxyRequestFlagsExtMode2 = 0x20000 + proxyRequestFlagsIntermediate = 0x20000000 + proxyRequestFlagsAbdridged = 0x40000000 + proxyRequestFlagsQuickAck = 0x80000000 +) + +var proxyRequestFlagsEncryptedPrefix [8]byte + +func (r proxyRequestFlags) Bytes() []byte { + converted := make([]byte, 4) + binary.LittleEndian.PutUint32(converted, uint32(r)) + + return converted +} + +func (r proxyRequestFlags) String() string { + flags := make([]string, 0, 7) + + if r&proxyRequestFlagsHasAdTag != 0 { + flags = append(flags, "HAS_AD_TAG") + } + if r&proxyRequestFlagsEncrypted != 0 { + flags = append(flags, "ENCRYPTED") + } + if r&proxyRequestFlagsMagic != 0 { + flags = append(flags, "MAGIC") + } + if r&proxyRequestFlagsExtMode2 != 0 { + flags = append(flags, "EXT_MODE_2") + } + if r&proxyRequestFlagsIntermediate != 0 { + flags = append(flags, "INTERMEDIATE") + } + if r&proxyRequestFlagsAbdridged != 0 { + flags = append(flags, "ABRIDGED") + } + if r&proxyRequestFlagsQuickAck != 0 { + flags = append(flags, "QUICK_ACK") + } + + return strings.Join(flags, " | ") +} diff --git a/mtproto/rpc/proxy_request.go b/mtproto/rpc/proxy_request.go new file mode 100644 index 0000000..40ab3cf --- /dev/null +++ b/mtproto/rpc/proxy_request.go @@ -0,0 +1,95 @@ +package rpc + +import ( + "bytes" + "crypto/rand" + "encoding/binary" + "fmt" + "net" + + "github.com/juju/errors" + + "github.com/9seconds/mtg/mtproto" +) + +type ProxyRequest struct { + Flags proxyRequestFlags + ConnectionID []byte + OurIPPort []byte + ClientIPPort []byte + ADTag []byte + Options *mtproto.ConnectionOpts +} + +func (r *ProxyRequest) MakeHeader(message []byte) (*bytes.Buffer, fmt.Stringer) { + bufferLength := len(TagProxyRequest) + + 4 + // len(flags) + len(r.ConnectionID) + + len(r.ClientIPPort) + + len(r.OurIPPort) + + len(ProxyRequestExtraSize) + + len(ProxyRequestProxyTag) + + 1 + // len(AdTag) + len(r.ADTag) + bufferLength += bufferLength % 4 + + buf := &bytes.Buffer{} + buf.Grow(bufferLength + len(message)) + + flags := r.Flags + if r.Options.ReadHacks.QuickAck { + flags |= proxyRequestFlagsQuickAck + } + + if bytes.HasPrefix(message, proxyRequestFlagsEncryptedPrefix[:]) { + flags |= proxyRequestFlagsEncrypted + } + + buf.Write(TagProxyRequest) + buf.Write(flags.Bytes()) + buf.Write(r.ConnectionID) + buf.Write(r.ClientIPPort) + buf.Write(r.OurIPPort) + buf.Write(ProxyRequestExtraSize) + buf.Write(ProxyRequestProxyTag) + buf.WriteByte(byte(len(r.ADTag))) + buf.Write(r.ADTag) + buf.Write(make([]byte, (4-buf.Len()%4)%4)) + + return buf, flags +} + +func NewProxyRequest(clientAddr, ownAddr *net.TCPAddr, opts *mtproto.ConnectionOpts, adTag []byte) (*ProxyRequest, error) { + flags := proxyRequestFlagsHasAdTag | proxyRequestFlagsMagic | proxyRequestFlagsExtMode2 + + switch opts.ConnectionType { + case mtproto.ConnectionTypeAbridged: + flags |= proxyRequestFlagsAbdridged + case mtproto.ConnectionTypeIntermediate: + flags |= proxyRequestFlagsIntermediate + } + + request := &ProxyRequest{ + Flags: flags, + ADTag: adTag, + Options: opts, + ConnectionID: make([]byte, 8), + ClientIPPort: make([]byte, 16+4), + OurIPPort: make([]byte, 16+4), + } + + if _, err := rand.Read(request.ConnectionID); err != nil { + return nil, errors.Annotate(err, "Cannot generate connection ID") + } + + port := [4]byte{} + copy(request.ClientIPPort[:16], clientAddr.IP.To16()) + binary.LittleEndian.PutUint32(port[:], uint32(clientAddr.Port)) + copy(request.ClientIPPort[16:], port[:]) + + copy(request.OurIPPort[:16], ownAddr.IP.To16()) + binary.LittleEndian.PutUint32(port[:], uint32(ownAddr.Port)) + copy(request.OurIPPort[16:], port[:]) + + return request, nil +} diff --git a/mtproto/rpc/rpc.go b/mtproto/rpc/rpc.go index f4edf71..bfb0555 100644 --- a/mtproto/rpc/rpc.go +++ b/mtproto/rpc/rpc.go @@ -1,30 +1,30 @@ package rpc const ( - RPCNonceSeqNo = -2 - RPCHandshakeSeqNo = -1 + SeqNoNonce = -2 + SeqNoHandshake = -1 ) var ( - RPCTagCloseExt = []byte{0xa2, 0x34, 0xb6, 0x5e} - RPCTagProxyAns = []byte{0x0d, 0xda, 0x03, 0x44} - RPCTagSimpleAck = []byte{0x9b, 0x40, 0xac, 0x3b} - RPCTagHandshake = []byte{0xf5, 0xee, 0x82, 0x76} - RPCTagNonce = []byte{0xaa, 0x87, 0xcb, 0x7a} - RPCTagProxyRequest = []byte{0xee, 0xf1, 0xce, 0x36} + TagCloseExt = []byte{0xa2, 0x34, 0xb6, 0x5e} + TagProxyAns = []byte{0x0d, 0xda, 0x03, 0x44} + TagSimpleAck = []byte{0x9b, 0x40, 0xac, 0x3b} + TagHandshake = []byte{0xf5, 0xee, 0x82, 0x76} + TagNonce = []byte{0xaa, 0x87, 0xcb, 0x7a} + TagProxyRequest = []byte{0xee, 0xf1, 0xce, 0x36} - RPCNonceCryptoAES = []byte{0x01, 0x00, 0x00, 0x00} + NonceCryptoAES = []byte{0x01, 0x00, 0x00, 0x00} - RPCHandshakeFlags = []byte{0x00, 0x00, 0x00, 0x00} + HandshakeFlags = []byte{0x00, 0x00, 0x00, 0x00} - RPCProxyRequestExtraSize = []byte{0x18, 0x00, 0x00, 0x00} - RPCProxyRequestProxyTag = []byte{0xae, 0x26, 0x1e, 0xdb} + ProxyRequestExtraSize = []byte{0x18, 0x00, 0x00, 0x00} + ProxyRequestProxyTag = []byte{0xae, 0x26, 0x1e, 0xdb} - RPCHandshakeSenderPID = []byte{} - RPCHandshakePeerPID = []byte{} + HandshakeSenderPID []byte + HandshakePeerPID []byte ) func init() { - RPCHandshakeSenderPID = []byte("IPIPPRPDTIME") - RPCHandshakePeerPID = []byte("IPIPPRPDTIME") + HandshakeSenderPID = []byte("IPIPPRPDTIME") + HandshakePeerPID = []byte("IPIPPRPDTIME") } diff --git a/mtproto/rpc/rpc_handshake_request.go b/mtproto/rpc/rpc_handshake_request.go deleted file mode 100644 index 4a21fe4..0000000 --- a/mtproto/rpc/rpc_handshake_request.go +++ /dev/null @@ -1,21 +0,0 @@ -package rpc - -import "bytes" - -type RPCHandshakeRequest struct { -} - -func (r *RPCHandshakeRequest) Bytes() []byte { - buf := &bytes.Buffer{} - - buf.Write(RPCTagHandshake) - buf.Write(RPCHandshakeFlags) - buf.Write(RPCHandshakeSenderPID) - buf.Write(RPCHandshakePeerPID) - - return buf.Bytes() -} - -func NewRPCHandshakeRequest() *RPCHandshakeRequest { - return &RPCHandshakeRequest{} -} diff --git a/mtproto/rpc/rpc_proxy_flags.go b/mtproto/rpc/rpc_proxy_flags.go deleted file mode 100644 index 4edd74d..0000000 --- a/mtproto/rpc/rpc_proxy_flags.go +++ /dev/null @@ -1,24 +0,0 @@ -package rpc - -import "encoding/binary" - -type RPCProxyRequestFlags uint32 - -const ( - RPCProxyRequestFlagsHasAdTag RPCProxyRequestFlags = 0x8 - RPCProxyRequestFlagsEncrypted = 0x2 - RPCProxyRequestFlagsMagic = 0x1000 - RPCProxyRequestFlagsExtMode2 = 0x20000 - RPCProxyRequestFlagsIntermediate = 0x20000000 - RPCProxyRequestFlagsAbdridged = 0x40000000 - RPCProxyRequestFlagsQuickAck = 0x80000000 -) - -var rpcProxyRequestFlagsEncryptedPrefix [8]byte - -func (r RPCProxyRequestFlags) Bytes() []byte { - converted := make([]byte, 4) - binary.LittleEndian.PutUint32(converted, uint32(r)) - - return converted -} diff --git a/mtproto/rpc/rpc_proxy_request.go b/mtproto/rpc/rpc_proxy_request.go deleted file mode 100644 index d2b8392..0000000 --- a/mtproto/rpc/rpc_proxy_request.go +++ /dev/null @@ -1,83 +0,0 @@ -package rpc - -import ( - "bytes" - "crypto/rand" - "encoding/binary" - "net" - - "github.com/juju/errors" - - "github.com/9seconds/mtg/mtproto" -) - -type RPCProxyRequest struct { - Flags RPCProxyRequestFlags - ConnectionID []byte - OurIPPort []byte - ClientIPPort []byte - ADTag []byte - Options *mtproto.ConnectionOpts -} - -func (r *RPCProxyRequest) Bytes(message []byte) []byte { - buf := &bytes.Buffer{} - - flags := r.Flags - if r.Options.QuickAck { - flags |= RPCProxyRequestFlagsQuickAck - } - - if bytes.HasPrefix(message, rpcProxyRequestFlagsEncryptedPrefix[:]) { - flags |= RPCProxyRequestFlagsEncrypted - } - - buf.Write(RPCTagProxyRequest) - buf.Write(flags.Bytes()) - buf.Write(r.ConnectionID[:]) - buf.Write(r.ClientIPPort[:]) - buf.Write(r.OurIPPort[:]) - buf.Write(RPCProxyRequestExtraSize) - buf.Write(RPCProxyRequestProxyTag) - buf.WriteByte(byte(len(r.ADTag))) - buf.Write(r.ADTag) - buf.Write(bytes.Repeat([]byte{0x00}, buf.Len()%4)) - buf.Write(message) - - return buf.Bytes() -} - -func NewRPCProxyRequest(clientAddr, ownAddr *net.TCPAddr, opts *mtproto.ConnectionOpts, adTag []byte) (*RPCProxyRequest, error) { - flags := RPCProxyRequestFlagsHasAdTag | RPCProxyRequestFlagsMagic | RPCProxyRequestFlagsExtMode2 - - switch opts.ConnectionType { - case mtproto.ConnectionTypeAbridged: - flags |= RPCProxyRequestFlagsAbdridged - case mtproto.ConnectionTypeIntermediate: - flags |= RPCProxyRequestFlagsIntermediate - } - - request := RPCProxyRequest{ - Flags: flags, - ADTag: adTag, - Options: opts, - ConnectionID: make([]byte, 8), - ClientIPPort: make([]byte, 16+4), - OurIPPort: make([]byte, 16+4), - } - - if _, err := rand.Read(request.ConnectionID); err != nil { - return nil, errors.Annotate(err, "Cannot generate connection ID") - } - - port := [4]byte{} - copy(request.ClientIPPort[:16], clientAddr.IP.To16()) - binary.LittleEndian.PutUint32(port[:], uint32(clientAddr.Port)) - copy(request.ClientIPPort[16:], port[:]) - - copy(request.OurIPPort[:16], ownAddr.IP.To16()) - binary.LittleEndian.PutUint32(port[:], uint32(ownAddr.Port)) - copy(request.OurIPPort[16:], port[:]) - - return &request, nil -} diff --git a/mtproto/wrappers/abridged.go b/mtproto/wrappers/abridged.go deleted file mode 100644 index 66b440d..0000000 --- a/mtproto/wrappers/abridged.go +++ /dev/null @@ -1,121 +0,0 @@ -package wrappers - -import ( - "bytes" - "encoding/binary" - "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 { - var msgLength uint8 - if err := binary.Read(a.conn, binary.LittleEndian, &msgLength); err != nil { - return errors.Annotate(err, "Cannot read message length") - } - - a.opts.QuickAck = false - if msgLength >= abridgedQuickAckLength { - a.opts.QuickAck = true - msgLength -= 0x80 - } - - msgLength32 := uint32(msgLength) - if msgLength == abridgedSmallPacketLength { - buf := &bytes.Buffer{} - buf.Grow(3) - - if _, err := io.CopyN(buf, a.conn, 3); err != nil { - return errors.Annotate(err, "Cannot read correct message length") - } - number := uint24{} - copy(number[:], buf.Bytes()) - msgLength32 = fromUint24(number) - } - msgLength32 *= 4 - - if _, err := io.CopyN(a.Buffer, a.conn, int64(msgLength32)); err != nil { - return errors.Annotate(err, "Cannot read message") - } - - 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.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()) - - default: - 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 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/mtproto/wrappers/crypt_test.go b/mtproto/wrappers/crypt_test.go deleted file mode 100644 index 62bf0b9..0000000 --- a/mtproto/wrappers/crypt_test.go +++ /dev/null @@ -1,45 +0,0 @@ -package wrappers - -import ( - "encoding/binary" - "net" - "testing" - - "github.com/stretchr/testify/assert" - - "github.com/9seconds/mtg/mtproto/rpc" -) - -var proxySecret = []byte{196, 249, 250, 202, 150, 120, 230, 187, 72, 173, - 108, 126, 44, 229, 192, 210, 68, 48, 100, 93, 85, 74, 221, 235, 85, 65, - 158, 3, 77, 166, 39, 33, 208, 70, 234, 171, 110, 82, 171, 20, 169, 90, 68, - 62, 207, 179, 70, 62, 121, 160, 90, 102, 97, 42, 223, 156, 174, 218, 139, - 233, 168, 13, 166, 152, 111, 176, 166, 255, 56, 122, 248, 77, 136, 239, - 58, 100, 19, 113, 62, 92, 51, 119, 246, 225, 163, 212, 125, 153, 245, 224, - 197, 110, 236, 232, 240, 92, 84, 196, 144, 176, 121, 227, 27, 239, 130, - 255, 14, 232, 242, 176, 163, 39, 86, 210, 73, 197, 242, 18, 105, 129, 108, - 183, 6, 27, 38, 93, 178, 18} - -func TestMakeKeys(t *testing.T) { - req, err := rpc.NewRPCNonceRequest(proxySecret) - assert.Nil(t, err) - - copy(req.Nonce[:], []byte{24, 49, 53, 111, 198, 10, 235, 180, 230, 112, 92, 78, 1, 201, 106, 105}) - binary.LittleEndian.PutUint32(req.CryptoTS[:], 1528396015) - - resp := &rpc.RPCNonceResponse{} - copy(resp.Nonce[:], []byte{247, 40, 210, 56, 65, 12, 101, 170, 216, 155, 14, 253, 250, 238, 219, 226}) - - cltAddr := &net.TCPAddr{ - IP: net.ParseIP("80.211.29.34"), - Port: 54208, - } - srvAddr := &net.TCPAddr{ - IP: net.ParseIP("149.154.162.38"), - Port: 80, - } - - key, iv := makeKeys(CipherPurposeClient, req, resp, cltAddr, srvAddr, proxySecret) - assert.Equal(t, key, []byte{165, 158, 127, 49, 41, 232, 187, 69, 38, 29, 163, 226, 183, 146, 28, 67, 225, 224, 134, 191, 207, 152, 255, 166, 152, 66, 169, 196, 54, 135, 50, 188}) - assert.Equal(t, iv, []byte{33, 110, 125, 221, 183, 121, 160, 116, 130, 180, 156, 249, 52, 111, 37, 178}) -} diff --git a/mtproto/wrappers/frame.go b/mtproto/wrappers/frame.go deleted file mode 100644 index d008d37..0000000 --- a/mtproto/wrappers/frame.go +++ /dev/null @@ -1,130 +0,0 @@ -package wrappers - -import ( - "bytes" - "crypto/aes" - "encoding/binary" - "hash/crc32" - "io" - "io/ioutil" - "net" - - "github.com/juju/errors" - - "github.com/9seconds/mtg/wrappers" -) - -// Frame: { MessageLength(4) | SequenceNumber(4) | Message(???) | CRC32(4) [| padding(4), ...] } -const ( - frameRWCMinMessageLength = 12 - frameRWCMaxMessageLength = 16777216 -) - -var frameRWCPadding = [4]byte{0x04, 0x00, 0x00, 0x00} - -type FrameRWC struct { - wrappers.BufferedReader - - conn wrappers.ReadWriteCloserWithAddr - readSeqNo int32 - writeSeqNo int32 -} - -func (f *FrameRWC) Write(buf []byte) (int, error) { - writeBuf := &bytes.Buffer{} - - // 4 - len bytes - // 4 - seq bytes - // . - message - // 4 - crc32 - messageLength := 4 + 4 + len(buf) + 4 - paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize - writeBuf.Grow(messageLength + paddingLength) - - binary.Write(writeBuf, binary.LittleEndian, uint32(messageLength)) - binary.Write(writeBuf, binary.LittleEndian, f.writeSeqNo) - writeBuf.Write(buf) - f.writeSeqNo++ - - checksum := crc32.ChecksumIEEE(writeBuf.Bytes()) - binary.Write(writeBuf, binary.LittleEndian, checksum) - writeBuf.Write(bytes.Repeat(frameRWCPadding[:], paddingLength/4)) - - _, err := f.conn.Write(writeBuf.Bytes()) - return len(buf), err -} - -func (f *FrameRWC) Read(p []byte) (int, error) { - return f.BufferedRead(p, func() error { - buf := &bytes.Buffer{} - sum := crc32.NewIEEE() - writer := io.MultiWriter(buf, sum) - - for { - buf.Reset() - sum.Reset() - if _, err := io.CopyN(writer, 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) - } - - buf.Reset() - buf.Grow(int(messageLength) - 4 - 4) - if _, err := io.CopyN(writer, f.conn, int64(messageLength)-4-4); err != nil { - return errors.Annotate(err, "Cannot read the message frame") - } - - 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++ - - data, _ := ioutil.ReadAll(buf) - buf.Reset() - // write to buf, not to writer. This is because we are going to fetch - // crc32 checksum. - if _, err := io.CopyN(buf, f.conn, 4); err != nil { - return errors.Annotate(err, "Cannot read checksum") - } - checksum := binary.LittleEndian.Uint32(buf.Bytes()) - - if checksum != sum.Sum32() { - return errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum) - - } - f.Buffer.Write(data) - - return nil - }) -} - -func (f *FrameRWC) Close() error { - return f.conn.Close() -} - -func (f *FrameRWC) LocalAddr() *net.TCPAddr { - return f.conn.LocalAddr() -} - -func (f *FrameRWC) RemoteAddr() *net.TCPAddr { - return f.conn.RemoteAddr() -} - -func NewFrameRWC(conn wrappers.ReadWriteCloserWithAddr, seqNo int32) wrappers.ReadWriteCloserWithAddr { - return &FrameRWC{ - BufferedReader: wrappers.NewBufferedReader(), - conn: conn, - readSeqNo: seqNo, - writeSeqNo: seqNo, - } -} diff --git a/mtproto/wrappers/intermediate.go b/mtproto/wrappers/intermediate.go deleted file mode 100644 index 1c2a6ab..0000000 --- a/mtproto/wrappers/intermediate.go +++ /dev/null @@ -1,83 +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 { - 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 deleted file mode 100644 index 657f120..0000000 --- a/mtproto/wrappers/proxy_request.go +++ /dev/null @@ -1,112 +0,0 @@ -package wrappers - -import ( - "bytes" - "io" - "io/ioutil" - "net" - - "github.com/juju/errors" - - "github.com/9seconds/mtg/mtproto" - "github.com/9seconds/mtg/mtproto/rpc" - "github.com/9seconds/mtg/wrappers" -) - -type ProxyRequestReadWriteCloserWithAddr struct { - wrappers.BufferedReader - - conn wrappers.ReadWriteCloserWithAddr - req *rpc.RPCProxyRequest -} - -func (p *ProxyRequestReadWriteCloserWithAddr) Read(buf []byte) (int, error) { - return p.BufferedRead(buf, func() error { - ansBuf := &bytes.Buffer{} - ansBuf.Grow(4) - - if _, err := io.CopyN(ansBuf, p.conn, 4); err != nil { - return errors.Annotate(err, "Cannot read RPC tag") - } - - if bytes.Equal(ansBuf.Bytes(), rpc.RPCTagCloseExt) { - return p.readCloseExt() - } else if bytes.Equal(ansBuf.Bytes(), rpc.RPCTagProxyAns) { - return p.readProxyAns(buf) - } else if bytes.Equal(ansBuf.Bytes(), rpc.RPCTagSimpleAck) { - return p.readSimpleAck() - } - - return nil - }) -} - -func (p *ProxyRequestReadWriteCloserWithAddr) readCloseExt() error { - return errors.New("Connection has been closed remotely") -} - -func (p *ProxyRequestReadWriteCloserWithAddr) readProxyAns(buf []byte) error { - if _, err := io.CopyN(ioutil.Discard, p.conn, 8+4); err != nil { - return errors.Annotate(err, "Cannot skip flags and connid") - } - - for { - n, err := p.conn.Read(buf) - if err != nil { - return errors.Annotate(err, "Cannot read proxy answer") - } - if n == 0 { - break - } - p.Buffer.Write(buf[:n]) - } - - return nil -} - -func (p *ProxyRequestReadWriteCloserWithAddr) readSimpleAck() error { - if _, err := io.CopyN(ioutil.Discard, p.conn, 8); err != nil { - return errors.Annotate(err, "Cannot skip connid") - } - if _, err := io.CopyN(p.Buffer, p.conn, 4); err != nil { - return errors.Annotate(err, "Cannot read simple ack") - } - p.req.Options.SimpleAck = true - - return nil -} - -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 -} - -func (p *ProxyRequestReadWriteCloserWithAddr) Close() error { - return p.conn.Close() -} - -func (p *ProxyRequestReadWriteCloserWithAddr) LocalAddr() *net.TCPAddr { - return p.conn.LocalAddr() -} - -func (p *ProxyRequestReadWriteCloserWithAddr) RemoteAddr() *net.TCPAddr { - return p.conn.RemoteAddr() -} - -func NewProxyRequestRWC(conn wrappers.ReadWriteCloserWithAddr, connOpts *mtproto.ConnectionOpts, adTag []byte) (wrappers.ReadWriteCloserWithAddr, error) { - req, err := rpc.NewRPCProxyRequest(connOpts.ClientAddr, conn.LocalAddr(), connOpts, adTag) - if err != nil { - return nil, errors.Annotate(err, "Cannot create new RPC proxy request") - } - - return &ProxyRequestReadWriteCloserWithAddr{ - BufferedReader: wrappers.NewBufferedReader(), - conn: conn, - req: req, - }, nil -} diff --git a/proxy/copy_pool.go b/proxy/copy_pool.go deleted file mode 100644 index 23477aa..0000000 --- a/proxy/copy_pool.go +++ /dev/null @@ -1,18 +0,0 @@ -package proxy - -import ( - "sync" - - "github.com/9seconds/mtg/config" -) - -var copyPool sync.Pool - -func init() { - copyPool = sync.Pool{ - New: func() interface{} { - data := make([]byte, config.BufferSizeCopy) - return &data - }, - } -} diff --git a/proxy/proxy.go b/proxy/proxy.go new file mode 100644 index 0000000..a8399eb --- /dev/null +++ b/proxy/proxy.go @@ -0,0 +1,154 @@ +package proxy + +import ( + "io" + "net" + "sync" + + "github.com/juju/errors" + uuid "github.com/satori/go.uuid" + "go.uber.org/zap" + + "github.com/9seconds/mtg/client" + "github.com/9seconds/mtg/config" + "github.com/9seconds/mtg/mtproto" + "github.com/9seconds/mtg/telegram" + "github.com/9seconds/mtg/wrappers" +) + +type Proxy struct { + clientInit client.Init + tg telegram.Telegram + conf *config.Config +} + +func (p *Proxy) Serve() error { + lsock, err := net.Listen("tcp", p.conf.BindAddr()) + if err != nil { + return errors.Annotate(err, "Cannot create listen socket") + } + + for { + if conn, err := lsock.Accept(); err != nil { + zap.S().Errorw("Cannot allocate incoming connection", "error", err) + } else { + go p.accept(conn) + } + } +} + +func (p *Proxy) accept(conn net.Conn) { + connID := uuid.NewV4().String() + log := zap.S().With("connection_id", connID) + + defer func() { + conn.Close() + + if err := recover(); err != nil { + log.Errorw("Crash of accept handler", "error", err) + } + }() + + log.Infow("Client connected", "addr", conn.RemoteAddr()) + + client, opts, err := p.clientInit(conn, connID, p.conf) + if err != nil { + log.Errorw("Cannot initialize client connection", "error", err) + return + } + defer client.(io.Closer).Close() + + server, err := p.getTelegramConn(opts, connID) + if err != nil { + log.Errorw("Cannot initialize server connection", "error", err) + return + } + defer server.(io.Closer).Close() + + wait := &sync.WaitGroup{} + wait.Add(2) + + if p.conf.UseMiddleProxy() { + clientPacket := client.(wrappers.PacketReadWriteCloser) + serverPacket := server.(wrappers.PacketReadWriteCloser) + go p.middlePipe(clientPacket, serverPacket, wait, &opts.ReadHacks) + go p.middlePipe(serverPacket, clientPacket, wait, &opts.WriteHacks) + } else { + clientStream := client.(wrappers.StreamReadWriteCloser) + serverStream := server.(wrappers.StreamReadWriteCloser) + go p.directPipe(clientStream, serverStream, wait) + go p.directPipe(serverStream, clientStream, wait) + } + + wait.Wait() + + log.Infow("Client disconnected", "addr", conn.RemoteAddr()) +} + +func (p *Proxy) getTelegramConn(opts *mtproto.ConnectionOpts, connID string) (wrappers.Wrap, error) { + streamConn, err := p.tg.Dial(connID, opts) + if err != nil { + return nil, errors.Annotate(err, "Cannot dial to Telegram") + } + + packetConn, err := p.tg.Init(opts, streamConn) + if err != nil { + return nil, errors.Annotate(err, "Cannot handshake telegram") + } + + return packetConn, nil +} + +func (p *Proxy) middlePipe(src wrappers.PacketReadCloser, dst wrappers.PacketWriteCloser, wait *sync.WaitGroup, hacks *mtproto.Hacks) { + defer func() { + src.Close() + dst.Close() + wait.Done() + }() + + for { + hacks.SimpleAck = false + hacks.QuickAck = false + + packet, err := src.Read() + if err != nil { + src.Logger().Warnw("Cannot read packet", "error", err) + return + } + if _, err = dst.Write(packet); err != nil { + src.Logger().Warnw("Cannot write packet", "error", err) + return + } + } +} + +func (p *Proxy) directPipe(src wrappers.StreamReadCloser, dst wrappers.StreamWriteCloser, wait *sync.WaitGroup) { + defer func() { + src.Close() + dst.Close() + wait.Done() + }() + + if _, err := io.Copy(dst, src); err != nil { + src.Logger().Warnw("Cannot pump sockets", "error", err) + } +} + +func NewProxy(conf *config.Config) *Proxy { + var clientInit client.Init + var tg telegram.Telegram + + if conf.UseMiddleProxy() { + clientInit = client.MiddleInit + tg = telegram.NewMiddleTelegram(conf) + } else { + clientInit = client.DirectInit + tg = telegram.NewDirectTelegram(conf) + } + + return &Proxy{ + conf: conf, + clientInit: clientInit, + tg: tg, + } +} diff --git a/proxy/server.go b/proxy/server.go deleted file mode 100644 index 314b3d9..0000000 --- a/proxy/server.go +++ /dev/null @@ -1,157 +0,0 @@ -package proxy - -import ( - "context" - "io" - "net" - "sync" - - "github.com/juju/errors" - uuid "github.com/satori/go.uuid" - "go.uber.org/zap" - - "github.com/9seconds/mtg/client" - "github.com/9seconds/mtg/config" - "github.com/9seconds/mtg/mtproto" - "github.com/9seconds/mtg/telegram" - "github.com/9seconds/mtg/wrappers" -) - -// Server is an insgtance of MTPROTO proxy. -type Server struct { - conf *config.Config - logger *zap.SugaredLogger - stats *Stats - tg telegram.Telegram - clientInit client.Init -} - -// Serve does MTPROTO proxying. -func (s *Server) Serve() error { - lsock, err := net.Listen("tcp", s.conf.BindAddr()) - if err != nil { - return errors.Annotate(err, "Cannot create listen socket") - } - - for { - if conn, err := lsock.Accept(); err != nil { - s.logger.Warn("Cannot allocate incoming connection", "error", err) - } else { - go s.accept(conn) - } - } -} - -func (s *Server) accept(conn net.Conn) { - defer func() { - s.stats.closeConnection() - conn.Close() // nolint: errcheck - - if r := recover(); r != nil { - s.logger.Errorw("Crash of accept handler", "error", r) - } - }() - - s.stats.newConnection() - ctx, cancel := context.WithCancel(context.Background()) - socketID := uuid.NewV4().String() - - s.logger.Debugw("Client connected", - "addr", conn.RemoteAddr().String(), - "socketid", socketID, - ) - - connOpts, clientConn, err := s.getClientStream(ctx, cancel, conn, socketID) - if err != nil { - s.logger.Warnw("Cannot initialize client connection", - "addr", conn.RemoteAddr().String(), - "socketid", socketID, - "error", err, - ) - return - } - defer clientConn.Close() // nolint: errcheck - - tgConn, err := s.getTelegramStream(ctx, cancel, connOpts, socketID) - if err != nil { - s.logger.Warnw("Cannot initialize Telegram connection", - "socketid", socketID, - "error", err, - ) - return - } - defer tgConn.Close() // nolint: errcheck - - wait := &sync.WaitGroup{} - wait.Add(2) - - go s.pipe(clientConn, tgConn, wait) - go s.pipe(tgConn, clientConn, wait) - - <-ctx.Done() - wait.Wait() - - s.logger.Debugw("Client disconnected", - "addr", conn.RemoteAddr().String(), - "socketid", socketID, - ) -} - -func (s *Server) getClientStream(ctx context.Context, cancel context.CancelFunc, conn net.Conn, socketID string) (*mtproto.ConnectionOpts, io.ReadWriteCloser, error) { - connOpts, socket, err := s.clientInit(conn, s.conf) - if err != nil { - return nil, nil, errors.Annotate(err, "Cannot init client connection") - } - - socket = wrappers.NewTrafficRWC(socket, s.stats.addIncomingTraffic, s.stats.addOutgoingTraffic) - socket = wrappers.NewLogRWC(socket, s.logger, socketID, "client") - socket = wrappers.NewCtxRWC(ctx, cancel, socket) - - return connOpts, socket, nil -} - -func (s *Server) getTelegramStream(ctx context.Context, cancel context.CancelFunc, connOpts *mtproto.ConnectionOpts, socketID string) (io.ReadWriteCloser, error) { - conn, err := s.tg.Dial(connOpts) - if err != nil { - return nil, errors.Annotate(err, "Cannot connect to Telegram") - } - - conn = wrappers.NewTrafficRWC(conn, s.stats.addIncomingTraffic, s.stats.addOutgoingTraffic) - conn, err = s.tg.Init(connOpts, conn) - if err != nil { - return nil, errors.Annotate(err, "Cannot handshake Telegram") - } - - conn = wrappers.NewLogRWC(conn, s.logger, socketID, "telegram") - conn = wrappers.NewCtxRWC(ctx, cancel, conn) - - return conn, nil -} - -func (s *Server) pipe(dst io.Writer, src io.Reader, wait *sync.WaitGroup) { - defer wait.Done() - - buf := copyPool.Get().(*[]byte) - defer copyPool.Put(buf) - - io.CopyBuffer(dst, src, *buf) // nolint: errcheck -} - -// NewServer creates new instance of MTPROTO proxy. -func NewServer(conf *config.Config, logger *zap.SugaredLogger, stat *Stats) *Server { - clientInit := client.DirectInit - tg := telegram.NewDirectTelegram - - if len(conf.AdTag) > 0 { - clientInit = client.MiddleInit - tg = telegram.NewMiddleTelegram - } - - return &Server{ - conf: conf, - logger: logger, - stats: stat, - tg: tg(conf, logger), - clientInit: clientInit, - } -} diff --git a/proxy/stats.go b/proxy/stats.go deleted file mode 100644 index 9469c2f..0000000 --- a/proxy/stats.go +++ /dev/null @@ -1,74 +0,0 @@ -package proxy - -import ( - "encoding/json" - "net/http" - "strconv" - "sync/atomic" - "time" - - "github.com/9seconds/mtg/config" -) - -type statsUptime time.Time - -func (s statsUptime) MarshalJSON() ([]byte, error) { - uptime := int(time.Since(time.Time(s)).Seconds()) - return []byte(strconv.Itoa(uptime)), nil -} - -// Stats is a datastructure for statistics on work of this proxy. -type Stats struct { - AllConnections uint64 `json:"all_connections"` - ActiveConnections uint32 `json:"active_connections"` - Traffic struct { - Incoming uint64 `json:"incoming"` - Outgoing uint64 `json:"outgoing"` - } `json:"traffic"` - URLs config.IPURLs `json:"urls"` - Uptime statsUptime `json:"uptime"` - - conf *config.Config -} - -func (s *Stats) newConnection() { - atomic.AddUint64(&s.AllConnections, 1) - atomic.AddUint32(&s.ActiveConnections, 1) -} - -func (s *Stats) closeConnection() { - atomic.AddUint32(&s.ActiveConnections, ^uint32(0)) -} - -func (s *Stats) addIncomingTraffic(n int) { - atomic.AddUint64(&s.Traffic.Incoming, uint64(n)) -} - -func (s *Stats) addOutgoingTraffic(n int) { - atomic.AddUint64(&s.Traffic.Outgoing, uint64(n)) -} - -// Serve runs statistics HTTP server. -func (s *Stats) Serve() { - http.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - - encoder := json.NewEncoder(w) - encoder.SetEscapeHTML(false) - encoder.SetIndent("", " ") - encoder.Encode(s) // nolint: errcheck, gas - }) - - http.ListenAndServe(s.conf.StatAddr(), nil) // nolint: errcheck, gas -} - -// NewStats returns new instance of statistics datastructure. -func NewStats(conf *config.Config) *Stats { - stat := &Stats{ - Uptime: statsUptime(time.Now()), - conf: conf, - } - stat.URLs = conf.GetURLs() - - return stat -} diff --git a/telegram/dialer.go b/telegram/dialer.go index 9cf048b..fc70d64 100644 --- a/telegram/dialer.go +++ b/telegram/dialer.go @@ -30,11 +30,12 @@ func (t *tgDialer) dial(addr string) (net.Conn, error) { return conn, nil } -func (t *tgDialer) dialRWC(addr string) (wrappers.ReadWriteCloserWithAddr, error) { +func (t *tgDialer) dialRWC(addr, connID string) (wrappers.StreamReadWriteCloser, error) { conn, err := t.dial(addr) if err != nil { return nil, err } + tgConn := wrappers.NewConn(conn, connID, wrappers.ConnPurposeTelegram, t.conf.PublicIPv4, t.conf.PublicIPv6) - return wrappers.NewTimeoutRWC(conn, t.conf.PublicIPv4, t.conf.PublicIPv6), nil + return tgConn, nil } diff --git a/telegram/direct.go b/telegram/direct.go index ae89276..749aa07 100644 --- a/telegram/direct.go +++ b/telegram/direct.go @@ -4,7 +4,6 @@ import ( "net" "github.com/juju/errors" - "go.uber.org/zap" "github.com/9seconds/mtg/config" "github.com/9seconds/mtg/mtproto" @@ -29,11 +28,11 @@ var ( } ) -type directTelegram struct { +type DirectTelegram struct { baseTelegram } -func (t *directTelegram) Dial(connOpts *mtproto.ConnectionOpts) (wrappers.ReadWriteCloserWithAddr, error) { +func (t *DirectTelegram) Dial(connID string, connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) { dc := connOpts.DC if dc < 0 { dc = -dc @@ -41,23 +40,23 @@ func (t *directTelegram) Dial(connOpts *mtproto.ConnectionOpts) (wrappers.ReadWr dc = 1 } - return t.baseTelegram.dial(dc-1, connOpts.ConnectionProto) + return t.baseTelegram.dial(dc-1, connID, connOpts.ConnectionProto) } -func (t *directTelegram) Init(connOpts *mtproto.ConnectionOpts, conn wrappers.ReadWriteCloserWithAddr) (wrappers.ReadWriteCloserWithAddr, error) { +func (t *DirectTelegram) Init(connOpts *mtproto.ConnectionOpts, conn wrappers.StreamReadWriteCloser) (wrappers.Wrap, error) { obfs2, frame := obfuscated2.MakeTelegramObfuscated2Frame(connOpts) - if n, err := conn.Write(frame); err != nil || n != obfuscated2.FrameLen { + if _, err := conn.Write(frame); err != nil { return nil, errors.Annotate(err, "Cannot write hadnshake frame") } - return wrappers.NewStreamCipherRWC(conn, obfs2.Encryptor, obfs2.Decryptor), nil + return wrappers.NewStreamCipher(conn, obfs2.Encryptor, obfs2.Decryptor), nil } // NewDirectTelegram returns Telegram instance which connects directly // to Telegram bypassing middleproxies. -func NewDirectTelegram(conf *config.Config, _ *zap.SugaredLogger) Telegram { - return &directTelegram{baseTelegram{ +func NewDirectTelegram(conf *config.Config) Telegram { + return &DirectTelegram{baseTelegram{ dialer: tgDialer{ Dialer: net.Dialer{Timeout: telegramDialTimeout}, conf: conf, diff --git a/telegram/middle.go b/telegram/middle.go index 8e1bd7b..7404e3d 100644 --- a/telegram/middle.go +++ b/telegram/middle.go @@ -1,30 +1,114 @@ package telegram import ( - "fmt" - "io" "net" "net/http" "sync" "github.com/juju/errors" - "go.uber.org/zap" "github.com/9seconds/mtg/config" "github.com/9seconds/mtg/mtproto" "github.com/9seconds/mtg/mtproto/rpc" - mtwrappers "github.com/9seconds/mtg/mtproto/wrappers" "github.com/9seconds/mtg/wrappers" ) -type middleTelegram struct { +type MiddleTelegram struct { middleTelegramCaller conf *config.Config } -func NewMiddleTelegram(conf *config.Config, logger *zap.SugaredLogger) Telegram { - tg := &middleTelegram{ +func (t *MiddleTelegram) Init(connOpts *mtproto.ConnectionOpts, conn wrappers.StreamReadWriteCloser) (wrappers.Wrap, error) { + rpcNonceConn := wrappers.NewMTProtoFrame(conn, rpc.SeqNoNonce) + + rpcNonceReq, err := t.sendRPCNonceRequest(rpcNonceConn) + if err != nil { + return nil, err + } + rpcNonceResp, err := t.receiveRPCNonceResponse(rpcNonceConn, rpcNonceReq) + if err != nil { + return nil, err + } + + secureConn := wrappers.NewMiddleProxyCipher(conn, rpcNonceReq, rpcNonceResp, t.proxySecret) + frameConn := wrappers.NewMTProtoFrame(secureConn, rpc.SeqNoHandshake) + + rpcHandshakeReq, err := t.sendRPCHandshakeRequest(frameConn) + if err != nil { + return nil, err + } + _, err = t.receiveRPCHandshakeResponse(frameConn, rpcHandshakeReq) + if err != nil { + return nil, err + } + + proxyConn, err := wrappers.NewMTProtoProxy(frameConn, connOpts, t.conf.AdTag) + if err != nil { + return nil, err + } + proxyConn.Logger().Infow("Telegram connection initialized") + + return proxyConn, nil +} + +func (t *MiddleTelegram) sendRPCNonceRequest(conn wrappers.PacketWriter) (*rpc.NonceRequest, error) { + rpcNonceReq, err := rpc.NewNonceRequest(t.proxySecret) + if err != nil { + return nil, errors.Annotate(err, "Cannot create RPC nonce request") + } + if _, err = conn.Write(rpcNonceReq.Bytes()); err != nil { + return nil, errors.Annotate(err, "Cannot send RPC nonce request") + } + + return rpcNonceReq, nil +} + +func (t *MiddleTelegram) receiveRPCNonceResponse(conn wrappers.PacketReader, req *rpc.NonceRequest) (*rpc.NonceResponse, error) { + packet, err := conn.Read() + if err != nil { + return nil, errors.Annotate(err, "Cannot read RPC nonce response") + } + + rpcNonceResp, err := rpc.NewNonceResponse(packet) + if err != nil { + return nil, errors.Annotate(err, "Cannot initialize RPC nonce response") + } + if err = rpcNonceResp.Valid(req); err != nil { + return nil, errors.Annotate(err, "Invalid RPC nonce response") + } + + return rpcNonceResp, nil +} + +func (t *MiddleTelegram) sendRPCHandshakeRequest(conn wrappers.PacketWriter) (*rpc.HandshakeRequest, error) { + req := rpc.NewHandshakeRequest() + if _, err := conn.Write(req.Bytes()); err != nil { + return nil, errors.Annotate(err, "Cannot send RPC handshake request") + } + + return req, nil +} + +func (t *MiddleTelegram) receiveRPCHandshakeResponse(conn wrappers.PacketReader, req *rpc.HandshakeRequest) (*rpc.HandshakeResponse, error) { + packet, err := conn.Read() + if err != nil { + return nil, errors.Annotate(err, "Cannot read RPC handshake response") + } + + rpcHandshakeResp, err := rpc.NewHandshakeResponse(packet) + if err != nil { + return nil, errors.Annotate(err, "Cannot initialize RPC handshake response") + } + if err = rpcHandshakeResp.Valid(req); err != nil { + return nil, errors.Annotate(err, "Invalid RPC handshake response") + } + + return rpcHandshakeResp, nil +} + +func NewMiddleTelegram(conf *config.Config) Telegram { + tg := &MiddleTelegram{ middleTelegramCaller: middleTelegramCaller{ baseTelegram: baseTelegram{ dialer: tgDialer{ @@ -32,7 +116,6 @@ func NewMiddleTelegram(conf *config.Config, logger *zap.SugaredLogger) Telegram conf: conf, }, }, - logger: logger, httpClient: &http.Client{ Timeout: middleTelegramHTTPClientTimeout, }, @@ -48,88 +131,3 @@ func NewMiddleTelegram(conf *config.Config, logger *zap.SugaredLogger) Telegram return tg } - -func (t *middleTelegram) Init(connOpts *mtproto.ConnectionOpts, conn wrappers.ReadWriteCloserWithAddr) (wrappers.ReadWriteCloserWithAddr, error) { - rpcNonceConn := mtwrappers.NewFrameRWC(conn, rpc.RPCNonceSeqNo) - - rpcNonceReq, err := t.sendRPCNonceRequest(rpcNonceConn) - if err != nil { - return nil, err - } - rpcNonceResp, err := t.receiveRPCNonceResponse(rpcNonceConn, rpcNonceReq) - if err != nil { - return nil, err - } - - secureConn := mtwrappers.NewMiddleProxyCipherRWC(conn, rpcNonceReq, rpcNonceResp, t.proxySecret) - secureConn = mtwrappers.NewFrameRWC(secureConn, rpc.RPCHandshakeSeqNo) - - rpcHandshakeReq, err := t.sendRPCHandshakeRequest(secureConn) - if err != nil { - return nil, err - } - _, err = t.receiveRPCHandshakeResponse(secureConn, rpcHandshakeReq) - if err != nil { - return nil, err - } - - return mtwrappers.NewProxyRequestRWC(secureConn, connOpts, t.conf.AdTag) -} - -func (t *middleTelegram) sendRPCNonceRequest(conn io.Writer) (*rpc.RPCNonceRequest, error) { - rpcNonceReq, err := rpc.NewRPCNonceRequest(t.proxySecret) - if err != nil { - return nil, errors.Annotate(err, "Cannot create RPC nonce request") - } - if _, err = conn.Write(rpcNonceReq.Bytes()); err != nil { - return nil, errors.Annotate(err, "Cannot send RPC nonce request") - } - - return rpcNonceReq, nil -} - -func (t *middleTelegram) receiveRPCNonceResponse(conn io.Reader, req *rpc.RPCNonceRequest) (*rpc.RPCNonceResponse, error) { - var ans [128]byte - - n, err := conn.Read(ans[:]) - if err != nil { - return nil, errors.Annotate(err, "Cannot read RPC nonce response") - } - rpcNonceResp, err := rpc.NewRPCNonceResponse(ans[:n]) - if err != nil { - return nil, errors.Annotate(err, "Cannot initialize RPC nonce response") - } - if err = rpcNonceResp.Valid(req); err != nil { - return nil, errors.Annotate(err, "Invalid RPC nonce response") - } - - return rpcNonceResp, nil -} - -func (t *middleTelegram) sendRPCHandshakeRequest(conn io.Writer) (*rpc.RPCHandshakeRequest, error) { - req := rpc.NewRPCHandshakeRequest() - if _, err := conn.Write(req.Bytes()); err != nil { - return nil, errors.Annotate(err, "Cannot send RPC handshake request") - } - - return req, nil -} - -func (t *middleTelegram) receiveRPCHandshakeResponse(conn io.Reader, req *rpc.RPCHandshakeRequest) (*rpc.RPCHandshakeResponse, error) { - var ans [128]byte - - n, err := conn.Read(ans[:]) - if err != nil { - return nil, errors.Annotate(err, "Cannot read RPC handshake response") - } - rpcHandshakeResp, err := rpc.NewRPCHandshakeResponse(ans[:n]) - if err != nil { - return nil, errors.Annotate(err, "Cannot initialize RPC handshake response") - } - if err = rpcHandshakeResp.Valid(req); err != nil { - return nil, errors.Annotate(err, "Invalid RPC handshake response") - } - fmt.Println("VICTORY") - - return rpcHandshakeResp, nil -} diff --git a/telegram/middle_caller.go b/telegram/middle_caller.go index acc1b27..832430a 100644 --- a/telegram/middle_caller.go +++ b/telegram/middle_caller.go @@ -28,18 +28,17 @@ const ( tgUserAgent = "mtg" ) -var middleTelegramProxyConfigSplitter *regexp.Regexp +var middleTelegramProxyConfigSplitter = regexp.MustCompile(`\s+`) type middleTelegramCaller struct { baseTelegram proxySecret []byte dialerMutex *sync.RWMutex - logger *zap.SugaredLogger httpClient *http.Client } -func (t *middleTelegramCaller) Dial(connOpts *mtproto.ConnectionOpts) (wrappers.ReadWriteCloserWithAddr, error) { +func (t *middleTelegramCaller) Dial(connID string, connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) { dc := connOpts.DC if dc == 0 { dc = 1 @@ -47,13 +46,13 @@ func (t *middleTelegramCaller) Dial(connOpts *mtproto.ConnectionOpts) (wrappers. t.dialerMutex.RLock() defer t.dialerMutex.RUnlock() - return t.baseTelegram.dial(dc, connOpts.ConnectionProto) + return t.baseTelegram.dial(dc, connID, connOpts.ConnectionProto) } func (t *middleTelegramCaller) autoUpdate() { for range time.Tick(middleTelegramAutoUpdateInterval) { if err := t.update(); err != nil { - t.logger.Warnw("Cannot update from Telegram", "error", err) + zap.S().Warnw("Cannot update from Telegram", "error", err) } } } @@ -80,7 +79,7 @@ func (t *middleTelegramCaller) update() error { t.v6Addresses = v6Addresses t.dialerMutex.Unlock() - t.logger.Infow("Telegram middle proxy data has been updated") + zap.S().Infow("Telegram middle proxy data has been updated") return nil } @@ -151,7 +150,3 @@ func (t *middleTelegramCaller) call(url string) (*http.Response, error) { return t.httpClient.Do(req) } - -func init() { - middleTelegramProxyConfigSplitter = regexp.MustCompile(`\s+`) -} diff --git a/telegram/telegram.go b/telegram/telegram.go index 5ce16ad..43f436d 100644 --- a/telegram/telegram.go +++ b/telegram/telegram.go @@ -9,12 +9,9 @@ import ( "github.com/9seconds/mtg/wrappers" ) -// Telegram defines an interface to connect to Telegram. This -// encapsulates logic of working with middleproxies or direct -// connections. type Telegram interface { - Dial(*mtproto.ConnectionOpts) (wrappers.ReadWriteCloserWithAddr, error) - Init(*mtproto.ConnectionOpts, wrappers.ReadWriteCloserWithAddr) (wrappers.ReadWriteCloserWithAddr, error) + Dial(string, *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) + Init(*mtproto.ConnectionOpts, wrappers.StreamReadWriteCloser) (wrappers.Wrap, error) } type baseTelegram struct { @@ -24,7 +21,7 @@ type baseTelegram struct { v6Addresses map[int16][]string } -func (b *baseTelegram) dial(dcIdx int16, proto mtproto.ConnectionProtocol) (wrappers.ReadWriteCloserWithAddr, error) { +func (b *baseTelegram) dial(dcIdx int16, connID string, proto mtproto.ConnectionProtocol) (wrappers.StreamReadWriteCloser, error) { addrs := make([]string, 2) if proto&mtproto.ConnectionProtocolIPv6 != 0 { @@ -39,7 +36,7 @@ func (b *baseTelegram) dial(dcIdx int16, proto mtproto.ConnectionProtocol) (wrap } for _, addr := range addrs { - if conn, err := b.dialer.dialRWC(addr); err == nil { + if conn, err := b.dialer.dialRWC(addr, connID); err == nil { return conn, err } } diff --git a/utils/read_current_data.go b/utils/read_current_data.go new file mode 100644 index 0000000..d66802b --- /dev/null +++ b/utils/read_current_data.go @@ -0,0 +1,20 @@ +package utils + +import "io" + +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) + n := readCurrentDataBufferSize + + for n == len(buf) { + n, err = src.Read(buf) + if err != nil { + return nil, err + } + rv = append(rv, buf[:n]...) + } + + return rv, nil +} 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..a4a7149 --- /dev/null +++ b/wrappers/blockcipher.go @@ -0,0 +1,90 @@ +package wrappers + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "net" + + "go.uber.org/zap" + + "github.com/9seconds/mtg/utils" + "github.com/juju/errors" +) + +type BlockCipher struct { + buf *bytes.Buffer + + logger *zap.SugaredLogger + conn StreamReadWriteCloser + encryptor cipher.BlockMode + decryptor cipher.BlockMode +} + +func (b *BlockCipher) Read(p []byte) (int, error) { + if b.buf.Len() > 0 { + return b.flush(p) + } + + buf := []byte{} + for len(buf) == 0 || len(buf)%aes.BlockSize != 0 { + rv, err := utils.ReadCurrentData(b.conn) + if err != nil { + return 0, errors.Annotate(err, "Cannot read from socket") + } + buf = append(buf, rv...) + } + + b.decryptor.CryptBlocks(buf, buf) + b.buf.Write(buf) + + return b.flush(p) +} + +func (b *BlockCipher) flush(p []byte) (int, error) { + if b.buf.Len() <= len(p) { + sizeToReturn := b.buf.Len() + copy(p, b.buf.Bytes()) + b.buf.Reset() + return sizeToReturn, nil + } + + return b.buf.Read(p) +} + +func (b *BlockCipher) 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)) + b.encryptor.CryptBlocks(encrypted, p) + + return b.conn.Write(encrypted) +} + +func (b *BlockCipher) Logger() *zap.SugaredLogger { + return b.logger +} + +func (b *BlockCipher) LocalAddr() *net.TCPAddr { + return b.conn.LocalAddr() +} + +func (b *BlockCipher) RemoteAddr() *net.TCPAddr { + return b.conn.RemoteAddr() +} + +func (b *BlockCipher) Close() error { + return b.conn.Close() +} + +func NewBlockCipher(conn StreamReadWriteCloser, encryptor, decryptor cipher.BlockMode) StreamReadWriteCloser { + return &BlockCipher{ + buf: &bytes.Buffer{}, + conn: conn, + logger: conn.Logger().Named("block-cipher"), + encryptor: encryptor, + decryptor: decryptor, + } +} diff --git a/wrappers/blockcipherrwc.go b/wrappers/blockcipherrwc.go deleted file mode 100644 index 9db024d..0000000 --- a/wrappers/blockcipherrwc.go +++ /dev/null @@ -1,66 +0,0 @@ -package wrappers - -import ( - "crypto/aes" - "crypto/cipher" - "net" - - "github.com/juju/errors" -) - -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 { - 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]) - bufferLength = c.Buffer.Len() - } - c.decryptor.CryptBlocks(c.Buffer.Bytes(), c.Buffer.Bytes()) - - 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 { - 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 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 deleted file mode 100644 index cb48d56..0000000 --- a/wrappers/buffered_reader.go +++ /dev/null @@ -1,32 +0,0 @@ -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) { - if b.Buffer.Len() < len(p) { - sizeToReturn := b.Buffer.Len() - copy(p, b.Buffer.Bytes()) - b.Buffer.Reset() - return sizeToReturn, nil - } - - return b.Buffer.Read(p) -} - -func NewBufferedReader() BufferedReader { - return BufferedReader{Buffer: &bytes.Buffer{}} -} diff --git a/wrappers/conn.go b/wrappers/conn.go new file mode 100644 index 0000000..ad853af --- /dev/null +++ b/wrappers/conn.go @@ -0,0 +1,105 @@ +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 Conn struct { + connID string + conn net.Conn + logger *zap.SugaredLogger + publicIPv4 net.IP + publicIPv6 net.IP +} + +func (c *Conn) Write(p []byte) (int, error) { + c.conn.SetWriteDeadline(time.Now().Add(connTimeoutWrite)) + n, err := c.conn.Write(p) + + c.logger.Debugw("Write to stream", "bytes", n, "error", err) + + return n, err +} + +func (c *Conn) Read(p []byte) (int, error) { + c.conn.SetReadDeadline(time.Now().Add(connTimeoutRead)) + n, err := c.conn.Read(p) + + c.logger.Debugw("Read from stream", "bytes", n, "error", err) + + return n, err +} + +func (c *Conn) Close() error { + defer c.logger.Debugw("Closed connection") + return c.conn.Close() +} + +func (c *Conn) LocalAddr() *net.TCPAddr { + addr := c.conn.LocalAddr().(*net.TCPAddr) + newAddr := *addr + + if c.RemoteAddr().IP.To4() != nil { + if c.publicIPv4 != nil { + newAddr.IP = c.publicIPv4 + } + } else if c.publicIPv6 != nil { + newAddr.IP = c.publicIPv6 + } + + return &newAddr +} + +func (c *Conn) RemoteAddr() *net.TCPAddr { + return c.conn.RemoteAddr().(*net.TCPAddr) +} + +func (c *Conn) Logger() *zap.SugaredLogger { + return c.logger +} + +func NewConn(conn net.Conn, connID string, purpose ConnPurpose, publicIPv4, publicIPv6 net.IP) StreamReadWriteCloser { + logger := zap.S().With( + "connection_id", connID, + "local_address", conn.LocalAddr(), + "remote_address", conn.RemoteAddr(), + "purpose", purpose, + ).Named("conn") + + wrapper := Conn{ + logger: logger, + connID: connID, + conn: conn, + publicIPv4: publicIPv4, + publicIPv6: publicIPv6, + } + wrapper.logger = logger.With("faked_local_addr", wrapper.LocalAddr()) + + return &wrapper +} diff --git a/wrappers/ctxrwc.go b/wrappers/ctxrwc.go deleted file mode 100644 index 64c1006..0000000 --- a/wrappers/ctxrwc.go +++ /dev/null @@ -1,67 +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() -} - -// 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 fd80ad4..0000000 --- a/wrappers/logrwc.go +++ /dev/null @@ -1,55 +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) - 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) - 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() -} - -// 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..2ed116f --- /dev/null +++ b/wrappers/mtproto_abridged.go @@ -0,0 +1,152 @@ +package wrappers + +import ( + "bytes" + "io" + "net" + + "github.com/juju/errors" + "go.uber.org/zap" + + "github.com/9seconds/mtg/mtproto" + "github.com/9seconds/mtg/utils" +) + +const ( + mtprotoAbridgedSmallPacketLength = 0x7f + mtprotoAbridgedQuickAckLength = 0x80 + mtprotoAbridgedLargePacketLength = 16777216 // 256 ^ 3 +) + +type MTProtoAbridged struct { + conn StreamReadWriteCloser + opts *mtproto.ConnectionOpts + logger *zap.SugaredLogger + + readCounter uint32 + writeCounter uint32 +} + +func (m *MTProtoAbridged) Read() ([]byte, error) { + defer func() { + m.readCounter++ + }() + + m.logger.Debugw("Read packet", + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.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 := uint32(buf.Bytes()[0]) + buf.Reset() + + m.logger.Debugw("Packet first byte", + "byte", msgLength, + "counter", m.readCounter, + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, + ) + + if msgLength >= mtprotoAbridgedQuickAckLength { + m.opts.ReadHacks.QuickAck = true + msgLength -= mtprotoAbridgedQuickAckLength + } + + if msgLength == mtprotoAbridgedSmallPacketLength { + 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()) + msgLength = utils.FromUint24(number) + } + msgLength *= 4 + + m.logger.Debugw("Packet length", + "length", msgLength, + "simple_ack", m.opts.ReadHacks.SimpleAck, + "quick_ack", m.opts.ReadHacks.QuickAck, + "counter", m.readCounter, + ) + + buf.Reset() + buf.Grow(int(msgLength)) + if _, err := io.CopyN(buf, m.conn, int64(msgLength)); err != nil { + return nil, errors.Annotate(err, "Cannot read message") + } + + return buf.Bytes(), nil +} + +func (m *MTProtoAbridged) Write(p []byte) (int, error) { + defer func() { + m.writeCounter++ + }() + + m.logger.Debugw("Write 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 < mtprotoAbridgedSmallPacketLength: + newData := append([]byte{byte(packetLength)}, p...) + return m.conn.Write(newData) + + case packetLength < mtprotoAbridgedLargePacketLength: + length24 := utils.ToUint24(uint32(packetLength)) + + buf := &bytes.Buffer{} + buf.Grow(1 + 3 + len(p)) + + buf.WriteByte(byte(mtprotoAbridgedSmallPacketLength)) + buf.Write(length24[:]) + buf.Write(p) + + return m.conn.Write(buf.Bytes()) + } + + return 0, errors.Errorf("Packet is too big %d", len(p)) +} + +func (m *MTProtoAbridged) Logger() *zap.SugaredLogger { + return m.logger +} + +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 StreamReadWriteCloser, opts *mtproto.ConnectionOpts) PacketReadWriteCloser { + return &MTProtoAbridged{ + conn: conn, + opts: opts, + logger: conn.Logger().Named("mtproto-abridged"), + } +} diff --git a/mtproto/wrappers/crypt.go b/wrappers/mtproto_cipher.go similarity index 64% rename from mtproto/wrappers/crypt.go rename to wrappers/mtproto_cipher.go index e28ad8e..72dad3f 100644 --- a/mtproto/wrappers/crypt.go +++ b/wrappers/mtproto_cipher.go @@ -10,7 +10,7 @@ import ( "net" "github.com/9seconds/mtg/mtproto/rpc" - "github.com/9seconds/mtg/wrappers" + "github.com/9seconds/mtg/utils" ) type CipherPurpose uint8 @@ -22,21 +22,20 @@ const ( var emptyIP = [4]byte{0x00, 0x00, 0x00, 0x00} -func NewMiddleProxyCipherRWC(conn wrappers.ReadWriteCloserWithAddr, req *rpc.RPCNonceRequest, resp *rpc.RPCNonceResponse, secret []byte) wrappers.ReadWriteCloserWithAddr { +func NewMiddleProxyCipher(conn StreamReadWriteCloser, req *rpc.NonceRequest, resp *rpc.NonceResponse, secret []byte) StreamReadWriteCloser { localAddr := conn.LocalAddr() remoteAddr := conn.RemoteAddr() - encKey, encIV := makeKeys(CipherPurposeClient, req, resp, localAddr, remoteAddr, secret) - decKey, decIV := makeKeys(CipherPurposeServer, req, resp, localAddr, remoteAddr, secret) + encKey, encIV := deriveKeys(CipherPurposeClient, req, resp, localAddr, remoteAddr, secret) + decKey, decIV := deriveKeys(CipherPurposeServer, req, resp, localAddr, remoteAddr, secret) enc, _ := makeEncrypterDecrypter(encKey, encIV) _, dec := makeEncrypterDecrypter(decKey, decIV) - return wrappers.NewBlockCipherRWC(conn, enc, dec) + return NewBlockCipher(conn, enc, dec) } -func makeKeys(purpose CipherPurpose, req *rpc.RPCNonceRequest, resp *rpc.RPCNonceResponse, - client *net.TCPAddr, remote *net.TCPAddr, secret []byte) ([]byte, []byte) { +func deriveKeys(purpose CipherPurpose, req *rpc.NonceRequest, resp *rpc.NonceResponse, client *net.TCPAddr, remote *net.TCPAddr, secret []byte) ([]byte, []byte) { message := bytes.Buffer{} message.Write(resp.Nonce[:]) message.Write(req.Nonce[:]) @@ -45,8 +44,8 @@ func makeKeys(purpose CipherPurpose, req *rpc.RPCNonceRequest, resp *rpc.RPCNonc clientIPv4 := emptyIP[:] serverIPv4 := emptyIP[:] if client.IP.To4() != nil { - clientIPv4 = reverseBytes(client.IP.To4()) - serverIPv4 = reverseBytes(remote.IP.To4()) + clientIPv4 = utils.ReverseBytes(client.IP.To4()) + serverIPv4 = utils.ReverseBytes(remote.IP.To4()) } message.Write(serverIPv4) @@ -93,16 +92,3 @@ func makeEncrypterDecrypter(key, iv []byte) (cipher.BlockMode, cipher.BlockMode) return cipher.NewCBCEncrypter(block, iv), cipher.NewCBCDecrypter(block, iv) } - -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/wrappers/mtproto_frame.go b/wrappers/mtproto_frame.go new file mode 100644 index 0000000..71ee084 --- /dev/null +++ b/wrappers/mtproto_frame.go @@ -0,0 +1,143 @@ +package wrappers + +import ( + "bytes" + "crypto/aes" + "encoding/binary" + "hash/crc32" + "io" + "io/ioutil" + "net" + + "github.com/juju/errors" + "go.uber.org/zap" +) + +const ( + mtprotoFrameMinMessageLength = 12 + mtprotoFrameMaxMessageLength = 16777216 +) + +var mtprotoFramePadding = []byte{0x04, 0x00, 0x00, 0x00} + +type MTProtoFrame struct { + conn StreamReadWriteCloser + logger *zap.SugaredLogger + + readSeqNo int32 + writeSeqNo int32 +} + +func (m *MTProtoFrame) Read() ([]byte, error) { + buf := &bytes.Buffer{} + sum := crc32.NewIEEE() + writer := io.MultiWriter(buf, sum) + + for { + buf.Reset() + sum.Reset() + if _, err := io.CopyN(writer, m.conn, 4); err != nil { + return nil, errors.Annotate(err, "Cannot read frame padding") + } + if !bytes.Equal(buf.Bytes(), mtprotoFramePadding) { + break + } + } + + messageLength := binary.LittleEndian.Uint32(buf.Bytes()) + m.logger.Debugw("Read MTProto frame", + "messageLength", messageLength, + "sequence_number", m.readSeqNo, + ) + if messageLength%4 != 0 || messageLength < mtprotoFrameMinMessageLength || messageLength > mtprotoFrameMaxMessageLength { + return nil, errors.Errorf("Incorrect frame message length %d", messageLength) + } + + buf.Reset() + buf.Grow(int(messageLength) - 4 - 4) + if _, err := io.CopyN(writer, m.conn, int64(messageLength)-4-4); err != nil { + return nil, errors.Annotate(err, "Cannot read the message frame") + } + + var seqNo int32 + binary.Read(buf, binary.LittleEndian, &seqNo) + if seqNo != m.readSeqNo { + return nil, errors.Errorf("Unexpected sequence number %d (wait for %d)", seqNo, m.readSeqNo) + } + + data, _ := ioutil.ReadAll(buf) + buf.Reset() + // write to buf, not to writer. This is because we are going to fetch + // crc32 checksum. + if _, err := io.CopyN(buf, m.conn, 4); err != nil { + return nil, errors.Annotate(err, "Cannot read checksum") + } + + checksum := binary.LittleEndian.Uint32(buf.Bytes()) + if checksum != sum.Sum32() { + return nil, errors.Errorf("CRC32 checksum mismatch. Wait for %d, got %d", sum.Sum32(), checksum) + } + + m.logger.Debugw("Read MTProto frame", + "messageLength", messageLength, + "sequence_number", m.readSeqNo, + "dataLength", len(data), + "checksum", checksum, + ) + m.readSeqNo++ + + return data, nil +} + +func (m *MTProtoFrame) Write(p []byte) (int, error) { + messageLength := 4 + 4 + len(p) + 4 + paddingLength := (aes.BlockSize - messageLength%aes.BlockSize) % aes.BlockSize + + buf := &bytes.Buffer{} + buf.Grow(messageLength + paddingLength) + + binary.Write(buf, binary.LittleEndian, uint32(messageLength)) + binary.Write(buf, binary.LittleEndian, m.writeSeqNo) + buf.Write(p) + + checksum := crc32.ChecksumIEEE(buf.Bytes()) + binary.Write(buf, binary.LittleEndian, checksum) + buf.Write(bytes.Repeat(mtprotoFramePadding, paddingLength/4)) + + m.logger.Debugw("Write MTProto frame", + "length", len(p), + "sequence_number", m.writeSeqNo, + "crc32", checksum, + "frame_length", buf.Len(), + ) + m.writeSeqNo++ + + _, err := m.conn.Write(buf.Bytes()) + + return len(p), err +} + +func (m *MTProtoFrame) Logger() *zap.SugaredLogger { + return m.logger +} + +func (m *MTProtoFrame) LocalAddr() *net.TCPAddr { + return m.conn.LocalAddr() +} + +func (m *MTProtoFrame) RemoteAddr() *net.TCPAddr { + return m.conn.RemoteAddr() +} + +func (m *MTProtoFrame) Close() error { + return m.conn.Close() +} + +func NewMTProtoFrame(conn StreamReadWriteCloser, seqNo int32) PacketReadWriteCloser { + return &MTProtoFrame{ + conn: conn, + logger: conn.Logger().Named("mtproto-frame"), + readSeqNo: seqNo, + writeSeqNo: seqNo, + } +} diff --git a/wrappers/mtproto_intermediate.go b/wrappers/mtproto_intermediate.go new file mode 100644 index 0000000..5c29c79 --- /dev/null +++ b/wrappers/mtproto_intermediate.go @@ -0,0 +1,113 @@ +package wrappers + +import ( + "bytes" + "encoding/binary" + "io" + "net" + + "github.com/juju/errors" + "go.uber.org/zap" + + "github.com/9seconds/mtg/mtproto" +) + +const mtprotoIntermediateQuickAckLength = 0x80000000 + +type MTProtoIntermediate struct { + conn StreamReadWriteCloser + opts *mtproto.ConnectionOpts + logger *zap.SugaredLogger + + readCounter uint32 + writeCounter uint32 +} + +func (m *MTProtoIntermediate) Read() ([]byte, error) { + defer func() { + m.readCounter++ + }() + + m.logger.Debugw("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.logger.Debugw("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 + } + + return buf.Bytes()[:length], nil +} + +func (m *MTProtoIntermediate) Write(p []byte) (int, error) { + defer func() { + m.writeCounter++ + }() + + m.logger.Debugw("Write packet", + "simple_ack", m.opts.WriteHacks.SimpleAck, + "quick_ack", m.opts.WriteHacks.QuickAck, + "counter", 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) Logger() *zap.SugaredLogger { + return m.logger +} + +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 StreamReadWriteCloser, opts *mtproto.ConnectionOpts) PacketReadWriteCloser { + return &MTProtoIntermediate{ + conn: conn, + logger: conn.Logger().Named("mtproto-intermediate"), + opts: opts, + } +} diff --git a/wrappers/mtproto_proxy.go b/wrappers/mtproto_proxy.go new file mode 100644 index 0000000..fbbb332 --- /dev/null +++ b/wrappers/mtproto_proxy.go @@ -0,0 +1,158 @@ +package wrappers + +import ( + "bytes" + "fmt" + "net" + + "github.com/juju/errors" + "go.uber.org/zap" + + "github.com/9seconds/mtg/mtproto" + "github.com/9seconds/mtg/mtproto/rpc" +) + +type MTProtoProxy struct { + conn PacketReadWriteCloser + req *rpc.ProxyRequest + logger *zap.SugaredLogger + + readCounter uint32 + writeCounter uint32 +} + +func (m *MTProtoProxy) Read() ([]byte, error) { + defer func() { + m.readCounter++ + }() + + m.logger.Debugw("Read packet", + "counter", m.readCounter, + "simple_ack", m.req.Options.WriteHacks.SimpleAck, + "quick_ack", m.req.Options.WriteHacks.QuickAck, + ) + + packet, err := m.conn.Read() + if err != nil { + return nil, errors.Annotate(err, "Cannot read packet") + } + + m.logger.Debugw("Read packet length", + "counter", m.readCounter, + "simple_ack", m.req.Options.WriteHacks.SimpleAck, + "quick_ack", m.req.Options.WriteHacks.QuickAck, + "length", len(packet), + ) + + if len(packet) < 4 { + return nil, errors.Annotate(err, "Incorrect packet length") + } + + tag, packet := packet[:4], packet[4:] + switch { + case bytes.Equal(tag, rpc.TagProxyAns): + return m.readProxyAns(packet) + case bytes.Equal(tag, rpc.TagSimpleAck): + return m.readSimpleAck(packet) + case bytes.Equal(tag, rpc.TagCloseExt): + return m.readCloseExt(packet) + } + + return nil, errors.Errorf("Unknown RPC answer %v", tag) +} + +func (m *MTProtoProxy) readProxyAns(data []byte) ([]byte, error) { + if len(data) < 12 { + return nil, errors.Errorf("Incorrect data of proxy answer: %d", len(data)) + } + data = data[12:] + + m.logger.Debugw("Read RPC_PROXY_ANS", + "counter", m.readCounter, + "length", len(data), + ) + + return data, nil +} + +func (m *MTProtoProxy) readSimpleAck(data []byte) ([]byte, error) { + if len(data) != 12 { + return nil, errors.Errorf("Incorrect data of simple ack: %d", len(data)) + } + data = data[8:12] + m.req.Options.WriteHacks.SimpleAck = true + + m.logger.Debugw("Read RPC_SIMPLE_ACK", + "counter", m.readCounter, + "length", len(data), + ) + + return data, nil +} + +func (m *MTProtoProxy) readCloseExt(data []byte) ([]byte, error) { + m.logger.Debugw("Read RPC_CLOSE_EXT", "counter", m.readCounter) + + return nil, errors.New("Connection has been closed remotely by RPC call") +} + +func (m *MTProtoProxy) Write(p []byte) (int, error) { + defer func() { + m.writeCounter++ + }() + + m.logger.Debugw("Write packet", + "length", len(p), + "counter", m.writeCounter, + "simple_ack", m.req.Options.ReadHacks.SimpleAck, + "quick_ack", m.req.Options.ReadHacks.QuickAck, + ) + + header, flags := m.req.MakeHeader(p) + if ce := m.logger.Desugar().Check(zap.DebugLevel, "RPC_PROXY_REQ header"); ce != nil { + ce.Write( + zap.Int("length", len(p)), + zap.Uint32("counter", m.writeCounter), + zap.Bool("simple_ack", m.req.Options.ReadHacks.QuickAck), + zap.Bool("quick_ack", m.req.Options.ReadHacks.SimpleAck), + zap.String("header", fmt.Sprintf("%v", header.Bytes())), + zap.Stringer("flags", flags), + ) + } + header.Write(p) + + if _, err := m.conn.Write(header.Bytes()); err != nil { + return 0, err + } + + return len(p), nil +} + +func (m *MTProtoProxy) Logger() *zap.SugaredLogger { + return m.logger +} + +func (m *MTProtoProxy) LocalAddr() *net.TCPAddr { + return m.conn.LocalAddr() +} + +func (m *MTProtoProxy) RemoteAddr() *net.TCPAddr { + return m.conn.RemoteAddr() +} + +func (m *MTProtoProxy) Close() error { + return m.conn.Close() +} + +func NewMTProtoProxy(conn PacketReadWriteCloser, connOpts *mtproto.ConnectionOpts, adTag []byte) (PacketReadWriteCloser, error) { + req, err := rpc.NewProxyRequest(connOpts.ClientAddr, conn.LocalAddr(), connOpts, adTag) + if err != nil { + return nil, errors.Annotate(err, "Cannot create new RPC proxy request") + } + + return &MTProtoProxy{ + conn: conn, + logger: conn.Logger().Named("mtproto-proxy"), + req: req, + }, nil +} diff --git a/wrappers/rwcaddr.go b/wrappers/rwcaddr.go deleted file mode 100644 index 6ca4299..0000000 --- a/wrappers/rwcaddr.go +++ /dev/null @@ -1,13 +0,0 @@ -package wrappers - -import ( - "io" - "net" -) - -type ReadWriteCloserWithAddr interface { - io.ReadWriteCloser - - LocalAddr() *net.TCPAddr - RemoteAddr() *net.TCPAddr -} diff --git a/wrappers/streamcipher.go b/wrappers/streamcipher.go new file mode 100644 index 0000000..da89535 --- /dev/null +++ b/wrappers/streamcipher.go @@ -0,0 +1,58 @@ +package wrappers + +import ( + "crypto/cipher" + "net" + + "github.com/juju/errors" + "go.uber.org/zap" +) + +type StreamCipher struct { + encryptor cipher.Stream + decryptor cipher.Stream + conn StreamReadWriteCloser + logger *zap.SugaredLogger +} + +func (s *StreamCipher) Read(p []byte) (int, error) { + n, err := s.conn.Read(p) + if err != nil { + return 0, errors.Annotate(err, "Cannot read stream ciphered data") + } + s.decryptor.XORKeyStream(p, p[:n]) + + return n, nil +} + +func (s *StreamCipher) Write(p []byte) (int, error) { + encrypted := make([]byte, len(p)) + s.encryptor.XORKeyStream(encrypted, p) + + return s.conn.Write(encrypted) +} + +func (s *StreamCipher) Logger() *zap.SugaredLogger { + return s.logger +} + +func (s *StreamCipher) LocalAddr() *net.TCPAddr { + return s.conn.LocalAddr() +} + +func (s *StreamCipher) RemoteAddr() *net.TCPAddr { + return s.conn.RemoteAddr() +} + +func (s *StreamCipher) Close() error { + return s.conn.Close() +} + +func NewStreamCipher(conn StreamReadWriteCloser, encryptor, decryptor cipher.Stream) StreamReadWriteCloser { + return &StreamCipher{ + conn: conn, + logger: conn.Logger().Named("stream-cipher"), + encryptor: encryptor, + decryptor: decryptor, + } +} diff --git a/wrappers/streamcipherrwc.go b/wrappers/streamcipherrwc.go deleted file mode 100644 index a5f15d1..0000000 --- a/wrappers/streamcipherrwc.go +++ /dev/null @@ -1,62 +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() -} - -// 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 952ade1..0000000 --- a/wrappers/timeoutrwc.go +++ /dev/null @@ -1,55 +0,0 @@ -package wrappers - -import ( - "net" - "time" - - "github.com/9seconds/mtg/config" -) - -type TimeoutReadWriteCloserWithAddr struct { - conn net.Conn - 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 NewTimeoutRWC(conn net.Conn, ipv4, ipv6 net.IP) ReadWriteCloserWithAddr { - return &TimeoutReadWriteCloserWithAddr{ - conn: conn, - publicIPv4: ipv4, - publicIPv6: ipv6, - } -} diff --git a/wrappers/trafficrwc.go b/wrappers/trafficrwc.go deleted file mode 100644 index 60dd8c2..0000000 --- a/wrappers/trafficrwc.go +++ /dev/null @@ -1,47 +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() -} - -// 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..2023694 --- /dev/null +++ b/wrappers/wrap.go @@ -0,0 +1,85 @@ +package wrappers + +import ( + "io" + "net" + + "go.uber.org/zap" +) + +type Wrap interface { + Logger() *zap.SugaredLogger + LocalAddr() *net.TCPAddr + RemoteAddr() *net.TCPAddr +} + +type Writer interface { + io.Writer + Wrap +} + +type Closer interface { + io.Closer + Wrap +} + +type WriteCloser interface { + io.Closer + Writer +} + +type StreamReader interface { + io.Reader + Wrap +} + +type StreamReadCloser interface { + io.Closer + StreamReader +} + +type StreamReadWriter interface { + io.Writer + StreamReader +} + +type StreamWriteCloser interface { + io.WriteCloser + Wrap +} + +type StreamReadWriteCloser interface { + io.Closer + StreamReadWriter +} + +type PacketReader interface { + Read() ([]byte, error) + Wrap +} + +type PacketWriter interface { + io.Writer + Wrap +} + +type PacketReadWriter interface { + io.Writer + PacketReader +} + +type PacketReadCloser interface { + io.Closer + PacketReader +} + +type PacketWriteCloser interface { + io.Writer + io.Closer + Wrap +} + +type PacketReadWriteCloser interface { + io.Closer + PacketReadWriter +}