fix goroutine leak and deadlock

This commit is contained in:
Page Fault
2020-06-12 06:02:06 +00:00
parent 2855a50f1d
commit c58251a099
22 changed files with 268 additions and 204 deletions
+2 -1
View File
@@ -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"
}
}
```
-1
View File
@@ -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"
)
+1
View File
@@ -4,4 +4,5 @@ package build
import (
_ "github.com/p4gefau1t/trojan-go/easy"
_ "github.com/p4gefau1t/trojan-go/log/golog"
)
+1 -27
View File
@@ -14,33 +14,7 @@ weight: 22
- Trojan-Go,可以从release页面下载
### 配置证书
为了伪装成一个正常的HTTPS站点,也为了保证传输的安全,我们需要一份经过权威证书机构签名的证书。Trojan-Go支持从Let's Encrypt自动申请证书。首先将你的域名正确解析到你的服务器IP。然后准备好一个邮箱地址,合乎邮箱地址规则即可,不需要真实邮箱地址。保证你的服务器443和80端口没有被其他程序(nginxapache,正在运行的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等机构免费申请签发
### 服务端配置
+3
View File
@@ -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{}) {
+41 -21
View File
@@ -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)
}
+44 -47
View File
@@ -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()
}
}
+8 -7
View File
@@ -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
+13 -10
View File
@@ -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")
}
}
+15 -11
View File
@@ -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
}
+30 -26
View File
@@ -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")
}
}
+11 -11
View File
@@ -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":
+6 -6
View File
@@ -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
}
}
+1
View File
@@ -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()
}
+6 -6
View File
@@ -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")
+6 -3
View File
@@ -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,
}
})
}
+5 -8
View File
@@ -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
}
+11 -4
View File
@@ -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
}
+14 -5
View File
@@ -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{},
},
+30
View File
@@ -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())
}()
}
+11 -6
View File
@@ -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")
}
}
+9 -4
View File
@@ -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 {