From c58251a0994945be2c670167b1c3ca034d0a40c7 Mon Sep 17 00:00:00 2001 From: Page Fault Date: Fri, 12 Jun 2020 06:02:06 +0000 Subject: [PATCH] fix goroutine leak and deadlock --- README.md | 3 +- build/base.go | 1 - build/other.go | 1 + docs/content/basic/config.md | 28 +---------- log/simplelog/simplelog.go | 3 ++ proxy/proxy.go | 62 +++++++++++++++--------- test/scenario/proxy_test.go | 91 +++++++++++++++++------------------- tunnel/dokodemo/conn.go | 15 +++--- tunnel/dokodemo/server.go | 23 +++++---- tunnel/mux/client.go | 26 ++++++----- tunnel/mux/server.go | 56 +++++++++++----------- tunnel/router/client.go | 22 ++++----- tunnel/router/conn.go | 12 ++--- tunnel/shadowsocks/conn.go | 1 + tunnel/simplesocks/server.go | 12 ++--- tunnel/socks/config.go | 9 ++-- tunnel/socks/conn.go | 13 ++---- tunnel/socks/server.go | 15 ++++-- tunnel/tproxy/server.go | 19 ++++++-- tunnel/tproxy/tproxy_test.go | 30 ++++++++++++ tunnel/transport/server.go | 17 ++++--- tunnel/trojan/server.go | 13 ++++-- 22 files changed, 268 insertions(+), 204 deletions(-) create mode 100644 tunnel/tproxy/tproxy_test.go diff --git a/README.md b/README.md index e45c552..99e2f79 100644 --- a/README.md +++ b/README.md @@ -163,7 +163,8 @@ server.json ], "ssl": { "cert": "your_cert.crt", - "key": "your_key.key" + "key": "your_key.key", + "sni": "www.your_awesome_domain_name.com" } } ``` diff --git a/build/base.go b/build/base.go index d6e8461..020140b 100644 --- a/build/base.go +++ b/build/base.go @@ -1,7 +1,6 @@ package build import ( - _ "github.com/p4gefau1t/trojan-go/log/golog" _ "github.com/p4gefau1t/trojan-go/statistic/memory" _ "github.com/p4gefau1t/trojan-go/version" ) diff --git a/build/other.go b/build/other.go index fc695d8..5fd1727 100644 --- a/build/other.go +++ b/build/other.go @@ -4,4 +4,5 @@ package build import ( _ "github.com/p4gefau1t/trojan-go/easy" + _ "github.com/p4gefau1t/trojan-go/log/golog" ) diff --git a/docs/content/basic/config.md b/docs/content/basic/config.md index 8cfe8b4..6dd7984 100644 --- a/docs/content/basic/config.md +++ b/docs/content/basic/config.md @@ -14,33 +14,7 @@ weight: 22 - Trojan-Go,可以从release页面下载 -### 配置证书 - -为了伪装成一个正常的HTTPS站点,也为了保证传输的安全,我们需要一份经过权威证书机构签名的证书。Trojan-Go支持从Let's Encrypt自动申请证书。首先将你的域名正确解析到你的服务器IP。然后准备好一个邮箱地址,合乎邮箱地址规则即可,不需要真实邮箱地址。保证你的服务器443和80端口没有被其他程序(nginx,apache,正在运行的Trojan等)占用。然后执行 - -```shell -sudo ./trojan-go -autocert request -``` - -按照屏幕提示填入相关信息。如果操作成功,当前目录下将得到四个文件 - -- server.key 服务器私钥 - -- server.crt 经过Let's Encrypt签名的服务器证书 - -- user.key 用户Email对应的私钥 - -- domain_info.json 域名和用户Email信息 - -备份好这些文件,不要将.key文件分享给其他任何人,否则你的身份可能被冒用。 - -证书的有效期通常是三个月,你可以使用 - -```shell -sudo ./trojan-go -autocert renew -``` - -进行证书更新。更新之前请确保同目录下有上述的四个文件。如果你没有指定ACME challenge使用的端口,Trojan-Go将默认使用443和80端口,请确保这两个端口没有被Trojan-Go或者其他程序(nginx, caddy等等)占用。 +- 证书密钥对,可以从letsencrpyt等机构免费申请签发 ### 服务端配置 diff --git a/log/simplelog/simplelog.go b/log/simplelog/simplelog.go index 4e4d84a..aea5175 100644 --- a/log/simplelog/simplelog.go +++ b/log/simplelog/simplelog.go @@ -3,6 +3,7 @@ package simplelog import ( "io" golog "log" + "os" "github.com/p4gefau1t/trojan-go/log" ) @@ -23,12 +24,14 @@ func (l *SimpleLogger) Fatal(v ...interface{}) { if l.logLevel <= log.FatalLevel { golog.Fatal(v...) } + os.Exit(1) } func (l *SimpleLogger) Fatalf(format string, v ...interface{}) { if l.logLevel <= log.FatalLevel { golog.Fatalf(format, v...) } + os.Exit(1) } func (l *SimpleLogger) Error(v ...interface{}) { diff --git a/proxy/proxy.go b/proxy/proxy.go index f881a09..401133a 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -5,6 +5,7 @@ import ( "io" "math/rand" "net" + "os" "strings" "github.com/p4gefau1t/trojan-go/common" @@ -23,7 +24,6 @@ const ( type Proxy struct { sources []tunnel.Server sink tunnel.Client - errChan chan error ctx context.Context cancel context.CancelFunc } @@ -31,15 +31,17 @@ type Proxy struct { func (p *Proxy) Run() error { p.relayConnLoop() p.relayPacketLoop() - return <-p.errChan + <-p.ctx.Done() + return nil } func (p *Proxy) Close() error { p.cancel() + p.sink.Close() for _, source := range p.sources { source.Close() } - return p.sink.Close() + return nil } func (p *Proxy) relayConnLoop() { @@ -73,9 +75,14 @@ func (p *Proxy) relayConnLoop() { } go copyConn(inbound, outbound) go copyConn(outbound, inbound) - err = <-errChan - if err != nil { - log.Error(err) + select { + case err = <-errChan: + if err != nil { + log.Error(err) + } + case <-p.ctx.Done(): + log.Debug("shutting down conn relay") + return } log.Debug("conn relay ends") }(inbound) @@ -109,23 +116,33 @@ func (p *Proxy) relayPacketLoop() { defer outbound.Close() errChan := make(chan error, 2) copyPacket := func(a, b tunnel.PacketConn) { - buf := make([]byte, MaxPacketSize) - n, metadata, err := a.ReadWithMetadata(buf) - if err != nil { - errChan <- err - return - } - n, err = b.WriteWithMetadata(buf[:n], metadata) - if err != nil { - errChan <- err - return + for { + buf := make([]byte, MaxPacketSize) + n, metadata, err := a.ReadWithMetadata(buf) + if err != nil { + errChan <- err + return + } + if n == 0 { + errChan <- nil + return + } + n, err = b.WriteWithMetadata(buf[:n], metadata) + if err != nil { + errChan <- err + return + } } } go copyPacket(inbound, outbound) go copyPacket(outbound, inbound) - err = <-errChan - if err != nil { - log.Error(err) + select { + case err = <-errChan: + if err != nil { + log.Error(err) + } + case <-p.ctx.Done(): + log.Debug("shutting down packet relay") } log.Debug("packet relay ends") }(inbound) @@ -139,7 +156,6 @@ func NewProxy(ctx context.Context, sources []tunnel.Server, sink tunnel.Client) return &Proxy{ sources: sources, sink: sink, - errChan: make(chan error, 32), ctx: ctx, cancel: cancel, } @@ -175,7 +191,11 @@ func NewProxyFromConfigData(data []byte, isJSON bool) (*Proxy, error) { } log.SetLogLevel(log.LogLevel(cfg.LogLevel)) if cfg.LogFile != "" { - + file, err := os.OpenFile(cfg.LogFile, os.O_APPEND|os.O_CREATE, 0600) + if err != nil { + return nil, common.NewError("failed to open log file").Base(err) + } + log.SetOutput(file) } return create(ctx) } diff --git a/test/scenario/proxy_test.go b/test/scenario/proxy_test.go index 77c9a50..5f6e6e0 100644 --- a/test/scenario/proxy_test.go +++ b/test/scenario/proxy_test.go @@ -6,6 +6,7 @@ import ( "github.com/p4gefau1t/trojan-go/test/util" "io/ioutil" "net" + "sync" "testing" "time" @@ -79,6 +80,47 @@ func init() { ioutil.WriteFile("server.key", []byte(key), 0777) } +func CheckClientServer(clientData, serverData string, socksPort int) (ok bool) { + server, err := proxy.NewProxyFromConfigData([]byte(clientData), false) + common.Must(err) + go server.Run() + + client, err := proxy.NewProxyFromConfigData([]byte(serverData), false) + common.Must(err) + go client.Run() + + time.Sleep(time.Second * 2) + dialer, err := netproxy.SOCKS5("tcp", fmt.Sprintf("127.0.0.1:%d", socksPort), nil, netproxy.Direct) + + ok = true + const num = 100 + wg := sync.WaitGroup{} + wg.Add(num) + for i := 0; i < num; i++ { + go func() { + const payloadSize = 1024 + payload := util.GeneratePayload(payloadSize) + buf := [payloadSize]byte{} + + conn, err := dialer.Dial("tcp", util.EchoAddr) + common.Must(err) + + common.Must2(conn.Write(payload)) + common.Must2(conn.Read(buf[:])) + + if !bytes.Equal(payload, buf[:]) { + ok = false + } + conn.Close() + wg.Done() + }() + } + wg.Wait() + client.Close() + server.Close() + return +} + func TestClientServerWebsocketSubTree(t *testing.T) { serverPort := common.PickPort("tcp", "127.0.0.1") socksPort := common.PickPort("tcp", "127.0.0.1") @@ -105,11 +147,6 @@ shadowsocks: mux: enabled: true `, socksPort, serverPort) - go func() { - proxy, err := proxy.NewProxyFromConfigData([]byte(clientData), false) - common.Must(err) - common.Must(proxy.Run()) - }() serverData := fmt.Sprintf(` run-type: server @@ -134,25 +171,8 @@ websocket: path: /ws hostname: 127.0.0.1 `, serverPort, util.HTTPPort) - go func() { - proxy, err := proxy.NewProxyFromConfigData([]byte(serverData), false) - common.Must(err) - common.Must(proxy.Run()) - }() - time.Sleep(time.Second * 2) - dialer, err := netproxy.SOCKS5("tcp", fmt.Sprintf("127.0.0.1:%d", socksPort), nil, netproxy.Direct) - - payload := util.GeneratePayload(1024) - buf := [1024]byte{} - - conn, err := dialer.Dial("tcp", util.EchoAddr) - common.Must(err) - - common.Must2(conn.Write(payload)) - common.Must2(conn.Read(buf[:])) - - if !bytes.Equal(payload, buf[:]) { + if !CheckClientServer(clientData, serverData, socksPort) { t.Fail() } } @@ -179,12 +199,6 @@ shadowsocks: mux: enabled: true `, socksPort, serverPort) - go func() { - proxy, err := proxy.NewProxyFromConfigData([]byte(clientData), false) - common.Must(err) - common.Must(proxy.Run()) - }() - serverData := fmt.Sprintf(` run-type: server local-addr: 127.0.0.1 @@ -204,25 +218,8 @@ shadowsocks: method: AEAD_CHACHA20_POLY1305 password: 12345678 `, serverPort, util.HTTPPort) - go func() { - proxy, err := proxy.NewProxyFromConfigData([]byte(serverData), false) - common.Must(err) - common.Must(proxy.Run()) - }() - time.Sleep(time.Second * 2) - dialer, err := netproxy.SOCKS5("tcp", fmt.Sprintf("127.0.0.1:%d", socksPort), nil, netproxy.Direct) - - payload := util.GeneratePayload(1024) - buf := [1024]byte{} - - conn, err := dialer.Dial("tcp", util.EchoAddr) - common.Must(err) - - common.Must2(conn.Write(payload)) - common.Must2(conn.Read(buf[:])) - - if !bytes.Equal(payload, buf[:]) { + if !CheckClientServer(clientData, serverData, socksPort) { t.Fail() } } diff --git a/tunnel/dokodemo/conn.go b/tunnel/dokodemo/conn.go index c01605f..13ebc75 100644 --- a/tunnel/dokodemo/conn.go +++ b/tunnel/dokodemo/conn.go @@ -2,9 +2,10 @@ package dokodemo import ( "context" - "github.com/p4gefau1t/trojan-go/tunnel" "io" "net" + + "github.com/p4gefau1t/trojan-go/tunnel" ) const MaxPacketSize = 1024 * 8 @@ -27,13 +28,13 @@ type PacketConn struct { Input chan []byte Output chan []byte Source net.Addr - context.Context - context.CancelFunc + Ctx context.Context + Cancel context.CancelFunc } func (c *PacketConn) Close() error { - c.CancelFunc() - return nil + c.Cancel() + return c.PacketConn.Close() } func (c *PacketConn) ReadFrom(p []byte) (int, net.Addr, error) { @@ -55,7 +56,7 @@ func (c *PacketConn) ReadWithMetadata(p []byte) (int, *tunnel.Metadata, error) { case payload := <-c.Input: n := copy(p, payload) return n, c.M, nil - case <-c.Done(): + case <-c.Ctx.Done(): return 0, nil, io.EOF } } @@ -63,7 +64,7 @@ func (c *PacketConn) ReadWithMetadata(p []byte) (int, *tunnel.Metadata, error) { func (c *PacketConn) WriteWithMetadata(p []byte, m *tunnel.Metadata) (int, error) { select { case c.Output <- p: - case <-c.Done(): + case <-c.Ctx.Done(): return 0, io.EOF } return len(p), nil diff --git a/tunnel/dokodemo/server.go b/tunnel/dokodemo/server.go index 2e6472d..1f8f50c 100644 --- a/tunnel/dokodemo/server.go +++ b/tunnel/dokodemo/server.go @@ -2,14 +2,14 @@ package dokodemo import ( "context" + "net" + "sync" + "time" + "github.com/p4gefau1t/trojan-go/common" "github.com/p4gefau1t/trojan-go/config" "github.com/p4gefau1t/trojan-go/log" "github.com/p4gefau1t/trojan-go/tunnel" - "io" - "net" - "sync" - "time" ) type Server struct { @@ -33,8 +33,11 @@ func (s *Server) dispatchLoop() { buf := make([]byte, MaxPacketSize) n, addr, err := s.udpListener.ReadFrom(buf) if err != nil { - s.cancel() - log.Debug(common.NewError("dokodemo udp read error, closing").Base(err)) + select { + case <-s.ctx.Done(): + default: + log.Fatal(common.NewError("dokodemo failed to read from udp socket").Base(err)) + } return } log.Debug("udp packet from", addr) @@ -51,8 +54,8 @@ func (s *Server) dispatchLoop() { M: fixedMetadata, Source: addr, PacketConn: s.udpListener, - Context: ctx, - CancelFunc: cancel, + Ctx: ctx, + Cancel: cancel, } s.mapping[addr.String()] = conn s.mappingLock.Unlock() @@ -88,7 +91,7 @@ func (s *Server) dispatchLoop() { func (s *Server) AcceptConn(tunnel.Tunnel) (tunnel.Conn, error) { conn, err := s.tcpListener.Accept() if err != nil { - return nil, err + log.Fatal(common.NewError("dokodemo failed to accept connection").Base(err)) } return &Conn{ Conn: conn, @@ -103,7 +106,7 @@ func (s *Server) AcceptPacket(tunnel.Tunnel) (tunnel.PacketConn, error) { case conn := <-s.packetChan: return conn, nil case <-s.ctx.Done(): - return nil, io.EOF + return nil, common.NewError("dokodemo server closed") } } diff --git a/tunnel/mux/client.go b/tunnel/mux/client.go index d87d412..904bbea 100644 --- a/tunnel/mux/client.go +++ b/tunnel/mux/client.go @@ -49,20 +49,23 @@ func (c *Client) Close() error { return nil } -func (c *Client) cleanWorker() { +func (c *Client) cleanLoop() { var checkDuration time.Duration if c.timeout <= 0 { checkDuration = time.Second * 10 - log.Warn("invalid mux timeout") + log.Warn("negative mux timeout") } else { checkDuration = c.timeout / 4 } + log.Debug("check duration:", checkDuration.Seconds(), "s") for { select { case <-time.After(checkDuration): c.clientPoolLock.Lock() for id, info := range c.clientPool { if info.client.IsClosed() { + info.client.Close() + info.underlayConn.Close() delete(c.clientPool, id) log.Info("mux client", id, "is dead") } else if info.client.NumStreams() == 0 && time.Now().Sub(info.lastActiveTime) > c.timeout { @@ -72,16 +75,18 @@ func (c *Client) cleanWorker() { log.Info("mux client", id, "is closed due to inactivity") } } - for id, info := range c.clientPool { - log.Debug(fmt.Sprintf(" %x: %d/%d", id, info.client.NumStreams(), c.concurrency)) - } log.Debug("current mux clients: ", len(c.clientPool)) + for id, info := range c.clientPool { + log.Debug(fmt.Sprintf(" - %x: %d/%d", id, info.client.NumStreams(), c.concurrency)) + } c.clientPoolLock.Unlock() case <-c.ctx.Done(): log.Debug("shutting down mux cleaner..") c.clientPoolLock.Lock() for id, info := range c.clientPool { info.client.Close() + info.underlayConn.Close() + delete(c.clientPool, id) log.Debug("mux client", id, "closed") } c.clientPoolLock.Unlock() @@ -121,16 +126,13 @@ func (c *Client) newMuxClient() (*smuxClientInfo, error) { } func (c *Client) DialConn(addr *tunnel.Address, _ tunnel.Tunnel) (tunnel.Conn, error) { - c.clientPoolLock.Lock() - defer c.clientPoolLock.Unlock() createNewConn := func(info *smuxClientInfo) (tunnel.Conn, error) { - info.lastActiveTime = time.Now() rwc, err := info.client.Open() info.lastActiveTime = time.Now() if err != nil { - c.clientPoolLock.Lock() - defer c.clientPoolLock.Unlock() + info.underlayConn.Close() + info.client.Close() delete(c.clientPool, info.id) return nil, common.NewError("mux failed to open stream from client").Base(err) } @@ -140,6 +142,8 @@ func (c *Client) DialConn(addr *tunnel.Address, _ tunnel.Tunnel) (tunnel.Conn, e }, nil } + c.clientPoolLock.Lock() + defer c.clientPoolLock.Unlock() for _, info := range c.clientPool { if info.client.IsClosed() { delete(c.clientPool, info.id) @@ -173,7 +177,7 @@ func NewClient(ctx context.Context, underlay tunnel.Client) (*Client, error) { cancel: cancel, clientPool: make(map[muxID]*smuxClientInfo), } - go client.cleanWorker() + go client.cleanLoop() log.Debug("mux client created") return client, nil } diff --git a/tunnel/mux/server.go b/tunnel/mux/server.go index 871b1cf..60f5800 100644 --- a/tunnel/mux/server.go +++ b/tunnel/mux/server.go @@ -13,7 +13,6 @@ import ( type Server struct { underlay tunnel.Server connChan chan tunnel.Conn - errChan chan error ctx context.Context cancel context.CancelFunc } @@ -30,29 +29,36 @@ func (s *Server) acceptConnWorker() { } continue } - smuxConfig := smux.DefaultConfig() - //smuxConfig.KeepAliveDisabled = true - smuxSession, err := smux.Server(conn, smuxConfig) - if err != nil { - s.errChan <- err - continue - } - // TODO context - go func(session *smux.Session, conn tunnel.Conn) { - defer session.Close() - defer conn.Close() - for { - stream, err := session.AcceptStream() - if err != nil { - s.errChan <- err - return - } - s.connChan <- &Conn{ - rwc: stream, - Conn: conn, - } + go func(conn tunnel.Conn) { + smuxConfig := smux.DefaultConfig() + //smuxConfig.KeepAliveDisabled = true + smuxSession, err := smux.Server(conn, smuxConfig) + if err != nil { + log.Error(err) + return } - }(smuxSession, conn) + // TODO context + go func(session *smux.Session, conn tunnel.Conn) { + defer session.Close() + defer conn.Close() + for { + stream, err := session.AcceptStream() + if err != nil { + log.Error(err) + return + } + select { + case s.connChan <- &Conn{ + rwc: stream, + Conn: conn, + }: + case <-s.ctx.Done(): + log.Debug("exiting") + return + } + } + }(smuxSession, conn) + }(conn) } } @@ -60,10 +66,8 @@ func (s *Server) AcceptConn(tunnel.Tunnel) (tunnel.Conn, error) { select { case conn := <-s.connChan: return conn, nil - case err := <-s.errChan: - return nil, err case <-s.ctx.Done(): - return nil, common.NewError("mux client closed") + return nil, common.NewError("mux server closed") } } diff --git a/tunnel/router/client.go b/tunnel/router/client.go index 49fbcc6..8291acc 100644 --- a/tunnel/router/client.go +++ b/tunnel/router/client.go @@ -122,8 +122,8 @@ type Client struct { defaultPolicy int domainStrategy int underlay tunnel.Client - context.Context - context.CancelFunc + ctx context.Context + cancel context.CancelFunc } func (c *Client) Route(address *tunnel.Address) int { @@ -198,19 +198,19 @@ func (c *Client) DialPacket(overlay tunnel.Tunnel) (tunnel.PacketConn, error) { if err != nil { return nil, common.NewError("router failed to dial udp (proxy)").Base(err) } - ctx, cancel := context.WithCancel(c.Context) + ctx, cancel := context.WithCancel(c.ctx) return &PacketConn{ Client: c, PacketConn: direct, proxy: proxy, - CancelFunc: cancel, - Context: ctx, + cancel: cancel, + ctx: ctx, packetChan: make(chan *packetInfo, 16), }, nil } func (c *Client) Close() error { - c.CancelFunc() + c.cancel() return c.underlay.Close() } @@ -252,11 +252,11 @@ func NewClient(ctx context.Context, underlay tunnel.Client) (*Client, error) { cfg := config.FromContext(ctx, Name).(*Config) ctx, cancel := context.WithCancel(ctx) client := &Client{ - domains: [3][]*v2router.Domain{}, - cidrs: [3][]*v2router.CIDR{}, - underlay: underlay, - Context: ctx, - CancelFunc: cancel, + domains: [3][]*v2router.Domain{}, + cidrs: [3][]*v2router.CIDR{}, + underlay: underlay, + ctx: ctx, + cancel: cancel, } switch cfg.Router.DomainStrategy { case "as_is": diff --git a/tunnel/router/conn.go b/tunnel/router/conn.go index af9caf5..dc5d29b 100644 --- a/tunnel/router/conn.go +++ b/tunnel/router/conn.go @@ -19,8 +19,8 @@ type PacketConn struct { net.PacketConn packetChan chan *packetInfo *Client - context.Context - context.CancelFunc + ctx context.Context + cancel context.CancelFunc } func (c *PacketConn) packetLoop() { @@ -30,7 +30,7 @@ func (c *PacketConn) packetLoop() { n, addr, err := c.proxy.ReadWithMetadata(buf) if err != nil { select { - case <-c.Done(): + case <-c.ctx.Done(): return default: log.Error("router packetConn error", err) @@ -48,7 +48,7 @@ func (c *PacketConn) packetLoop() { n, addr, err := c.PacketConn.ReadFrom(buf) if err != nil { select { - case <-c.Done(): + case <-c.ctx.Done(): return default: log.Error("router packetConn error", err) @@ -66,7 +66,7 @@ func (c *PacketConn) packetLoop() { } func (c *PacketConn) Close() error { - c.CancelFunc() + c.cancel() c.proxy.Close() return c.PacketConn.Close() } @@ -105,7 +105,7 @@ func (c *PacketConn) ReadWithMetadata(p []byte) (int, *tunnel.Metadata, error) { case info := <-c.packetChan: n := copy(p, info.payload) return n, info.src, nil - case <-c.Done(): + case <-c.ctx.Done(): return 0, nil, io.EOF } } diff --git a/tunnel/shadowsocks/conn.go b/tunnel/shadowsocks/conn.go index f362e89..3f280eb 100644 --- a/tunnel/shadowsocks/conn.go +++ b/tunnel/shadowsocks/conn.go @@ -19,6 +19,7 @@ func (c *Conn) Write(p []byte) (n int, err error) { } func (c *Conn) Close() error { + c.Conn.Close() return c.aeadConn.Close() } diff --git a/tunnel/simplesocks/server.go b/tunnel/simplesocks/server.go index e0b8812..bafe898 100644 --- a/tunnel/simplesocks/server.go +++ b/tunnel/simplesocks/server.go @@ -15,11 +15,12 @@ type Server struct { underlay tunnel.Server connChan chan tunnel.Conn packetChan chan tunnel.PacketConn - errChan chan error ctx context.Context + cancel context.CancelFunc } func (s *Server) Close() error { + s.cancel() return s.underlay.Close() } @@ -37,7 +38,7 @@ func (s *Server) acceptLoop() { } metadata := new(tunnel.Metadata) if err := metadata.ReadFrom(conn); err != nil { - s.errChan <- common.NewError("simplesocks server faield to read header").Base(err) + log.Error(common.NewError("simplesocks server faield to read header").Base(err)) conn.Close() continue } @@ -54,7 +55,7 @@ func (s *Server) acceptLoop() { }, } default: - s.errChan <- common.NewError(fmt.Sprintf("simplesocks unknown command %d", metadata.Command)) + log.Error(common.NewError(fmt.Sprintf("simplesocks unknown command %d", metadata.Command))) conn.Close() } } @@ -64,8 +65,6 @@ func (s *Server) AcceptConn(tunnel.Tunnel) (tunnel.Conn, error) { select { case conn := <-s.connChan: return conn, nil - case err := <-s.errChan: - return nil, err case <-s.ctx.Done(): return nil, common.NewError("simplesocks server closed") } @@ -81,12 +80,13 @@ func (s *Server) AcceptPacket(tunnel.Tunnel) (tunnel.PacketConn, error) { } func NewServer(ctx context.Context, underlay tunnel.Server) (*Server, error) { + ctx, cancel := context.WithCancel(ctx) server := &Server{ underlay: underlay, ctx: ctx, connChan: make(chan tunnel.Conn, 32), packetChan: make(chan tunnel.PacketConn, 32), - errChan: make(chan error, 32), + cancel: cancel, } go server.acceptLoop() log.Debug("simplesocks server created") diff --git a/tunnel/socks/config.go b/tunnel/socks/config.go index ef489c9..d90a590 100644 --- a/tunnel/socks/config.go +++ b/tunnel/socks/config.go @@ -3,12 +3,15 @@ package socks import "github.com/p4gefau1t/trojan-go/config" type Config struct { - LocalHost string `json:"local_addr" yaml:"local-addr"` - LocalPort int `json:"local_port" yaml:"local-port"` + LocalHost string `json:"local_addr" yaml:"local-addr"` + LocalPort int `json:"local_port" yaml:"local-port"` + UDPTimeout int `json:"udp_timeout" yaml:"udp-timeout"` } func init() { config.RegisterConfigCreator(Name, func() interface{} { - return new(Config) + return &Config{ + UDPTimeout: 30, + } }) } diff --git a/tunnel/socks/conn.go b/tunnel/socks/conn.go index 839b6df..3c854d5 100644 --- a/tunnel/socks/conn.go +++ b/tunnel/socks/conn.go @@ -21,13 +21,11 @@ func (c *Conn) Metadata() *tunnel.Metadata { type PacketConn struct { net.PacketConn - srcAddr net.Addr - timeout time.Duration - shutdownChan chan struct{} + srcAddr net.Addr + timeout time.Duration } func (c *PacketConn) Close() error { - c.shutdownChan <- struct{}{} return c.PacketConn.Close() } @@ -72,11 +70,10 @@ func (c *PacketConn) ReadWithMetadata(payload []byte) (int, *tunnel.Metadata, er }, nil } -func NewPacketConn(packet net.PacketConn) *PacketConn { +func NewPacketConn(packet net.PacketConn, timeout time.Duration) *PacketConn { conn := &PacketConn{ - PacketConn: packet, - timeout: time.Second * 10, - shutdownChan: make(chan struct{}), + PacketConn: packet, + timeout: timeout, } return conn } diff --git a/tunnel/socks/server.go b/tunnel/socks/server.go index c9dad07..69b3553 100644 --- a/tunnel/socks/server.go +++ b/tunnel/socks/server.go @@ -7,6 +7,7 @@ import ( "io" "io/ioutil" "net" + "time" "github.com/p4gefau1t/trojan-go/common" "github.com/p4gefau1t/trojan-go/config" @@ -23,13 +24,15 @@ const ( MaxPacketSize = 1024 * 8 ) -// Server is a socks4/5 server +// Server is a socks5 server type Server struct { connChan chan tunnel.Conn packetChan chan tunnel.PacketConn tcpListener net.Listener - ctx context.Context localHost string + timeout time.Duration + ctx context.Context + cancel context.CancelFunc } func (s *Server) AcceptConn(tunnel.Tunnel) (tunnel.Conn, error) { @@ -51,6 +54,7 @@ func (s *Server) AcceptPacket(tunnel.Tunnel) (tunnel.PacketConn, error) { } func (s *Server) Close() error { + s.cancel() return s.tcpListener.Close() } @@ -136,7 +140,7 @@ func (s *Server) acceptLoop() { log.Error(common.NewError("socks5 failed to bind udp").Base(err)) return } - s.packetChan <- NewPacketConn(l) + s.packetChan <- NewPacketConn(l, s.timeout) log.Info("socks5 udp session") if err := s.associate(newConn, associateAddr); err != nil { log.Error(common.NewError("socks5 failed to respond to associate request").Base(err)) @@ -161,14 +165,17 @@ func NewServer(ctx context.Context, underlay tunnel.Server) (tunnel.Server, erro if err != nil { return nil, common.NewError("socks5 failed to listen").Base(err) } - log.Info("socks5 server is listening on tcp:", l.Addr().String()) + ctx, cancel := context.WithCancel(ctx) server := &Server{ tcpListener: l, ctx: ctx, + cancel: cancel, connChan: make(chan tunnel.Conn, 32), packetChan: make(chan tunnel.PacketConn, 32), + timeout: time.Duration(cfg.UDPTimeout) * time.Second, } go server.acceptLoop() + log.Info("socks5 server is listening on tcp:", l.Addr().String()) log.Debug("socks server created") return server, nil } diff --git a/tunnel/tproxy/server.go b/tunnel/tproxy/server.go index c827487..1153826 100644 --- a/tunnel/tproxy/server.go +++ b/tunnel/tproxy/server.go @@ -38,7 +38,12 @@ func (s *Server) Close() error { func (s *Server) AcceptConn(tunnel.Tunnel) (tunnel.Conn, error) { conn, err := s.tcpListener.Accept() if err != nil { - return nil, common.NewError("tproxy failed to accept connection").Base(err) + select { + case <-s.ctx.Done(): + default: + log.Fatal(common.NewError("tproxy failed to accept connection").Base(err)) + } + return nil, common.NewError("tproxy failed to accept conn") } addr, err := getOriginalTCPDest(conn.(*tproxy.Conn).TCPConn) if err != nil { @@ -59,8 +64,12 @@ func (s *Server) packetDispatchLoop() { buf := make([]byte, MaxPacketSize) n, src, dst, err := tproxy.ReadFromUDP(s.udpListener, buf) if err != nil { - s.cancel() - log.Error("tproxy failed to read from udp") + select { + case <-s.ctx.Done(): + default: + log.Fatal("tproxy failed to read from udp") + } + s.Close() return } log.Debug("udp packet from", src, "to", dst) @@ -78,8 +87,8 @@ func (s *Server) packetDispatchLoop() { Output: make(chan []byte, 16), Source: src, PacketConn: s.udpListener, - Context: ctx, - CancelFunc: cancel, + Ctx: ctx, + Cancel: cancel, M: &tunnel.Metadata{ Address: &tunnel.Address{}, }, diff --git a/tunnel/tproxy/tproxy_test.go b/tunnel/tproxy/tproxy_test.go new file mode 100644 index 0000000..8cd1fdb --- /dev/null +++ b/tunnel/tproxy/tproxy_test.go @@ -0,0 +1,30 @@ +package tproxy + +import ( + "context" + "fmt" + "github.com/p4gefau1t/trojan-go/common" + "github.com/p4gefau1t/trojan-go/config" + "os" + "testing" +) + +func TestTProxy(t *testing.T) { + if os.Getuid() != 0 { + t.Skip() + } + port := common.PickPort("tcp", "127.0.0.1") + cfg := &Config{ + LocalHost: "127.0.0.1", + LocalPort: port, + UDPTimeout: 0, + } + ctx := config.WithConfig(context.Background(), Name, cfg) + s, err := NewServer(ctx, nil) + common.Must(err) + go func() { + conn, err := s.AcceptConn(nil) + common.Must(err) + fmt.Println(conn.Metadata()) + }() +} diff --git a/tunnel/transport/server.go b/tunnel/transport/server.go index 31b2453..f9e9199 100644 --- a/tunnel/transport/server.go +++ b/tunnel/transport/server.go @@ -6,7 +6,6 @@ import ( "crypto/tls" "crypto/x509" "encoding/pem" - "github.com/p4gefau1t/trojan-go/tunnel/websocket" "io" "io/ioutil" "net" @@ -16,6 +15,8 @@ import ( "strconv" "strings" + "github.com/p4gefau1t/trojan-go/tunnel/websocket" + "github.com/p4gefau1t/trojan-go/common" "github.com/p4gefau1t/trojan-go/config" "github.com/p4gefau1t/trojan-go/log" @@ -39,10 +40,10 @@ type Server struct { sessionTicket bool curve []tls.CurveID keyLogger io.WriteCloser - redir *redirector.Redirector connChan chan tunnel.Conn wsChan chan tunnel.Conn plugin bool + redir *redirector.Redirector cmd *exec.Cmd ctx context.Context cancel context.CancelFunc @@ -63,8 +64,12 @@ func (s *Server) acceptLoop() { for { tcpConn, err := s.tcpListener.Accept() if err != nil { - s.cancel() - log.Error(common.NewError("transport accept error")) + select { + case <-s.ctx.Done(): + return + default: + log.Fatal(common.NewError("transport accept error")) + } return } go func(tcpConn net.Conn) { @@ -161,7 +166,7 @@ func (s *Server) AcceptConn(overlay tunnel.Tunnel) (tunnel.Conn, error) { case conn := <-s.wsChan: return conn, nil case <-s.ctx.Done(): - return nil, io.EOF + return nil, common.NewError("transport server closed") } } // trojan overlay @@ -169,7 +174,7 @@ func (s *Server) AcceptConn(overlay tunnel.Tunnel) (tunnel.Conn, error) { case conn := <-s.connChan: return conn, nil case <-s.ctx.Done(): - return nil, io.EOF + return nil, common.NewError("transport server closed") } } diff --git a/tunnel/trojan/server.go b/tunnel/trojan/server.go index 7da9169..ec9c584 100644 --- a/tunnel/trojan/server.go +++ b/tunnel/trojan/server.go @@ -3,11 +3,12 @@ package trojan import ( "context" "fmt" + "io" + "net" + "github.com/p4gefau1t/trojan-go/api" "github.com/p4gefau1t/trojan-go/statistic/memory" "github.com/p4gefau1t/trojan-go/statistic/mysql" - "io" - "net" "github.com/p4gefau1t/trojan-go/common" "github.com/p4gefau1t/trojan-go/config" @@ -100,9 +101,11 @@ type Server struct { muxChan chan tunnel.Conn packetChan chan tunnel.PacketConn ctx context.Context + cancel context.CancelFunc } func (s *Server) Close() error { + s.cancel() return s.underlay.Close() } @@ -110,7 +113,7 @@ func (s *Server) acceptLoop() { for { conn, err := s.underlay.AcceptConn(&Tunnel{}) if err != nil { // Closing - log.Debug(err) + log.Error(err) select { case <-s.ctx.Done(): return @@ -201,14 +204,16 @@ func NewServer(ctx context.Context, underlay tunnel.Server) (tunnel.Server, erro return nil, common.NewError("failed to create authenticator").Base(err) } redirAddr := tunnel.NewAddressFromHostPort("tcp", cfg.RemoteHost, cfg.RemotePort) + ctx, cancel := context.WithCancel(ctx) s := &Server{ underlay: underlay, auth: auth, - ctx: ctx, redirAddr: redirAddr, connChan: make(chan tunnel.Conn, 32), muxChan: make(chan tunnel.Conn, 32), packetChan: make(chan tunnel.PacketConn, 32), + ctx: ctx, + cancel: cancel, } if !cfg.DisableHTTPCheck {