From 83b46d1b8081005c8bc75b7e5b884ce2dd1caaae Mon Sep 17 00:00:00 2001 From: 9seconds Date: Wed, 30 May 2018 21:02:50 +0300 Subject: [PATCH] Add timeoutreadwritecloser --- main.go | 14 +++++++++++++- server/server.go | 41 ++++++++++++++++++++++++----------------- server/telegram.go | 7 ++++--- server/timeoutrwc.go | 35 +++++++++++++++++++++++++++++++++++ 4 files changed, 76 insertions(+), 21 deletions(-) create mode 100644 server/timeoutrwc.go diff --git a/main.go b/main.go index e998d48..4903af2 100644 --- a/main.go +++ b/main.go @@ -39,6 +39,16 @@ var ( Envar("MTG_PORT"). Default("3128"). Uint16() + readTimeout = app.Flag("read-timeout", "Socket read timeout"). + Short('r'). + Envar("MTG_READ_TIMEOUT"). + Default("30s"). + Duration() + writeTimeout = app.Flag("write-timeout", "Socket write timeout"). + Short('w'). + Envar("MTG_WRITE_TIMEOUT"). + Default("30s"). + Duration() serverName = app.Flag("server-name", "Which server name to use. Default is IP address resolved by ipify."). Short('s'). @@ -86,7 +96,9 @@ func main() { )).Sugar() printURLs() - if err := server.NewServer(*bindIP, int(*bindPort), secretBytes, logger).Serve(); err != nil { + srv := server.NewServer(*bindIP, int(*bindPort), secretBytes, logger, + *readTimeout, *writeTimeout) + if err := srv.Serve(); err != nil { logger.Fatal(err.Error()) } } diff --git a/server/server.go b/server/server.go index 89dc6ec..bd5ac00 100644 --- a/server/server.go +++ b/server/server.go @@ -6,6 +6,7 @@ import ( "net" "strconv" "sync" + "time" "github.com/9seconds/mtg/obfuscated2" "github.com/juju/errors" @@ -16,12 +17,14 @@ import ( const bufferSize = 4096 type Server struct { - ip net.IP - port int - secret []byte - logger *zap.SugaredLogger - lsock net.Listener - ctx context.Context + ip net.IP + port int + secret []byte + logger *zap.SugaredLogger + lsock net.Listener + ctx context.Context + readTimeout time.Duration + writeTimeout time.Duration } func (s *Server) Serve() error { @@ -98,7 +101,8 @@ func (s *Server) makeSocketID() string { } func (s *Server) getClientStream(conn net.Conn, ctx context.Context, cancel context.CancelFunc, socketID string) (io.ReadWriteCloser, int16, error) { - frame, err := obfuscated2.ExtractFrame(conn) + wConn := newTimeoutReadWriteCloser(conn, s.readTimeout, s.writeTimeout) + frame, err := obfuscated2.ExtractFrame(wConn) if err != nil { return nil, 0, errors.Annotate(err, "Cannot create client stream") } @@ -108,7 +112,7 @@ func (s *Server) getClientStream(conn net.Conn, ctx context.Context, cancel cont return nil, 0, errors.Annotate(err, "Cannot create client stream") } - wConn := newLogReadWriteCloser(conn, s.logger, socketID, "client") + wConn = newLogReadWriteCloser(wConn, s.logger, socketID, "client") wConn = newCipherReadWriteCloser(conn, obfs2) wConn = newCtxReadWriteCloser(wConn, ctx, cancel) @@ -116,17 +120,18 @@ func (s *Server) getClientStream(conn net.Conn, ctx context.Context, cancel cont } func (s *Server) getTelegramStream(dc int16, ctx context.Context, cancel context.CancelFunc, socketID string) (io.ReadWriteCloser, error) { - socket, err := dialToTelegram(dc) + socket, err := dialToTelegram(dc, s.readTimeout) if err != nil { return nil, errors.Annotate(err, "Cannot dial") } + wConn := newTimeoutReadWriteCloser(socket, s.readTimeout, s.writeTimeout) obfs2, frame := obfuscated2.MakeTelegramObfuscated2Frame() if n, err := socket.Write(frame); err != nil || n != len(frame) { return nil, errors.Annotate(err, "Cannot write hadnshake frame") } - wConn := newLogReadWriteCloser(socket, s.logger, socketID, "telegram") + wConn = newLogReadWriteCloser(socket, s.logger, socketID, "telegram") wConn = newCipherReadWriteCloser(wConn, obfs2) wConn = newCtxReadWriteCloser(wConn, ctx, cancel) @@ -138,15 +143,17 @@ func (s *Server) pipe(wait *sync.WaitGroup, reader io.Reader, writer io.Writer) buf := make([]byte, bufferSize) io.CopyBuffer(writer, reader, buf) - } -func NewServer(ip net.IP, port int, secret []byte, logger *zap.SugaredLogger) *Server { +func NewServer(ip net.IP, port int, secret []byte, logger *zap.SugaredLogger, + readTimeout, writeTimeout time.Duration) *Server { return &Server{ - ip: ip, - port: port, - secret: secret, - ctx: context.Background(), - logger: logger, + ip: ip, + port: port, + secret: secret, + ctx: context.Background(), + logger: logger, + readTimeout: readTimeout, + writeTimeout: writeTimeout, } } diff --git a/server/telegram.go b/server/telegram.go index 1d36cec..22e1335 100644 --- a/server/telegram.go +++ b/server/telegram.go @@ -17,13 +17,14 @@ var telegramDCIPs = [5]string{ const telegramKeepAlive = 30 * time.Second -func dialToTelegram(dcIdx int16) (net.Conn, error) { +func dialToTelegram(dcIdx int16, timeout time.Duration) (net.Conn, error) { if dcIdx < 0 || dcIdx >= 5 { return nil, errors.New("Incorrect DC IDX") } - tcpAddr, _ := net.ResolveTCPAddr("tcp", telegramDCIPs[dcIdx]) - conn, err := net.DialTCP("tcp", nil, tcpAddr) + dialer := net.Dialer{Timeout: timeout} + rawConn, err := dialer.Dial("tcp", telegramDCIPs[dcIdx]) + conn := rawConn.(*net.TCPConn) if err != nil { return nil, errors.Annotate(err, "Cannot dial") } diff --git a/server/timeoutrwc.go b/server/timeoutrwc.go new file mode 100644 index 0000000..2e56b40 --- /dev/null +++ b/server/timeoutrwc.go @@ -0,0 +1,35 @@ +package server + +import ( + "io" + "net" + "time" +) + +type TimeoutReadWriteCloser struct { + conn net.Conn + readTimeout time.Duration + writeTimeout time.Duration +} + +func (t *TimeoutReadWriteCloser) Read(p []byte) (int, error) { + t.conn.SetReadDeadline(time.Now().Add(t.readTimeout)) + return t.conn.Read(p) +} + +func (t *TimeoutReadWriteCloser) Write(p []byte) (int, error) { + t.conn.SetWriteDeadline(time.Now().Add(t.writeTimeout)) + return t.conn.Write(p) +} + +func (t *TimeoutReadWriteCloser) Close() error { + return t.conn.Close() +} + +func newTimeoutReadWriteCloser(conn net.Conn, readTimeout, writeTimeout time.Duration) io.ReadWriteCloser { + return &TimeoutReadWriteCloser{ + conn: conn, + readTimeout: readTimeout, + writeTimeout: writeTimeout, + } +}