Remove contexts

This commit is contained in:
9seconds
2018-07-08 12:10:27 +03:00
parent 4fc6b07448
commit 1fe3319988
6 changed files with 22 additions and 100 deletions
+1 -2
View File
@@ -1,7 +1,6 @@
package client package client
import ( import (
"context"
"net" "net"
"github.com/9seconds/mtg/config" "github.com/9seconds/mtg/config"
@@ -9,4 +8,4 @@ import (
"github.com/9seconds/mtg/wrappers" "github.com/9seconds/mtg/wrappers"
) )
type Init func(context.Context, context.CancelFunc, net.Conn, string, *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) type Init func(net.Conn, string, *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error)
+1 -4
View File
@@ -1,7 +1,6 @@
package client package client
import ( import (
"context"
"net" "net"
"time" "time"
@@ -15,8 +14,7 @@ import (
const handshakeTimeout = 10 * time.Second const handshakeTimeout = 10 * time.Second
func DirectInit(ctx context.Context, cancel context.CancelFunc, socket net.Conn, connID string, func DirectInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) {
conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) {
if err := config.SetSocketOptions(socket); err != nil { if err := config.SetSocketOptions(socket); err != nil {
return nil, nil, errors.Annotate(err, "Cannot set socket options") return nil, nil, errors.Annotate(err, "Cannot set socket options")
} }
@@ -37,7 +35,6 @@ func DirectInit(ctx context.Context, cancel context.CancelFunc, socket net.Conn,
connOpts.ClientAddr = conn.RemoteAddr() connOpts.ClientAddr = conn.RemoteAddr()
conn = wrappers.NewStreamCipher(conn, obfs2.Encryptor, obfs2.Decryptor) conn = wrappers.NewStreamCipher(conn, obfs2.Encryptor, obfs2.Decryptor)
conn = wrappers.NewCtx(ctx, cancel, conn)
return conn, connOpts, nil return conn, connOpts, nil
} }
+2 -4
View File
@@ -1,7 +1,6 @@
package client package client
import ( import (
"context"
"net" "net"
"github.com/9seconds/mtg/config" "github.com/9seconds/mtg/config"
@@ -9,9 +8,8 @@ import (
"github.com/9seconds/mtg/wrappers" "github.com/9seconds/mtg/wrappers"
) )
func MiddleInit(ctx context.Context, cancel context.CancelFunc, socket net.Conn, connID string, func MiddleInit(socket net.Conn, connID string, conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) {
conf *config.Config) (wrappers.Wrap, *mtproto.ConnectionOpts, error) { conn, opts, err := DirectInit(socket, connID, conf)
conn, opts, err := DirectInit(ctx, cancel, socket, connID, conf)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
+17 -13
View File
@@ -1,7 +1,6 @@
package proxy package proxy
import ( import (
"context"
"io" "io"
"net" "net"
"sync" "sync"
@@ -52,15 +51,14 @@ func (p *Proxy) accept(conn net.Conn) {
log.Infow("Client connected", "addr", conn.RemoteAddr()) log.Infow("Client connected", "addr", conn.RemoteAddr())
ctx, cancel := context.WithCancel(context.Background()) client, opts, err := p.clientInit(conn, connID, p.conf)
client, 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
} }
defer client.(io.Closer).Close() defer client.(io.Closer).Close()
server, err := p.getTelegramConn(ctx, cancel, opts, connID) server, err := p.getTelegramConn(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
@@ -82,19 +80,16 @@ func (p *Proxy) accept(conn net.Conn) {
go p.directPipe(serverStream, clientStream, wait) go p.directPipe(serverStream, clientStream, wait)
} }
<-ctx.Done()
wait.Wait() wait.Wait()
log.Infow("Client disconnected", "addr", conn.RemoteAddr()) log.Infow("Client disconnected", "addr", conn.RemoteAddr())
} }
func (p *Proxy) getTelegramConn(ctx context.Context, cancel context.CancelFunc, opts *mtproto.ConnectionOpts, func (p *Proxy) getTelegramConn(opts *mtproto.ConnectionOpts, connID string) (wrappers.Wrap, error) {
connID string) (wrappers.Wrap, error) {
streamConn, err := p.tg.Dial(connID, opts) streamConn, err := p.tg.Dial(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")
} }
streamConn = wrappers.NewCtx(ctx, cancel, streamConn)
packetConn, err := p.tg.Init(opts, streamConn) packetConn, err := p.tg.Init(opts, streamConn)
if err != nil { if err != nil {
@@ -104,8 +99,13 @@ func (p *Proxy) getTelegramConn(ctx context.Context, cancel context.CancelFunc,
return packetConn, nil return packetConn, nil
} }
func (p *Proxy) middlePipe(src wrappers.PacketReader, dst wrappers.PacketWriter, wait *sync.WaitGroup, hacks *mtproto.Hacks) { func (p *Proxy) middlePipe(src wrappers.PacketReadCloser, dst wrappers.PacketWriteCloser, wait *sync.WaitGroup, hacks *mtproto.Hacks) {
defer wait.Done() defer func() {
src.Close()
dst.Close()
wait.Done()
}()
for { for {
hacks.SimpleAck = false hacks.SimpleAck = false
hacks.QuickAck = false hacks.QuickAck = false
@@ -120,10 +120,14 @@ func (p *Proxy) middlePipe(src wrappers.PacketReader, dst wrappers.PacketWriter,
} }
} }
func (p *Proxy) directPipe(src io.Reader, dst io.Writer, wait *sync.WaitGroup) { func (p *Proxy) directPipe(src io.ReadCloser, dst io.WriteCloser, wait *sync.WaitGroup) {
defer wait.Done() defer func() {
io.Copy(dst, src) src.Close()
dst.Close()
wait.Done()
}()
io.Copy(dst, src)
} }
func NewProxy(conf *config.Config) *Proxy { func NewProxy(conf *config.Config) *Proxy {
-76
View File
@@ -1,76 +0,0 @@
package wrappers
import (
"context"
"net"
"github.com/juju/errors"
)
type Ctx struct {
cancel context.CancelFunc
conn StreamReadWriteCloser
ctx context.Context
}
func (c *Ctx) 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
}
}
func (c *Ctx) 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
}
}
func (c *Ctx) LogDebug(msg string, data ...interface{}) {
c.conn.LogDebug(msg, data...)
}
func (c *Ctx) LogInfo(msg string, data ...interface{}) {
c.conn.LogInfo(msg, data...)
}
func (c *Ctx) LogWarn(msg string, data ...interface{}) {
c.conn.LogWarn(msg, data...)
}
func (c *Ctx) LogError(msg string, data ...interface{}) {
c.conn.LogError(msg, data...)
}
func (c *Ctx) LocalAddr() *net.TCPAddr {
return c.conn.LocalAddr()
}
func (c *Ctx) RemoteAddr() *net.TCPAddr {
return c.conn.RemoteAddr()
}
func (c *Ctx) Close() error {
return c.conn.Close()
}
func NewCtx(ctx context.Context, cancel context.CancelFunc, conn StreamReadWriteCloser) StreamReadWriteCloser {
return &Ctx{
ctx: ctx,
cancel: cancel,
conn: conn,
}
}
+1 -1
View File
@@ -71,7 +71,7 @@ type PacketReadWriter interface {
PacketReader PacketReader
} }
type BlockReadCloser interface { type PacketReadCloser interface {
io.Closer io.Closer
PacketReader PacketReader
} }