diff --git a/go.mod b/go.mod index 8593663..c5cda61 100644 --- a/go.mod +++ b/go.mod @@ -12,6 +12,7 @@ require ( github.com/mediocregopher/radix/v3 v3.5.0 github.com/onsi/ginkgo v1.10.1 // indirect github.com/onsi/gomega v1.7.0 // indirect + github.com/patrickmn/go-cache v2.1.0+incompatible github.com/proullon/ramsql v0.0.0-20181213202341-817cee58a244 github.com/refraction-networking/utls v0.0.0-20190909200633-43c36d3c1f57 github.com/smartystreets/goconvey v1.6.4 diff --git a/go.sum b/go.sum index cb53be9..9c7d97c 100644 --- a/go.sum +++ b/go.sum @@ -225,6 +225,8 @@ github.com/onsi/gomega v1.7.0/go.mod h1:ex+gbHU/CVuBBDIJjb2X0qEXbFg53c61hWP/1Cpa github.com/openzipkin/zipkin-go v0.1.6/go.mod h1:QgAqvLzwWbR/WpD4A3cGpPtJrZXNIiJc5AZX7/PBEpw= github.com/oracle/oci-go-sdk v7.0.0+incompatible/go.mod h1:VQb79nF8Z2cwLkLS35ukwStZIg5F66tcBccjip/j888= github.com/ovh/go-ovh v0.0.0-20181109152953-ba5adb4cf014/go.mod h1:joRatxRJaZBsY3JAOEMcoOp05CnZzsx4scTxi95DHyQ= +github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc= +github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ= github.com/pierrec/lz4 v2.0.5+incompatible/go.mod h1:pdkljMzZIN41W+lC3N2tnIh5sFi+IEE17M5jbnwPHcY= github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= diff --git a/protocol/direct/outbound.go b/protocol/direct/outbound.go index 66bca87..f5fa774 100644 --- a/protocol/direct/outbound.go +++ b/protocol/direct/outbound.go @@ -13,8 +13,11 @@ import ( "github.com/p4gefau1t/trojan-go/conf" "github.com/p4gefau1t/trojan-go/log" "github.com/p4gefau1t/trojan-go/protocol" + "github.com/patrickmn/go-cache" ) +var dnsCache = cache.New(5*time.Minute, 1*time.Minute) + type DirectOutboundConnSession struct { protocol.ConnSession conn io.ReadWriteCloser @@ -35,8 +38,19 @@ func (o *DirectOutboundConnSession) Close() error { func NewOutboundConnSession(ctx context.Context, req *protocol.Request, config *conf.GlobalConfig) (protocol.ConnSession, error) { var newConn net.Conn - //custom dns server - if req.AddressType == common.DomainName && len(config.DNS) != 0 { + var err error + //look up the domain name in cache first + if req.AddressType == common.DomainName && len(config.DNS) != 0 { //customized dns server + ip, found := dnsCache.Get(req.DomainName) + if found { + newConn, err = net.DialTCP("tcp", nil, &net.TCPAddr{ + IP: ip.(net.IP), + }) + if err != nil { + return nil, err + } + goto done + } //find a avaliable dns server for _, s := range config.DNS { var dnsType conf.DNSType @@ -86,6 +100,7 @@ func NewOutboundConnSession(ctx context.Context, req *protocol.Request, config * if err != nil { return nil, err } + dnsCache.Set(req.DomainName, net.ParseIP(ip), cache.DefaultExpiration) break } } else { @@ -129,6 +144,7 @@ func NewOutboundConnSession(ctx context.Context, req *protocol.Request, config * if err != nil { return nil, err } + dnsCache.Set(req.DomainName, ip.IP, cache.DefaultExpiration) break } } @@ -140,11 +156,12 @@ func NewOutboundConnSession(ctx context.Context, req *protocol.Request, config * } else { //default resolver var err error - newConn, err = net.Dial(req.Network(), req.String()) + newConn, err = net.Dial("tcp", req.String()) if err != nil { return nil, err } } +done: o := &DirectOutboundConnSession{ request: req, conn: newConn, diff --git a/protocol/direct/outbound_test.go b/protocol/direct/outbound_test.go index 428ae70..533ff67 100644 --- a/protocol/direct/outbound_test.go +++ b/protocol/direct/outbound_test.go @@ -42,7 +42,6 @@ func TestUDPDirectOutbound(t *testing.T) { func TestDNS(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) - go test.RunEchoTCPServer(ctx) config := &conf.GlobalConfig{ DNS: []string{"114.114.114.114:53"}, } @@ -67,7 +66,6 @@ func TestDNS(t *testing.T) { func TestDOT(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) - go test.RunEchoTCPServer(ctx) config := &conf.GlobalConfig{ DNS: []string{"dot://223.5.5.5:853"}, } @@ -92,7 +90,6 @@ func TestDOT(t *testing.T) { func TestDOH(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) - go test.RunEchoTCPServer(ctx) config := &conf.GlobalConfig{ DNS: []string{"https://223.5.5.5:443"}, } @@ -114,3 +111,25 @@ func TestDOH(t *testing.T) { fmt.Println(string(buf[:])) cancel() } + +func TestCache(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + config := &conf.GlobalConfig{ + DNS: []string{"223.5.5.5:53"}, + } + req := &protocol.Request{ + Address: &common.Address{ + DomainName: "www.baidu.com", + Port: 80, + AddressType: common.DomainName, + NetworkType: "tcp", + }, + } + conn, err := NewOutboundConnSession(ctx, req, config) + common.Must(err) + conn.Close() + conn, err = NewOutboundConnSession(ctx, req, config) + common.Must(err) + conn.Close() + cancel() +} diff --git a/protocol/trojan/outbound.go b/protocol/trojan/outbound.go index d7985d3..20aeeae 100644 --- a/protocol/trojan/outbound.go +++ b/protocol/trojan/outbound.go @@ -14,13 +14,14 @@ import ( type TrojanOutboundConnSession struct { protocol.ConnSession - config *conf.GlobalConfig - rwc io.ReadWriteCloser - request *protocol.Request - sent uint64 - recv uint64 - auth stat.Authenticator - meter stat.TrafficMeter + config *conf.GlobalConfig + rwc io.ReadWriteCloser + request *protocol.Request + sent uint64 + recv uint64 + auth stat.Authenticator + meter stat.TrafficMeter + trojanHeader []byte } func (o *TrojanOutboundConnSession) SetMeter(meter stat.TrafficMeter) { @@ -28,6 +29,19 @@ func (o *TrojanOutboundConnSession) SetMeter(meter stat.TrafficMeter) { } func (o *TrojanOutboundConnSession) Write(p []byte) (int, error) { + if o.trojanHeader != nil { + //glue the payload with the trojan request header + fullRequest := o.trojanHeader + fullRequest = append(fullRequest, p...) + n, err := o.rwc.Write(fullRequest) + if n >= len(o.trojanHeader) { + n -= len(o.trojanHeader) + } + o.meter.Count(n, 0) + o.sent += uint64(n) + o.trojanHeader = nil + return n, err + } n, err := o.rwc.Write(p) o.meter.Count(n, 0) o.sent += uint64(n) @@ -57,8 +71,8 @@ func (o *TrojanOutboundConnSession) writeRequest() error { buf.WriteByte(byte(o.request.Command)) protocol.WriteAddress(buf, o.request) buf.Write(crlf) - _, err := o.rwc.Write(buf.Bytes()) - return err + o.trojanHeader = buf.Bytes() + return nil } func NewOutboundConnSession(req *protocol.Request, rwc io.ReadWriteCloser, config *conf.GlobalConfig, auth stat.Authenticator) (protocol.ConnSession, error) {