Use contexts for Conn wrapper

This commit is contained in:
9seconds
2018-07-28 13:21:57 +03:00
parent b0d86abc74
commit 9f20e8749a
9 changed files with 73 additions and 32 deletions
+3 -1
View File
@@ -1,6 +1,7 @@
package client package client
import ( import (
"context"
"net" "net"
"github.com/9seconds/mtg/config" "github.com/9seconds/mtg/config"
@@ -9,4 +10,5 @@ import (
) )
// Init defines common method for initializing client connections. // Init defines common method for initializing client connections.
type Init func(net.Conn, string, *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) type Init func(context.Context, context.CancelFunc, net.Conn, string,
*config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error)
+4 -2
View File
@@ -1,6 +1,7 @@
package client package client
import ( import (
"context"
"net" "net"
"time" "time"
@@ -16,7 +17,8 @@ const handshakeTimeout = 10 * time.Second
// DirectInit initializes client connection for proxy which connects to // DirectInit initializes client connection for proxy which connects to
// Telegram directly. // Telegram directly.
func DirectInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) { func DirectInit(ctx context.Context, cancel context.CancelFunc, socket net.Conn,
connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) {
tcpSocket := socket.(*net.TCPConn) tcpSocket := socket.(*net.TCPConn)
if err := tcpSocket.SetNoDelay(false); err != nil { if err := tcpSocket.SetNoDelay(false); err != nil {
return nil, nil, errors.Annotate(err, "Cannot disable NO_DELAY to client socket") return nil, nil, errors.Annotate(err, "Cannot disable NO_DELAY to client socket")
@@ -35,7 +37,7 @@ func DirectInit(socket net.Conn, connID string, conf *config.Config) (wrappers.W
} }
socket.SetReadDeadline(time.Time{}) // nolint: errcheck socket.SetReadDeadline(time.Time{}) // nolint: errcheck
conn := wrappers.NewConn(socket, connID, wrappers.ConnPurposeClient, conf.PublicIPv4, conf.PublicIPv6) conn := wrappers.NewConn(ctx, cancel, socket, connID, wrappers.ConnPurposeClient, conf.PublicIPv4, conf.PublicIPv6)
obfs2, connOpts, err := obfuscated2.ParseObfuscated2ClientFrame(conf.Secret, frame) obfs2, connOpts, err := obfuscated2.ParseObfuscated2ClientFrame(conf.Secret, frame)
if err != nil { if err != nil {
return nil, nil, errors.Annotate(err, "Cannot parse obfuscated frame") return nil, nil, errors.Annotate(err, "Cannot parse obfuscated frame")
+4 -2
View File
@@ -1,6 +1,7 @@
package client package client
import ( import (
"context"
"net" "net"
"github.com/9seconds/mtg/config" "github.com/9seconds/mtg/config"
@@ -10,8 +11,9 @@ import (
// MiddleInit initializes client connection for proxy which has to // MiddleInit initializes client connection for proxy which has to
// support promoted channels, connect to Telegram middle proxies etc. // support promoted channels, connect to Telegram middle proxies etc.
func MiddleInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) { func MiddleInit(ctx context.Context, cancel context.CancelFunc, socket net.Conn,
conn, opts, err := DirectInit(socket, connID, conf) connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) {
conn, opts, err := DirectInit(ctx, cancel, socket, connID, conf)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
+7 -4
View File
@@ -1,6 +1,7 @@
package proxy package proxy
import ( import (
"context"
"io" "io"
"net" "net"
"sync" "sync"
@@ -43,6 +44,7 @@ func (p *Proxy) Serve() error {
func (p *Proxy) accept(conn net.Conn) { func (p *Proxy) accept(conn net.Conn) {
connID := uuid.NewV4().String() connID := uuid.NewV4().String()
log := zap.S().With("connection_id", connID).Named("main") log := zap.S().With("connection_id", connID).Named("main")
ctx, cancel := context.WithCancel(context.Background())
defer func() { defer func() {
conn.Close() // nolint: errcheck conn.Close() // nolint: errcheck
@@ -55,7 +57,7 @@ func (p *Proxy) accept(conn net.Conn) {
log.Infow("Client connected", "addr", conn.RemoteAddr()) log.Infow("Client connected", "addr", conn.RemoteAddr())
clientConn, opts, err := p.clientInit(conn, connID, p.conf) clientConn, opts, err := p.clientInit(ctx, cancel, conn, connID, p.conf)
if err != nil { if err != nil {
log.Errorw("Cannot initialize client connection", "error", err) log.Errorw("Cannot initialize client connection", "error", err)
return return
@@ -65,7 +67,7 @@ func (p *Proxy) accept(conn net.Conn) {
stats.ClientConnected(opts.ConnectionType, clientConn.RemoteAddr()) stats.ClientConnected(opts.ConnectionType, clientConn.RemoteAddr())
defer stats.ClientDisconnected(opts.ConnectionType, clientConn.RemoteAddr()) defer stats.ClientDisconnected(opts.ConnectionType, clientConn.RemoteAddr())
serverConn, err := p.getTelegramConn(opts, connID) serverConn, err := p.getTelegramConn(ctx, cancel, opts, connID)
if err != nil { if err != nil {
log.Errorw("Cannot initialize server connection", "error", err) log.Errorw("Cannot initialize server connection", "error", err)
return return
@@ -92,8 +94,9 @@ func (p *Proxy) accept(conn net.Conn) {
log.Infow("Client disconnected", "addr", conn.RemoteAddr()) log.Infow("Client disconnected", "addr", conn.RemoteAddr())
} }
func (p *Proxy) getTelegramConn(opts *mtproto.ConnectionOpts, connID string) (wrappers.Wrap, error) { func (p *Proxy) getTelegramConn(ctx context.Context, cancel context.CancelFunc,
streamConn, err := p.tg.Dial(connID, opts) opts *mtproto.ConnectionOpts, connID string) (wrappers.Wrap, error) {
streamConn, err := p.tg.Dial(ctx, cancel, connID, opts)
if err != nil { if err != nil {
return nil, errors.Annotate(err, "Cannot dial to Telegram") return nil, errors.Annotate(err, "Cannot dial to Telegram")
} }
+5 -2
View File
@@ -1,6 +1,7 @@
package telegram package telegram
import ( import (
"context"
"net" "net"
"time" "time"
@@ -38,12 +39,14 @@ func (t *tgDialer) dial(addr string) (net.Conn, error) {
return conn, nil return conn, nil
} }
func (t *tgDialer) dialRWC(addr, connID string) (wrappers.StreamReadWriteCloser, error) { func (t *tgDialer) dialRWC(ctx context.Context, cancel context.CancelFunc,
addr, connID string) (wrappers.StreamReadWriteCloser, error) {
conn, err := t.dial(addr) conn, err := t.dial(addr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
tgConn := wrappers.NewConn(conn, connID, wrappers.ConnPurposeTelegram, t.conf.PublicIPv4, t.conf.PublicIPv6) tgConn := wrappers.NewConn(ctx, cancel, conn, connID,
wrappers.ConnPurposeTelegram, t.conf.PublicIPv4, t.conf.PublicIPv6)
return tgConn, nil return tgConn, nil
} }
+4 -2
View File
@@ -1,6 +1,7 @@
package telegram package telegram
import ( import (
"context"
"net" "net"
"github.com/juju/errors" "github.com/juju/errors"
@@ -32,7 +33,8 @@ type directTelegram struct {
baseTelegram baseTelegram
} }
func (t *directTelegram) Dial(connID string, connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) { func (t *directTelegram) Dial(ctx context.Context, cancel context.CancelFunc,
connID string, connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) {
dc := connOpts.DC dc := connOpts.DC
if dc < 0 { if dc < 0 {
dc = -dc dc = -dc
@@ -40,7 +42,7 @@ func (t *directTelegram) Dial(connID string, connOpts *mtproto.ConnectionOpts) (
dc = 1 dc = 1
} }
return t.baseTelegram.dial(dc-1, connID, connOpts.ConnectionProto) return t.baseTelegram.dial(ctx, cancel, dc-1, connID, connOpts.ConnectionProto)
} }
func (t *directTelegram) Init(connOpts *mtproto.ConnectionOpts, func (t *directTelegram) Init(connOpts *mtproto.ConnectionOpts,
+3 -2
View File
@@ -2,6 +2,7 @@ package telegram
import ( import (
"bufio" "bufio"
"context"
"io/ioutil" "io/ioutil"
"net" "net"
"net/http" "net/http"
@@ -38,7 +39,7 @@ type middleTelegramCaller struct {
httpClient *http.Client httpClient *http.Client
} }
func (t *middleTelegramCaller) Dial(connID string, func (t *middleTelegramCaller) Dial(ctx context.Context, cancel context.CancelFunc, connID string,
connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) { connOpts *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) {
dc := connOpts.DC dc := connOpts.DC
if dc == 0 { if dc == 0 {
@@ -47,7 +48,7 @@ func (t *middleTelegramCaller) Dial(connID string,
t.dialerMutex.RLock() t.dialerMutex.RLock()
defer t.dialerMutex.RUnlock() defer t.dialerMutex.RUnlock()
return t.baseTelegram.dial(dc, connID, connOpts.ConnectionProto) return t.baseTelegram.dial(ctx, cancel, dc, connID, connOpts.ConnectionProto)
} }
func (t *middleTelegramCaller) autoUpdate() { func (t *middleTelegramCaller) autoUpdate() {
+4 -3
View File
@@ -1,6 +1,7 @@
package telegram package telegram
import ( import (
"context"
"math/rand" "math/rand"
"github.com/juju/errors" "github.com/juju/errors"
@@ -11,7 +12,7 @@ import (
// Telegram is an interface for different Telegram work modes. // Telegram is an interface for different Telegram work modes.
type Telegram interface { type Telegram interface {
Dial(string, *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error) Dial(context.Context, context.CancelFunc, string, *mtproto.ConnectionOpts) (wrappers.StreamReadWriteCloser, error)
Init(*mtproto.ConnectionOpts, wrappers.StreamReadWriteCloser) (wrappers.Wrap, error) Init(*mtproto.ConnectionOpts, wrappers.StreamReadWriteCloser) (wrappers.Wrap, error)
} }
@@ -22,7 +23,7 @@ type baseTelegram struct {
v6Addresses map[int16][]string v6Addresses map[int16][]string
} }
func (b *baseTelegram) dial(dcIdx int16, connID string, func (b *baseTelegram) dial(ctx context.Context, cancel context.CancelFunc, dcIdx int16, connID string,
proto mtproto.ConnectionProtocol) (wrappers.StreamReadWriteCloser, error) { proto mtproto.ConnectionProtocol) (wrappers.StreamReadWriteCloser, error) {
addrs := make([]string, 2) addrs := make([]string, 2)
@@ -38,7 +39,7 @@ func (b *baseTelegram) dial(dcIdx int16, connID string,
} }
for _, addr := range addrs { for _, addr := range addrs {
if conn, err := b.dialer.dialRWC(addr, connID); err == nil { if conn, err := b.dialer.dialRWC(ctx, cancel, addr, connID); err == nil {
return conn, err return conn, err
} }
} }
+39 -14
View File
@@ -1,12 +1,14 @@
package wrappers package wrappers
import ( import (
"context"
"net" "net"
"time" "time"
"go.uber.org/zap" "go.uber.org/zap"
"github.com/9seconds/mtg/stats" "github.com/9seconds/mtg/stats"
"github.com/juju/errors"
) )
// ConnPurpose is intended to be identifier of connection purpose. We // ConnPurpose is intended to be identifier of connection purpose. We
@@ -39,8 +41,10 @@ const (
// Conn is a basic wrapper for net.Conn providing the most low-level // Conn is a basic wrapper for net.Conn providing the most low-level
// logic and management as possible. // logic and management as possible.
type Conn struct { type Conn struct {
connID string
conn net.Conn conn net.Conn
ctx context.Context
cancel context.CancelFunc
connID string
logger *zap.SugaredLogger logger *zap.SugaredLogger
publicIPv4 net.IP publicIPv4 net.IP
@@ -48,28 +52,46 @@ type Conn struct {
} }
func (c *Conn) Write(p []byte) (int, error) { func (c *Conn) Write(p []byte) (int, error) {
c.conn.SetWriteDeadline(time.Now().Add(connTimeoutWrite)) // nolint: errcheck select {
n, err := c.conn.Write(p) case <-c.ctx.Done():
return 0, errors.Annotate(c.ctx.Err(), "Cannot write because context was closed")
default:
c.conn.SetWriteDeadline(time.Now().Add(connTimeoutWrite)) // nolint: errcheck
n, err := c.conn.Write(p)
if err != nil {
c.cancel()
}
c.logger.Debugw("Write to stream", "bytes", n, "error", err) c.logger.Debugw("Write to stream", "bytes", n, "error", err)
stats.EgressTraffic(n) stats.EgressTraffic(n)
return n, err return n, err
}
} }
func (c *Conn) Read(p []byte) (int, error) { func (c *Conn) Read(p []byte) (int, error) {
c.conn.SetReadDeadline(time.Now().Add(connTimeoutRead)) // nolint: errcheck select {
n, err := c.conn.Read(p) case <-c.ctx.Done():
return 0, errors.Annotate(c.ctx.Err(), "Cannot read because context was closed")
default:
c.conn.SetReadDeadline(time.Now().Add(connTimeoutRead)) // nolint: errcheck
n, err := c.conn.Read(p)
if err != nil {
c.cancel()
}
c.logger.Debugw("Read from stream", "bytes", n, "error", err) c.logger.Debugw("Read from stream", "bytes", n, "error", err)
stats.IngressTraffic(n) stats.IngressTraffic(n)
return n, err return n, err
}
} }
// Close closes underlying net.Conn instance. // Close closes underlying net.Conn instance.
func (c *Conn) Close() error { func (c *Conn) Close() error {
defer c.logger.Debugw("Close connection") defer c.logger.Debugw("Close connection")
c.cancel()
return c.conn.Close() return c.conn.Close()
} }
@@ -100,7 +122,8 @@ func (c *Conn) RemoteAddr() *net.TCPAddr {
} }
// NewConn initializes Conn wrapper for net.Conn. // NewConn initializes Conn wrapper for net.Conn.
func NewConn(conn net.Conn, connID string, purpose ConnPurpose, publicIPv4, publicIPv6 net.IP) StreamReadWriteCloser { func NewConn(ctx context.Context, cancel context.CancelFunc, conn net.Conn,
connID string, purpose ConnPurpose, publicIPv4, publicIPv6 net.IP) StreamReadWriteCloser {
logger := zap.S().With( logger := zap.S().With(
"connection_id", connID, "connection_id", connID,
"local_address", conn.LocalAddr(), "local_address", conn.LocalAddr(),
@@ -109,9 +132,11 @@ func NewConn(conn net.Conn, connID string, purpose ConnPurpose, publicIPv4, publ
).Named("conn") ).Named("conn")
wrapper := Conn{ wrapper := Conn{
logger: logger,
connID: connID,
conn: conn, conn: conn,
ctx: ctx,
cancel: cancel,
connID: connID,
logger: logger,
publicIPv4: publicIPv4, publicIPv4: publicIPv4,
publicIPv6: publicIPv6, publicIPv6: publicIPv6,
} }