From 83071288446bedaa22eaa38a4f6eba420d36d9b6 Mon Sep 17 00:00:00 2001 From: wweir Date: Sat, 17 Jul 2021 08:05:02 +0800 Subject: [PATCH] refactor, start v0.8 --- .github/Dockerfile | 17 -- .gitignore | 8 +- Makefile | 26 ++- cmd/client/main.go | 219 +++++++++++++++++++++++ cmd/client/socks5.go | 196 ++++++++++++++++++++ cmd/server/main.go | 163 +++++++++++++++++ conf/conf.go | 196 -------------------- conf/conf_darwin.go | 100 ----------- conf/conf_linux.go | 105 ----------- conf/conf_windows.go | 188 ------------------- conf/sower.toml | 65 ------- go.mod | 12 +- go.sum | 75 ++++---- main.go | 45 ----- {dhcp => pkg/dhcp}/dhcp.go | 2 +- {dhcp => pkg/dhcp}/dhcp_test.go | 4 +- {dhcp => pkg/dhcp}/pick_iface_other.go | 0 {dhcp => pkg/dhcp}/pick_iface_test.go | 2 +- {dhcp => pkg/dhcp}/pick_iface_windows.go | 0 {dhcp => pkg/dhcp}/util.go | 0 pkg/mem/mem.go | 142 +++++++++++++++ pkg/teeconn/tee_conn.go | 56 ++++++ proxy/dns.go | 118 ------------ proxy/http_proxy.go | 84 --------- proxy/proxy.go | 140 --------------- proxy/socks5.go | 44 ----- proxy/util.go | 64 ------- router/dns.go | 82 +++++++++ router/http_ping.go | 124 ------------- router/http_ping_test.go | 72 -------- router/mmdb.go | 52 ++++++ router/ping.go | 27 +++ router/router.go | 211 +++++++++------------- transport/proxy_conn.go | 68 ------- transport/socks5_client.go | 124 ------------- transport/socks5_rfc.go | 115 ------------ transport/socks5_server.go | 74 -------- transport/sower/sower.go | 79 ++++++++ transport/transport.go | 12 ++ transport/transport_test.go | 46 +++++ transport/trojan.go | 122 ------------- transport/trojan/trojan.go | 182 +++++++++++++++++++ transport/util.go | 89 --------- util/tee_conn.go | 50 ------ util/util.go | 51 +++++- 45 files changed, 1461 insertions(+), 2190 deletions(-) delete mode 100644 .github/Dockerfile create mode 100644 cmd/client/main.go create mode 100644 cmd/client/socks5.go create mode 100644 cmd/server/main.go delete mode 100644 conf/conf.go delete mode 100644 conf/conf_darwin.go delete mode 100644 conf/conf_linux.go delete mode 100644 conf/conf_windows.go delete mode 100644 conf/sower.toml delete mode 100644 main.go rename {dhcp => pkg/dhcp}/dhcp.go (97%) rename {dhcp => pkg/dhcp}/dhcp_test.go (61%) rename {dhcp => pkg/dhcp}/pick_iface_other.go (100%) rename {dhcp => pkg/dhcp}/pick_iface_test.go (81%) rename {dhcp => pkg/dhcp}/pick_iface_windows.go (100%) rename {dhcp => pkg/dhcp}/util.go (100%) create mode 100644 pkg/mem/mem.go create mode 100644 pkg/teeconn/tee_conn.go delete mode 100644 proxy/dns.go delete mode 100644 proxy/http_proxy.go delete mode 100644 proxy/proxy.go delete mode 100644 proxy/socks5.go delete mode 100644 proxy/util.go create mode 100644 router/dns.go delete mode 100644 router/http_ping.go delete mode 100644 router/http_ping_test.go create mode 100644 router/mmdb.go create mode 100644 router/ping.go delete mode 100644 transport/proxy_conn.go delete mode 100644 transport/socks5_client.go delete mode 100644 transport/socks5_rfc.go delete mode 100644 transport/socks5_server.go create mode 100644 transport/sower/sower.go create mode 100644 transport/transport.go create mode 100644 transport/transport_test.go delete mode 100644 transport/trojan.go create mode 100644 transport/trojan/trojan.go delete mode 100644 transport/util.go delete mode 100644 util/tee_conn.go diff --git a/.github/Dockerfile b/.github/Dockerfile deleted file mode 100644 index 76ad4cf..0000000 --- a/.github/Dockerfile +++ /dev/null @@ -1,17 +0,0 @@ -# Compile -FROM golang:1.14-alpine AS compiler - -RUN apk add --no-cache git make - -# enable go modules -WORKDIR /src -COPY . . - -RUN CGO_ENABLED=0 make build - - -# Build image -FROM scratch - -COPY --from=compiler /src/sower /sower -ENTRYPOINT [ "/sower" ] \ No newline at end of file diff --git a/.gitignore b/.gitignore index d24cfc8..a022dbd 100644 --- a/.gitignore +++ b/.gitignore @@ -17,6 +17,8 @@ *.swo .vscode .idea -sower -sower.exe* -/sower.toml + +/cmd/client/client +/client +/cmd/server/server +/server \ No newline at end of file diff --git a/Makefile b/Makefile index da8f9ce..7096ad8 100644 --- a/Makefile +++ b/Makefile @@ -1,12 +1,22 @@ +CPUS ?= $(shell nproc) +MAKEFLAGS += --jobs=$(CPUS) +GO:=CGO_ENABLED=0 go + default: test build test: - go vet ./... - go list ./... | grep -v internal | xargs go test -build: - go build -ldflags "-w -s \ - -X conf.version=$(shell git describe --tags --always) \ - -X conf.date=$(shell date +%Y-%m-%d)" -image: - docker build -t sower -f .github/Dockerfile . + ${GO} vet ./... + ${GO} test ./... +build: client server + +.PHONY: client +client: + ${GO} build -ldflags "\ + -X main.version=$(shell git describe --tags --always) \ + -X main.date=$(shell date +%Y-%m-%d)" ./cmd/client +.PHONY: server +server: + ${GO} build -ldflags "\ + -X main.version=$(shell git describe --tags --always) \ + -X main.date=$(shell date +%Y-%m-%d)" ./cmd/server \ No newline at end of file diff --git a/cmd/client/main.go b/cmd/client/main.go new file mode 100644 index 0000000..569c71f --- /dev/null +++ b/cmd/client/main.go @@ -0,0 +1,219 @@ +package main + +import ( + "bufio" + "crypto/tls" + "net" + "net/http" + "os" + "strings" + "time" + + "github.com/cristalhq/aconfig" + "github.com/miekg/dns" + "github.com/pkg/errors" + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/rs/zerolog/pkgerrors" + "github.com/wweir/sower/pkg/teeconn" + "github.com/wweir/sower/router" + "github.com/wweir/sower/transport" + "github.com/wweir/sower/transport/sower" + "github.com/wweir/sower/transport/trojan" + "github.com/wweir/sower/util" +) + +var ( + version, date string + + conf = struct { + Proxy struct { + Type string `default:"sower" usage:"remote proxy protocol, sower/trojan"` + Addr string `required:"true" usage:"remote proxy address, eg: proxy.com"` + Password string `required:"true" usage:"remote proxy password"` + } + + Socks5Addr string `default:":1080" usage:"socks5 listen address"` + FallbackDNS string `default:"223.5.5.5" usage:"fallback dns server"` + DNSServeIP string `usage:"dns server ip, eg: 127.0.0.1"` + + Router struct { + ProxyList []string + ProxyRefs []string + DirectList []string + DirectRefs []string + BlockList []string + BlockRefs []string + } + }{} +) + +func init() { + zerolog.ErrorStackMarshaler = func(err error) interface{} { + return pkgerrors.MarshalStack(err) + } + log.Logger = zerolog.New(zerolog.ConsoleWriter{ + Out: os.Stderr, + TimeFormat: time.StampMilli, + FormatCaller: func(i interface{}) string { + caller := i.(string) + if idx := strings.Index(caller, "/pkg/mod/"); idx > 0 { + return caller[idx+9:] + } + if idx := strings.LastIndexByte(caller, '/'); idx > 0 { + return caller[idx+1:] + } + return caller + }, + }).With().Timestamp().Caller().Logger() + + if err := aconfig.LoaderFor(&conf, aconfig.Config{}).Load(); err != nil { + log.Fatal().Err(err). + Msg("Load config") + } + + log.Info(). + Str("version", version). + Str("date", date). + Interface("config", conf). + Msg("Starting") +} + +func main() { + + r := router.NewRouter(genProxyDial()) + + go func() { + lnHTTP, err := net.Listen("tcp", net.JoinHostPort(conf.DNSServeIP, "80")) + if err != nil { + log.Fatal().Err(err).Msg("listen port") + } + go ServeHTTP(lnHTTP, r) + + lnHTTPS, err := net.Listen("tcp", net.JoinHostPort(conf.DNSServeIP, "443")) + if err != nil { + log.Fatal().Err(err).Msg("listen port") + } + go ServeHTTPS(lnHTTPS, r) + + if err := dns.ListenAndServe(conf.DNSServeIP, "udp", r); err != nil { + log.Fatal().Err(err).Msg("serve dns") + } + }() + + ln, err := net.Listen("tcp", conf.Socks5Addr) + if err != nil { + log.Fatal().Err(err).Msg("listen port") + } + go ServeSocks5(ln, r) + + select {} +} + +func genProxyDial() func(network, host string, port uint16) (net.Conn, error) { + var ( + proxyAddr = net.JoinHostPort(conf.Proxy.Addr, "443") + tlsCfg = &tls.Config{} + proxy transport.Transport + ) + switch conf.Proxy.Type { + case "sower": + proxy = sower.New(conf.Proxy.Password) + case "trojan": + proxy = trojan.New(conf.Proxy.Password) + default: + log.Fatal(). + Str("type", conf.Proxy.Type). + Msg("unknown proxy type") + } + + return func(network, host string, port uint16) (net.Conn, error) { + if host == "" || port == 0 { + return nil, errors.Errorf("invalid addr(%s:%d)", host, port) + } + + c, err := tls.Dial("tcp", proxyAddr, tlsCfg) + if err != nil { + return nil, err + } + + if err := proxy.Wrap(c, host, port); err != nil { + return nil, err + } + + return c, nil + } +} + +func ServeHTTP(ln net.Listener, r *router.Router) { + conn, err := ln.Accept() + if err != nil { + log.Fatal().Err(err). + Msg("serve socks5") + } + + go ServeHTTP(ln, r) + start := time.Now() + teeconn := teeconn.New(conn) + defer teeconn.Close() + + req, err := http.ReadRequest(bufio.NewReader(teeconn)) + if err != nil { + log.Error().Err(err).Msg("read http request") + return + } + + rc, err := r.ProxyDial("tcp", req.Host, 80) + if err != nil { + log.Error().Err(err). + Str("host", req.Host). + Interface("req", req.URL). + Msg("dial proxy") + return + } + defer rc.Close() + + teeconn.Stop().Reread() + util.Relay(teeconn, rc) + log.Info(). + Str("host", req.Host). + Dur("spend", time.Since(start)). + Msg("serve http") +} + +func ServeHTTPS(ln net.Listener, r *router.Router) { + conn, err := ln.Accept() + if err != nil { + log.Fatal().Err(err). + Msg("serve socks5") + } + + go ServeHTTPS(ln, r) + start := time.Now() + teeconn := teeconn.New(conn) + defer teeconn.Close() + + var domain string + tls.Server(teeconn, &tls.Config{ + GetConfigForClient: func(hello *tls.ClientHelloInfo) (*tls.Config, error) { + domain = hello.ServerName + return nil, nil + }, + }).Handshake() + + rc, err := r.ProxyDial("tcp", domain, 443) + if err != nil { + log.Error().Err(err). + Str("host", domain). + Msg("dial proxy") + return + } + defer rc.Close() + + teeconn.Stop().Reread() + util.Relay(teeconn, rc) + log.Info(). + Str("host", domain). + Dur("spend", time.Since(start)). + Msg("serve http") +} diff --git a/cmd/client/socks5.go b/cmd/client/socks5.go new file mode 100644 index 0000000..ee9ee6f --- /dev/null +++ b/cmd/client/socks5.go @@ -0,0 +1,196 @@ +package main + +import ( + "encoding/binary" + "io" + "net" + "time" + + "github.com/rs/zerolog/log" + "github.com/wweir/sower/router" +) + +func ServeSocks5(ln net.Listener, r *router.Router) { + conn, err := ln.Accept() + if err != nil { + log.Fatal().Err(err). + Msg("serve socks5") + } + go ServeSocks5(ln, r) + defer conn.Close() + start := time.Now() + + { + auth := new(socks5AuthReq) + if err := auth.Fulfill(conn); err != nil { + log.Error().Err(err). + Interface("request", auth). + Msg("socks5 auth") + return + } + + if err := binary.Write(conn, binary.BigEndian, socks5AuthResp); err != nil { + log.Error().Err(err). + Msg("socks5 auth") + return + } + } + + var addr addrType + { + head := new(socks5HeadReq) + if err := binary.Read(conn, binary.BigEndian, head); err != nil || !head.IsValid() { + return + } + switch head.ATYP { + case 0x01: // IPv4 + addr = &addrTypeIPv4{} + case 0x03: // domain name + addr = &addrTypeDomain{} + case 0x04: // IPv6 + addr = &addrTypeIPv6{} + default: + log.Error().Err(err). + Interface("head", head). + Msg("socks5 connect") + return + } + + if err := addr.Fulfill(conn); err != nil { + return + } + + if err := binary.Write(conn, binary.BigEndian, socks5HeadResp); err != nil { + log.Error().Err(err). + Msg("socks5 head") + return + } + } + + host, port := addr.Addr() + log.Err(r.RouteHandle(conn, host, port)). + Str("host", host). + Uint16("port", port). + Dur("spend", time.Since(start)). + Msg("serve socsk5") +} + +/******************* https://tools.ietf.org/html/rfc1928 *******************/ + +// 1. client send auth request +type socks5AuthReq struct { + VER byte + NMETHODS uint8 + METHODS []byte +} + +func (req *socks5AuthReq) Fulfill(r io.Reader) error { + buf := make([]byte, 2) + if n, err := r.Read(buf); err != nil || n != 2 { + return err + } + + req.VER = buf[0] + req.NMETHODS = buf[1] + + req.METHODS = make([]byte, int(req.NMETHODS)) + if n, err := r.Read(req.METHODS); err != nil || n != len(req.METHODS) { + return err + } + + return nil +} + +func (r *socks5AuthReq) IsValid() bool { + return r.VER == 5 && r.METHODS[0] == 0 +} + +// 2. server response auth request +var socks5AuthResp = struct { + VER byte + METHOD byte +}{VER: 5, METHOD: 0} + +// 3. client request with target address +type socks5HeadReq struct { + VER byte + CMD byte + RSV byte + ATYP byte +} + +func (r *socks5HeadReq) IsValid() bool { + return r.VER == 5 && r.CMD == 1 +} + +// 4. server response with the address that assigned to connect to target address +var socks5HeadResp = struct { + VER byte + REP byte + RSV byte + ATYP byte + BIND struct { + ADDR [4]byte + PORT uint16 + } +}{VER: 5, REP: 0, RSV: 0, ATYP: 1} + +type addrType interface { + Fulfill(r io.Reader) error + Addr() (domain string, port uint16) +} + +// ATYP: +// 0x01 -> net.IPv4len +// 0x03 -> first byte is length +// 0x04 -> net.IPv6len +type addrTypeIPv4 struct { + DST_ADDR [4]byte + DST_PORT uint16 +} + +func (a *addrTypeIPv4) Fulfill(r io.Reader) error { + return binary.Read(r, binary.BigEndian, &a) +} +func (a *addrTypeIPv4) Addr() (string, uint16) { + return net.IP(a.DST_ADDR[:]).String(), a.DST_PORT +} + +type addrTypeIPv6 struct { + DST_ADDR [16]byte + DST_PORT uint16 +} + +func (a *addrTypeIPv6) Fulfill(r io.Reader) error { + return binary.Read(r, binary.BigEndian, &a) +} +func (a *addrTypeIPv6) Addr() (string, uint16) { + return net.IP(a.DST_ADDR[:]).String(), a.DST_PORT +} + +type addrTypeDomain struct { + DST_ADDR_LEN uint8 + DST_ADDR []byte + DST_PORT uint16 +} + +func (a *addrTypeDomain) Fulfill(r io.Reader) error { + buf := make([]byte, 1) + if _, err := io.ReadFull(r, buf); err != nil { + return err + } + + a.DST_ADDR_LEN = uint8(buf[0]) + buf = make([]byte, a.DST_ADDR_LEN+2) + if _, err := io.ReadFull(r, buf); err != nil { + return err + } + + a.DST_ADDR = buf[:int(a.DST_ADDR_LEN)] + a.DST_PORT = binary.BigEndian.Uint16(buf[int(a.DST_ADDR_LEN):]) + + return nil +} +func (a *addrTypeDomain) Addr() (string, uint16) { + return string(a.DST_ADDR[:]), a.DST_PORT +} diff --git a/cmd/server/main.go b/cmd/server/main.go new file mode 100644 index 0000000..adedd26 --- /dev/null +++ b/cmd/server/main.go @@ -0,0 +1,163 @@ +package main + +import ( + "crypto/tls" + "net" + "net/http" + "os" + "path/filepath" + "strings" + "time" + + "github.com/cristalhq/aconfig" + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/rs/zerolog/pkgerrors" + "github.com/wweir/sower/pkg/teeconn" + "github.com/wweir/sower/transport/sower" + "github.com/wweir/sower/transport/trojan" + "github.com/wweir/sower/util" + "golang.org/x/crypto/acme/autocert" +) + +var ( + version, date string + + conf = struct { + ServeIP string `usage:"listen to port 80 443 of this IP, eg: 0.0.0.0"` + Password string `required:"true"` + FakeSite string `required:"true" default:"127.0.0.1:8080" usage:"fake site address"` + + Cert struct { + Email string + Cert string + Key string + } + }{} +) + +func init() { + zerolog.ErrorStackMarshaler = func(err error) interface{} { + return pkgerrors.MarshalStack(err) + } + log.Logger = zerolog.New(zerolog.ConsoleWriter{ + Out: os.Stderr, + TimeFormat: time.StampMilli, + FormatCaller: func(i interface{}) string { + caller := i.(string) + if idx := strings.Index(caller, "/pkg/mod/"); idx > 0 { + return caller[idx+9:] + } + if idx := strings.LastIndexByte(caller, '/'); idx > 0 { + return caller[idx+1:] + } + return caller + }, + }).With().Timestamp().Caller().Logger() + + if err := aconfig.LoaderFor(&conf, aconfig.Config{}).Load(); err != nil { + log.Fatal().Err(err).Msg("Load config") + } + + log.Info(). + Str("version", version). + Str("date", date). + Interface("config", conf). + Msg("Starting") +} + +func main() { + cacheDir, _ := os.UserCacheDir() + cacheDir = filepath.Join(cacheDir, "sower") + if err := os.MkdirAll(cacheDir, 0600); err != nil { + log.Fatal().Err(err). + Str("dir", cacheDir). + Msg("make cache dir") + } + + certManager := autocert.Manager{ + Prompt: autocert.AcceptTOS, + Email: conf.Cert.Email, + Cache: autocert.DirCache(cacheDir), + } + + tlsConf := &tls.Config{ + GetCertificate: certManager.GetCertificate, + MinVersion: tls.VersionTLS12, + NextProtos: []string{"http/1.1", "h2"}, + } + if conf.Cert.Cert != "" || conf.Cert.Key != "" { + cert, err := tls.LoadX509KeyPair(conf.Cert.Cert, conf.Cert.Key) + if err != nil { + log.Fatal().Err(err).Msg("load certificate") + } + + tlsConf.GetCertificate = nil + tlsConf.Certificates = []tls.Certificate{cert} + } + + // Redirect 80 to 443 + go http.ListenAndServe(net.JoinHostPort(conf.ServeIP, "80"), + certManager.HTTPHandler(http.HandlerFunc(redirectToHTTPS))) + + ln, err := tls.Listen("tcp", net.JoinHostPort(conf.ServeIP, "443"), tlsConf) + if err != nil { + log.Fatal().Err(err).Msg("listen tcp 443") + } + + go serve443(ln, conf.FakeSite, sower.New(conf.Password), trojan.New(conf.Password)) + select {} +} + +func redirectToHTTPS(w http.ResponseWriter, r *http.Request) { + r.URL.Scheme = "https" + if host, _, err := net.SplitHostPort(r.Host); err != nil { + r.URL.Host = r.Host + } else { + r.URL.Host = host + } + + http.Redirect(w, r, r.URL.String(), 301) +} + +func serve443(ln net.Listener, fakeSite string, sower *sower.Sower, trojan *trojan.Trojan) { + conn, err := ln.Accept() + if err != nil { + log.Fatal().Err(err).Msg("serve 443 port") + } + + go serve443(ln, fakeSite, sower, trojan) + teeconn := teeconn.New(conn) + defer teeconn.Close() + + teeconn.Reread() + if addr := sower.Unwrap(teeconn); addr != nil { + teeconn.Stop() + + dur, err := util.RelayTo(teeconn, addr.String()) + log.Err(err). + Dur("spend", dur). + Str("target", addr.String()). + Msg("relay sower conn") + return + } + + teeconn.Reread() + if addr := trojan.Unwrap(teeconn); addr != nil { + teeconn.Stop() + + dur, err := util.RelayTo(teeconn, addr.String()) + log.Err(err). + Dur("spend", dur). + Str("target", addr.String()). + Msg("relay trojan conn") + return + } + + teeconn.Stop().Reread() + dur, err := util.RelayTo(teeconn, fakeSite) + log.Err(err). + Dur("spend", dur). + Str("target", fakeSite). + Msg("relay fake site") +} diff --git a/conf/conf.go b/conf/conf.go deleted file mode 100644 index 488b242..0000000 --- a/conf/conf.go +++ /dev/null @@ -1,196 +0,0 @@ -package conf - -import ( - "bufio" - "flag" - "io" - "net/http" - "os" - "path/filepath" - "strings" - "sync" - - toml "github.com/pelletier/go-toml" - "github.com/wweir/sower/util" - "github.com/wweir/util-go/log" - "golang.org/x/xerrors" -) - -type Client struct { - Address string `toml:"address"` - DNSUpstream string `toml:"dns_upstream"` - Socks5Proxy string `toml:"socks5"` - HTTPProxy string `toml:"http_proxy"` - PortForward map[string]string `toml:"port_forward"` - - Router struct { - DetectLevel int `toml:"detect_level"` - ProxyList []string `toml:"proxy_list"` - ProxyRefs []string `toml:"proxy_refs"` - DirectList []string `toml:"direct_list"` - DirectRefs []string `toml:"direct_refs"` - BlockList []string `toml:"block_list"` - BlockRefs []string `toml:"block_refs"` - } `toml:"router"` -} -type Server struct { - Upstream string `toml:"upstream"` - CertFile string `toml:"cert_file"` - KeyFile string `toml:"key_file"` - CertEmail string `toml:"cert_email"` -} - -var ( - version, date string - - execFile, _ = os.Executable() - execDir, _ = filepath.Abs(filepath.Dir(execFile)) - // Conf full config, include common and server / client - conf = struct { - file string - Password string `toml:"password"` - Client Client `toml:"client"` - Server Server `toml:"server"` - }{} -) - -func Init() (*Client, *Server, string) { - beforeInitFlag() - defer afterInitFlag() - flag.StringVar(&conf.Password, "password", "", "password") - flag.StringVar(&conf.Server.Upstream, "s", "", "upstream http service, eg: 127.0.0.1:8080") - flag.StringVar(&conf.Server.CertFile, "s_cert", "", "tls cert file, gen cert from letsencrypt if empty") - flag.StringVar(&conf.Server.KeyFile, "s_key", "", "tls key file, gen cert from letsencrypt if empty") - flag.StringVar(&conf.Client.Address, "c", "", "remote server domain, eg: aa.bb.cc, socks5h://127.0.0.1:1080") - flag.StringVar(&conf.Client.HTTPProxy, "http_proxy", ":8080", "http proxy, empty to disable") - - if !flag.Parsed() { - flag.Parse() - } - - defer log.Infow("starting", "version", version, "date", date, "config", &conf) - if conf.file == "" { - return &conf.Client, &conf.Server, conf.Password - } - - for i := range loadConfigFns { - if err := loadConfigFns[i].fn(); err != nil { - log.Fatalw("load config", "config", conf.file, "step", loadConfigFns[i].step, "err", err) - } - } - return &conf.Client, &conf.Server, conf.Password -} - -// refreshFns will be executed while init and write new config -var loadConfigFns = []struct { - step string - fn func() error -}{{"parse file", func() error { - f, err := os.OpenFile(conf.file, os.O_RDONLY, 0644) - if err != nil { - return xerrors.New(err.Error()) - } - defer f.Close() - - return toml.NewDecoder(f).Decode(&conf) -}}, {"load referenced rule", func() error { - for _, addr := range conf.Client.Router.BlockRefs { - lines, err := getRemoteRuleLines(addr) - if err != nil { - return err - } - conf.Client.Router.BlockList = append(conf.Client.Router.BlockList, lines...) - } - for _, addr := range conf.Client.Router.ProxyRefs { - lines, err := getRemoteRuleLines(addr) - if err != nil { - return err - } - conf.Client.Router.ProxyList = append(conf.Client.Router.ProxyList, lines...) - } - for _, addr := range conf.Client.Router.DirectRefs { - lines, err := getRemoteRuleLines(addr) - if err != nil { - return err - } - conf.Client.Router.DirectList = append(conf.Client.Router.DirectList, lines...) - } - - return nil -}}} - -func getRemoteRuleLines(addr string) ([]string, error) { - resp, err := http.Get(addr) - if err != nil { - return nil, xerrors.New(err.Error()) - } - defer resp.Body.Close() - - br := bufio.NewReader(resp.Body) - lines := []string{} - for { - line, _, err := br.ReadLine() - if err == io.EOF { - return lines, nil - } else if err != nil { - return nil, xerrors.New(err.Error()) - } - - lines = append(lines, "**."+strings.TrimSpace(string(line))) - } -} - -// flushCh to avoid parallel persist -var flushCh = make(chan struct{}) -var flushOnce = sync.Once{} - -// PersistRule persist rule into config file -func PersistRule(domain string) { - flushOnce.Do(func() { - go flushConfDaemon() - }) - - log.Infow("persist direct rule into config", "domain", domain) - conf.Client.Router.DirectList = append(conf.Client.Router.DirectList, domain) - select { - case flushCh <- struct{}{}: - default: - } -} -func flushConfDaemon() { - for range flushCh { - // safe write file - if conf.file != "" { - f, err := os.OpenFile(conf.file+"~", os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644) - if err != nil { - log.Errorw("flush config", "step", "flush", "err", err) - continue - } - - conf.Client.Router.DirectList = - util.NewReverseSecSlice(conf.Client.Router.DirectList).Sort().Uniq() - - if err := toml.NewEncoder(f).ArraysWithOneElementPerLine(true).Encode(&conf); err != nil { - log.Errorw("flush config", "step", "flush", "err", err) - f.Close() - continue - } - f.Close() - - if stat, err := os.Stat(conf.file); err != nil { - log.Warnw("get file stat", "file", conf.file, "err", err) - } else { - // There is no common way to transfer ownership for a file - // cross-platform. Drop the ownership support but file mod. - if err = os.Chmod(conf.file+"~", stat.Mode()); err != nil { - log.Warnw("set file mod", "file", conf.file+"~", "err", err) - } - } - - if err = os.Rename(conf.file+"~", conf.file); err != nil { - log.Errorw("flush config", "step", "flush", "err", err) - continue - } - } - } -} diff --git a/conf/conf_darwin.go b/conf/conf_darwin.go deleted file mode 100644 index 817e98b..0000000 --- a/conf/conf_darwin.go +++ /dev/null @@ -1,100 +0,0 @@ -// +build darwin - -package conf - -import ( - "context" - "flag" - "fmt" - "io/ioutil" - "os" - "os/exec" - "path/filepath" - "time" - - "github.com/wweir/util-go/log" -) - -const svcPath = "/Library/LaunchDaemons/sower.plist" -const svcFile = ` - - - - Label - sower - ProgramArguments - - /bin/sh - -c - %s %s - - KeepAlive - - RunAtLoad - - -` - -var ( - ConfigDir = "" - installCmd = "" - uninstallFlag = false -) - -func beforeInitFlag() { - if _, err := os.Stat(execDir + "/sower.toml"); err == nil { - ConfigDir = execDir - } else { - dir, _ := os.UserConfigDir() - ConfigDir = filepath.Join("/", dir, "sower") - } - - if _, err := os.Stat(ConfigDir + "/sower.toml"); err != nil { - flag.StringVar(&conf.file, "f", "", "config file, rewrite all other parameters if set") - } else { - flag.StringVar(&conf.file, "f", ConfigDir+"/sower.toml", "config file, rewrite all other parameters if set") - } - - flag.StringVar(&installCmd, "install", "", "install service with cmd, eg: '-f \""+ConfigDir+"/sower.toml\"'") - flag.BoolVar(&uninstallFlag, "uninstall", false, "uninstall service") -} - -func afterInitFlag() { - switch { - case installCmd != "": - install() - case uninstallFlag: - uninstall() - default: - return - } - os.Exit(0) -} - -func install() { - if err := ioutil.WriteFile(svcPath, []byte(fmt.Sprintf(svcFile, execFile, installCmd)), 0644); err != nil { - log.Fatalw("write service file", "err", err) - } - - execute("launchctl unload " + svcPath) - if err := execute("launchctl load -wF " + svcPath); err != nil { - log.Fatalw("install service", "err", err) - } -} - -func uninstall() { - execute("launchctl unload " + svcPath) - os.Remove(svcPath) - os.RemoveAll("/etc/sower") -} - -func execute(cmd string) error { - ctx, cancel := context.WithTimeout(context.TODO(), 5*time.Second) - defer cancel() - - out, err := exec.CommandContext(ctx, "sh", "-c", cmd).CombinedOutput() - if err != nil { - return fmt.Errorf("cmd: %s, err: %s, output: %s", cmd, err, out) - } - return nil -} diff --git a/conf/conf_linux.go b/conf/conf_linux.go deleted file mode 100644 index 936cfeb..0000000 --- a/conf/conf_linux.go +++ /dev/null @@ -1,105 +0,0 @@ -// +build linux - -package conf - -import ( - "context" - "flag" - "fmt" - "io/ioutil" - "os" - "os/exec" - "path/filepath" - "time" - - "github.com/wweir/util-go/log" -) - -const svcPath = "/etc/systemd/system/sower.service" -const svcFile = `[Unit] -Description=Sower client service -After=network.target - -[Install] -WantedBy=multi-user.target - -[Service] -Type=simple -User=root -WorkingDirectory=/tmp -ExecStart=%s %s -RestartSec=3 -Restart=on-failure` - -var ( - ConfigDir = "" - installCmd = "" - uninstallFlag = false -) - -func beforeInitFlag() { - if _, err := os.Stat(execDir + "/sower.toml"); err == nil { - ConfigDir = execDir - - } else if stat, err := os.Stat("/etc/sower"); err == nil && stat.IsDir() { - ConfigDir = "/etc/sower" - - } else { - dir, _ := os.UserConfigDir() - ConfigDir = filepath.Join("/", dir, "sower") - } - - if _, err := os.Stat(ConfigDir + "/sower.toml"); err != nil { - flag.StringVar(&conf.file, "f", "", "config file, rewrite all other parameters if set") - } else { - flag.StringVar(&conf.file, "f", ConfigDir+"/sower.toml", "config file, rewrite all other parameters if set") - } - - flag.StringVar(&installCmd, "install", "", "install service with cmd, eg: '-f "+ConfigDir+"/sower.toml'") - flag.BoolVar(&uninstallFlag, "uninstall", false, "uninstall service") -} - -func afterInitFlag() { - switch { - case installCmd != "": - install() - case uninstallFlag: - uninstall() - default: - return - } - os.Exit(0) -} - -func install() { - if err := ioutil.WriteFile(svcPath, []byte(fmt.Sprintf(svcFile, execFile, installCmd)), 0644); err != nil { - log.Fatalw("write service file", "err", err) - } - if err := execute("systemctl daemon-reload"); err != nil { - log.Fatalw("install service", "err", err) - } - if err := execute("systemctl enable sower"); err != nil { - log.Fatalw("install service", "err", err) - } - if err := execute("systemctl start sower"); err != nil { - log.Fatalw("install service", "err", err) - } -} - -func uninstall() { - execute("systemctl stop sower") - execute("systemctl disable sower") - os.Remove(svcPath) - os.RemoveAll("/etc/sower") -} - -func execute(cmd string) error { - ctx, cancel := context.WithTimeout(context.TODO(), 5*time.Second) - defer cancel() - - out, err := exec.CommandContext(ctx, "sh", "-c", cmd).CombinedOutput() - if err != nil { - return fmt.Errorf("cmd: %s, err: %s, output: %s", cmd, err, out) - } - return nil -} diff --git a/conf/conf_windows.go b/conf/conf_windows.go deleted file mode 100644 index 4854f65..0000000 --- a/conf/conf_windows.go +++ /dev/null @@ -1,188 +0,0 @@ -// +build windows - -package conf - -import ( - "context" - "flag" - "fmt" - "os" - "os/exec" - "path/filepath" - "strings" - "syscall" - "time" - - "github.com/wweir/util-go/log" - "golang.org/x/sys/windows" - "golang.org/x/sys/windows/svc" - "golang.org/x/sys/windows/svc/eventlog" - "golang.org/x/sys/windows/svc/mgr" -) - -const name = "sower" -const cmdsAccepted = svc.AcceptStop | svc.AcceptShutdown | svc.AcceptPauseAndContinue - -var ( - ConfigDir = "" - installCmd = false - uninstallFlag = false -) - -func beforeInitFlag() { - flag.StringVar(&conf.file, "f", filepath.Join(execDir, "sower.toml"), "config file, rewrite all other parameters if set") - flag.BoolVar(&installCmd, "install", false, "put any character to install as a service, eg: true") -} - -func afterInitFlag() { - switch { - case installCmd: - install() - case uninstallFlag: - uninstall() - default: - runAsService() - return - } - os.Exit(0) -} - -func runAsService() { - os.Chdir(filepath.Dir(os.Args[0])) - if active, err := svc.IsAnInteractiveSession(); err != nil { - log.Fatalw("failed to determine if we are running in an interactive session", "err", err) - } else if !active { - go func() { - elog, err := eventlog.Open(name) - if err != nil { - log.Fatalw("install service", "err", err) - } - defer elog.Close() - - if err := svc.Run(name, &myservice{}); err != nil { - elog.Error(1, fmt.Sprintf("%s service failed: %v", name, err)) - log.Fatalw("install service", "err", err) - } - elog.Info(1, fmt.Sprintf("winsvc.RunAsService: %s service stopped", name)) - os.Exit(0) - }() - } -} - -func install() { - mgrDo(func(m *mgr.Mgr) error { - s, err := m.OpenService(name) - if err == nil { - s.Close() - return fmt.Errorf("service %s already exists", name) - } - s, err = m.CreateService(name, execFile, mgr.Config{ - DisplayName: "Sower Proxy", - StartType: windows.SERVICE_AUTO_START, - }) - if err != nil { - return err - } - defer s.Close() - err = eventlog.InstallAsEventCreate(name, eventlog.Error|eventlog.Warning|eventlog.Info) - if err != nil { - s.Delete() - return fmt.Errorf("SetupEventLogSource() failed: %s", err) - } - - return s.Start() - }) -} - -func uninstall() { - serviceDo(func(s *mgr.Service) error { - err := s.Delete() - if err != nil { - return err - } - return eventlog.Remove(name) - }) -} - -func serviceDo(fn func(*mgr.Service) error) { - mgrDo(func(m *mgr.Mgr) error { - s, err := m.OpenService(name) - if err != nil { - return fmt.Errorf("could not access service: %v", err) - } - defer s.Close() - return fn(s) - }) -} -func mgrDo(fn func(m *mgr.Mgr) error) { - m, err := mgr.Connect() - if err != nil { - log.Fatalw("install service", "err", err) - } - defer m.Disconnect() - - if err := fn(m); err != nil { - log.Fatalw("install service", "err", err) - } -} - -type myservice struct{} - -func (m *myservice) Execute(args []string, r <-chan svc.ChangeRequest, changes chan<- svc.Status) (ssec bool, errno uint32) { - elog, err := eventlog.Open(name) - if err != nil { - log.Errorw("install service", "err", err) - return - } - defer elog.Close() - elog.Info(1, strings.Join(args, "-")) - - changes <- svc.Status{State: svc.StartPending} - changes <- svc.Status{State: svc.Running, Accepts: cmdsAccepted} - for { - c := <-r - switch c.Cmd { - case svc.Interrogate: - changes <- c.CurrentStatus - // Testing deadlock from https://code.google.com/p/winsvc/issues/detail?id=4 - time.Sleep(100 * time.Millisecond) - changes <- c.CurrentStatus - case svc.Stop, svc.Shutdown: - changes <- svc.Status{State: svc.StopPending} - return - case svc.Pause: - changes <- svc.Status{State: svc.Paused, Accepts: cmdsAccepted} - case svc.Continue: - changes <- svc.Status{State: svc.Running, Accepts: cmdsAccepted} - default: - elog.Error(1, fmt.Sprintf("unexpected control request #%d", c)) - } - } -} - -func execute(cmd string) error { - ctx, cancel := context.WithTimeout(context.TODO(), 5*time.Second) - defer cancel() - - var cmds []string - for _, cmd := range strings.Split(cmd, " ") { - if cmd == "" { - continue - } - if strings.HasPrefix(cmd, "/") { - cmd = strings.Replace(cmd, "/", "-", 1) - } - cmds = append(cmds, cmd) - } - - if len(cmds) != 0 { - return nil - } - - command := exec.CommandContext(ctx, cmds[0], cmds[1:]...) - command.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} - if out, err := command.CombinedOutput(); err != nil { - return fmt.Errorf("cmd: %s, output: %s, err: %w", cmd, out, err) - } - return nil -} diff --git a/conf/sower.toml b/conf/sower.toml deleted file mode 100644 index 203ebf9..0000000 --- a/conf/sower.toml +++ /dev/null @@ -1,65 +0,0 @@ -password = "" # sower password - -[client] -address = "" # aa.bb.cc, socks5h://127.0.0.1:1080 -dns_upstream = "" # keep empty to set via dhcp, not effective in any environment -http_proxy = "" -socks5 = "" - -[client.port_forward] -# eg: ":2222"="aa.bb.cc:22" - -[client.router] -block_refs = [ - "https://cdn.jsdelivr.net/gh/pexcn/daily@gh-pages/adlist/adlist.txt", -] -detect_level = 0 # [-4, 4], the bigger the harder to add -direct_cird_refs = [ - "https://cdn.jsdelivr.net/gh/pexcn/daily@gh-pages/chnroute/chnroute.txt", -] -direct_list = [ - "**.in-addr.arpa", - "imap.*.*", - "imap.*.*.*", - "smtp.*.*", - "smtp.*.*.*", - "pop.*.*", - "pop.*.*.*", - "**.cn", -] -direct_refs = [ - "https://cdn.jsdelivr.net/gh/pexcn/daily@gh-pages/chinalist/chinalist.txt", -] -proxy_list = [ - "**.google.*", - "**.goo.gl", - "**.googleusercontent.com", - "**.googleapis.com", - "*.googlesource.com", - "**.youtube.com", - "**.ytimg.com", - "**.ggpht.com", - "**.googlevideo.com", - "**.facebook.com", - "**.fbcdn.net", - "**.twitter.com", - "**.twimg.com", - "**.blogspot.com", - "**.appspot.com", - "**.wikipedia.org", - "*.cloudfront.net", - "**.amazon.com", - "**.amazonaws.com", - "*.githubusercontent.com", - "*.githubassets.com", - "*.github.*", -] -proxy_refs = [ - "https://cdn.jsdelivr.net/gh/pexcn/daily@gh-pages/gfwlist/gfwlist.txt", -] - -[server] -cert_email = "" # eg: user@aa.bb.cc -cert_file = "" # eg: /etc/ssl/server.crt -key_file = "" # eg: /etc/ssl/server.key -upstream = "" # eg: 127.0.0.1:8080 diff --git a/go.mod b/go.mod index 90a4df5..8fe5968 100644 --- a/go.mod +++ b/go.mod @@ -3,14 +3,16 @@ module github.com/wweir/sower go 1.14 require ( + github.com/cristalhq/aconfig v0.16.1 + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/kr/pretty v0.1.0 // indirect github.com/krolaw/dhcp4 v0.0.0-20190909130307-a50d88189771 github.com/libp2p/go-reuseport v0.0.2 github.com/miekg/dns v1.1.30 - github.com/pelletier/go-toml v1.8.0 + github.com/oschwald/geoip2-golang v1.5.0 github.com/pkg/errors v0.9.1 - github.com/wweir/util-go/log v0.0.0-20200701032032-3cff7b4a46ea - github.com/wweir/util-go/mem v0.0.0-20200701032032-3cff7b4a46ea + github.com/rs/zerolog v1.23.0 + github.com/ulule/deepcopier v0.0.0-20200430083143-45decc6639b6 golang.org/x/crypto v0.0.0-20200709230013-948cd5f35899 - golang.org/x/sys v0.0.0-20190924154521-2837fb4f24fe - golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543 + gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127 // indirect ) diff --git a/go.sum b/go.sum index 0e55d3f..a2ff01f 100644 --- a/go.sum +++ b/go.sum @@ -1,10 +1,14 @@ -github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= +github.com/coreos/go-systemd/v22 v22.3.2/go.mod h1:Y58oyj3AT4RCenI/lSvhwexgC+NSVTIJ3seZv2GcEnc= +github.com/cristalhq/aconfig v0.16.1 h1:P3noQOujaPU5Of8E9fA1YYU3s/HUoPCIhpKkEgIQtfc= +github.com/cristalhq/aconfig v0.16.1/go.mod h1:NXaRp+1e6bkO4dJn+wZ71xyaihMDYPtCSvEhMTm/H3E= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/google/renameio v0.1.0/go.mod h1:KWCgfxg9yswjAJkECMjeO8J8rahYeXnNhOm40UhjYkI= -github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= +github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA= +github.com/kr/pretty v0.1.0 h1:L/CwN0zerZDmRFUapSPitk6f+Q3+0za1rQkzVuMiMFI= github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= +github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/krolaw/dhcp4 v0.0.0-20190909130307-a50d88189771 h1:t2c2B9g1ZVhMYduqmANSEGVD3/1WlsrEYNPtVoFlENk= github.com/krolaw/dhcp4 v0.0.0-20190909130307-a50d88189771/go.mod h1:0AqAH3ZogsCrvrtUpvc6EtVKbc3w6xwZhkvGLuqyi3o= @@ -12,64 +16,61 @@ github.com/libp2p/go-reuseport v0.0.2 h1:XSG94b1FJfGA01BUrT82imejHQyTxO4jEWqheyC github.com/libp2p/go-reuseport v0.0.2/go.mod h1:SPD+5RwGC7rcnzngoYC86GjPzjSywuQyMVAheVBD9nQ= github.com/miekg/dns v1.1.30 h1:Qww6FseFn8PRfw07jueqIXqodm0JKiiKuK0DeXSqfyo= github.com/miekg/dns v1.1.30/go.mod h1:KNUDUusw/aVsxyTYZM1oqvCicbwhgbNgztCETuNZ7xM= -github.com/pelletier/go-toml v1.8.0 h1:Keo9qb7iRJs2voHvunFtuuYFsbWeOBh8/P9v/kVMFtw= -github.com/pelletier/go-toml v1.8.0/go.mod h1:D6yutnOGMveHEPV7VQOuvI/gXY61bv+9bAOTRnLElKs= -github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/oschwald/geoip2-golang v1.5.0 h1:igg2yQIrrcRccB1ytFXqBfOHCjXWIoMv85lVJ1ONZzw= +github.com/oschwald/geoip2-golang v1.5.0/go.mod h1:xdvYt5xQzB8ORWFqPnqMwZpCpgNagttWdoZLlJQzg7s= +github.com/oschwald/maxminddb-golang v1.8.0 h1:Uh/DSnGoxsyp/KYbY1AuP0tYEwfs0sCph9p/UMXK/Hk= +github.com/oschwald/maxminddb-golang v1.8.0/go.mod h1:RXZtst0N6+FY/3qCNmZMBApR19cdQj43/NM9VkrNAis= github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/rogpeppe/go-internal v1.3.0/go.mod h1:M8bDsm7K2OlrFYOpmOWEs/qY81heoFRclV5y23lUDJ4= +github.com/rs/xid v1.2.1/go.mod h1:+uKXf+4Djp6Md1KODXJxgGQPKngRmWyn10oCKFzNHOQ= +github.com/rs/zerolog v1.23.0 h1:UskrK+saS9P9Y789yNNulYKdARjPZuS35B8gJF2x60g= +github.com/rs/zerolog v1.23.0/go.mod h1:6c7hFfxPOy7TacJc4Fcdi24/J0NKYGzjG8FWRI916Qo= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= -github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA= +github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/ulule/deepcopier v0.0.0-20200430083143-45decc6639b6 h1:TtyC78WMafNW8QFfv3TeP3yWNDG+uxNkk9vOrnDu6JA= github.com/ulule/deepcopier v0.0.0-20200430083143-45decc6639b6/go.mod h1:h8272+G2omSmi30fBXiZDMkmHuOgonplfKIKjQWzlfs= -github.com/wweir/util-go v0.0.0-20200701032032-3cff7b4a46ea h1:V/bIKmS5Bv4zMwQ5qE7ugJiHpcAVFnAxvxRbb7OCwug= -github.com/wweir/util-go/log v0.0.0-20200701032032-3cff7b4a46ea h1:IlO7d9R0I/l2syxJkx4/PURW9yVf4hy/Qf2IzdR4y38= -github.com/wweir/util-go/log v0.0.0-20200701032032-3cff7b4a46ea/go.mod h1:OObmMboiahxgJ1P3twG+AN8EAZ/wz14vstI4E5kaZzI= -github.com/wweir/util-go/mem v0.0.0-20200701032032-3cff7b4a46ea h1:nXgUp7CflDJnuR+zd43DHKxRA/hvuMkFtueSRLY7mEc= -github.com/wweir/util-go/mem v0.0.0-20200701032032-3cff7b4a46ea/go.mod h1:gp1gaDO1K4uV/B8NfrTjVim99Xq5xAmm1SHVamRsVW0= -go.uber.org/atomic v1.6.0 h1:Ezj3JGmsOnG1MoRWQkPBsKLe9DwWD9QeXzTRzzldNVk= -go.uber.org/atomic v1.6.0/go.mod h1:sABNBOSYdrvTF6hTgEIbc7YasKWGhgEQZyfxyTvoXHQ= -go.uber.org/multierr v1.5.0 h1:KCa4XfM8CWFCpxXRGok+Q0SS/0XBhMDbHHGABQLvD2A= -go.uber.org/multierr v1.5.0/go.mod h1:FeouvMocqHpRaaGuG9EjoKcStLC43Zu/fmqdUMPcKYU= -go.uber.org/tools v0.0.0-20190618225709-2cfd321de3ee/go.mod h1:vJERXedbb3MVM5f9Ejo0C68/HhF8uaILCdgjnY+goOA= -go.uber.org/zap v1.15.0 h1:ZZCA22JRF2gQE5FoNmhmrf7jeJJ2uhqDUNRYKm8dvmM= -go.uber.org/zap v1.15.0/go.mod h1:Mb2vm2krFEG5DV0W9qcHBYFtp/Wku1cvYaqPsS/WYfc= +github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= -golang.org/x/crypto v0.0.0-20190510104115-cbcb75029529/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= +golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20200709230013-948cd5f35899 h1:DZhuSZLsGlFL4CmhA8BcRA0mnthyA/nZ00AqCUo7vHg= golang.org/x/crypto v0.0.0-20200709230013-948cd5f35899/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/lint v0.0.0-20190930215403-16217165b5de/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= -golang.org/x/mod v0.0.0-20190513183733-4bf6d317e70e/go.mod h1:mXi4GBBbnImb6dmsKGUJ2LatrhH/nqhxcFungHvyanc= golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg= -golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= -golang.org/x/net v0.0.0-20190923162816-aa69164e4478 h1:l5EDrHhldLYb3ZRHDUhXF7Om7MvYXnkV9/iQNo1lX6g= golang.org/x/net v0.0.0-20190923162816-aa69164e4478/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= +golang.org/x/net v0.0.0-20201021035429-f5854403a974 h1:IX6qOQeG5uLjB/hjjwjedwfjND0hgjPMMyO1RoIXQNI= +golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9 h1:SQFwaSi55rU7vdNs9Yr0Z324VNlrF+0wMqRXT4St8ck= +golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190228124157-a34e9553db1e/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20190924154521-2837fb4f24fe h1:6fAMxZRR6sl1Uq8U61gxU+kPTs2tR8uOySCbBP7BN/M= golang.org/x/sys v0.0.0-20190924154521-2837fb4f24fe/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/text v0.3.0 h1:g61tztE5qeGQ89tm6NTjjM9VPIm088od1l6aSorWRWg= +golang.org/x/sys v0.0.0-20191224085550-c709ea063b76/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210119212857-b64e53b001e4 h1:myAQVi0cGEoqQVR5POX+8RR2mrocKqNN1hmeMqhX27k= +golang.org/x/sys v0.0.0-20210119212857-b64e53b001e4/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= -golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= -golang.org/x/tools v0.0.0-20190621195816-6e04913cbbac/go.mod h1:/rFqwRUd4F7ZHNgwSSTFct+R/Kf4OFW1sUzUTQQTgfc= -golang.org/x/tools v0.0.0-20191029041327-9cc4af7d6b2c/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= -golang.org/x/tools v0.0.0-20191029190741-b9c20aec41a5/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= +golang.org/x/text v0.3.3 h1:cokOdA+Jmi5PJGXLlLllQSgYigAEfHXJAERHVMaCc2k= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191216052735-49a3e744a425/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28= +golang.org/x/tools v0.1.0/go.mod h1:xkSsbof2nBLbhDlRMhhhyNLN/zl3eTqcnHD5viDpcZ0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543 h1:E7g+9GITq07hpfrRu66IVDexMakfv52eLZ2CXBWiKr4= -golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127 h1:qIbj1fsPNlZgppZ+VLlY7N33q108Sa+fhmuc+sWQYwY= gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -honnef.co/go/tools v0.0.1-2019.2.3/go.mod h1:a3bituU0lyd329TUQxRnasdCoJDkEUEAqEt0JzvZhAg= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/main.go b/main.go deleted file mode 100644 index d81aa43..0000000 --- a/main.go +++ /dev/null @@ -1,45 +0,0 @@ -package main - -import ( - "flag" - "fmt" - - "github.com/wweir/sower/conf" - "github.com/wweir/sower/proxy" - "github.com/wweir/sower/router" - "github.com/wweir/sower/transport" -) - -func main() { - client, server, password := conf.Init() - - switch { - case server.Upstream != "": - proxy.StartServer(server.Upstream, password, conf.ConfigDir, - server.CertFile, server.KeyFile, server.CertEmail) - - case client.Address != "": - route := router.NewRoute(client.Address, password, client.Router.DetectLevel, - client.Router.BlockList, client.Router.ProxyList, client.Router.DirectList, - conf.PersistRule) - - if client.Socks5Proxy != "" { - go proxy.StartSocks5Proxy(client.Socks5Proxy, client.Address, []byte(password)) - } - - if client.HTTPProxy != "" { - go proxy.StartHTTPProxy(client.HTTPProxy, client.Address, - []byte(password), route.GenProxyCheck(true)) - } - - transport.SetDNS(nil, client.DNSUpstream) - go proxy.StartDNS(client.DNSUpstream, route.GenProxyCheck(false)) - - proxy.StartClient(client.Address, password, - client.PortForward, route.GenProxyCheck(true)) - - default: - fmt.Println() - flag.Usage() - } -} diff --git a/dhcp/dhcp.go b/pkg/dhcp/dhcp.go similarity index 97% rename from dhcp/dhcp.go rename to pkg/dhcp/dhcp.go index 6edc6c0..60ccc88 100644 --- a/dhcp/dhcp.go +++ b/pkg/dhcp/dhcp.go @@ -14,7 +14,7 @@ import ( var xid = make([]byte, 4) var broadcastAddr, _ = net.ResolveUDPAddr("udp", "255.255.255.255:67") -func GetDefaultDNSServer() (string, error) { +func GetDNSServer() (string, error) { iface, err := PickInternetInterface() if err != nil { return "", errors.Wrap(err, "pick interface") diff --git a/dhcp/dhcp_test.go b/pkg/dhcp/dhcp_test.go similarity index 61% rename from dhcp/dhcp_test.go rename to pkg/dhcp/dhcp_test.go index 2f802c1..fdeca3c 100644 --- a/dhcp/dhcp_test.go +++ b/pkg/dhcp/dhcp_test.go @@ -3,11 +3,11 @@ package dhcp_test import ( "fmt" - "github.com/wweir/sower/dhcp" + "github.com/wweir/sower/pkg/dhcp" ) func Example_dns() { - got, err := dhcp.GetDefaultDNSServer() + got, err := dhcp.GetDNSServer() if err != nil { panic(err) } diff --git a/dhcp/pick_iface_other.go b/pkg/dhcp/pick_iface_other.go similarity index 100% rename from dhcp/pick_iface_other.go rename to pkg/dhcp/pick_iface_other.go diff --git a/dhcp/pick_iface_test.go b/pkg/dhcp/pick_iface_test.go similarity index 81% rename from dhcp/pick_iface_test.go rename to pkg/dhcp/pick_iface_test.go index b1fc743..424fa53 100644 --- a/dhcp/pick_iface_test.go +++ b/pkg/dhcp/pick_iface_test.go @@ -3,7 +3,7 @@ package dhcp_test import ( "fmt" - "github.com/wweir/sower/dhcp" + "github.com/wweir/sower/pkg/dhcp" ) func Example_iface() { diff --git a/dhcp/pick_iface_windows.go b/pkg/dhcp/pick_iface_windows.go similarity index 100% rename from dhcp/pick_iface_windows.go rename to pkg/dhcp/pick_iface_windows.go diff --git a/dhcp/util.go b/pkg/dhcp/util.go similarity index 100% rename from dhcp/util.go rename to pkg/dhcp/util.go diff --git a/pkg/mem/mem.go b/pkg/mem/mem.go new file mode 100644 index 0000000..9d83778 --- /dev/null +++ b/pkg/mem/mem.go @@ -0,0 +1,142 @@ +package mem + +import ( + "errors" + "fmt" + "reflect" + "sync" + "time" + + "github.com/ulule/deepcopier" +) + +// Data define the type which can speed up by mem cache +type Data interface { + Fulfill(key interface{}) error +} + +// Cache is the definition of cache, be careful of the memory usage +type Cache struct { + old *sync.Map + now *sync.Map + barrier *sync.Map + rotate <-chan time.Time + rwmutex *sync.RWMutex +} + +// DefaultCache is default cache for surge +var DefaultCache = New(time.Minute) + +// Remember is a surge, it provides a quite simple way to use cache +func Remember(dst Data, key interface{}) error { + return DefaultCache.Remember(dst, key) +} + +// Delete is a surge, it delete a specified data in DefaultCache +func Delete(dst Data, key interface{}) { + DefaultCache.Delete(dst, key) +} + +// New create a cache entity with a custom expiration time +func New(rotateInterval time.Duration) *Cache { + return &Cache{ + old: &sync.Map{}, + now: &sync.Map{}, + barrier: &sync.Map{}, + rotate: time.NewTicker(rotateInterval).C, + rwmutex: &sync.RWMutex{}, + } +} + +// Remember automatically save and retrieve data from a cache entity +func (c *Cache) Remember(dst Data, key interface{}) error { + rv := reflect.ValueOf(dst) + if rv.Kind() != reflect.Ptr { + panic("invalid not pointor type: " + reflect.TypeOf(dst).Name()) + } else if rv.IsNil() { + return errors.New("invalid nil pointor") + } + + c.rwmutex.RLock() + defer c.rwmutex.RUnlock() + + // rotate logic, rwlock just protect fields in Cache, but not field content. + // So that, write lock just take a very short time, and simple read lock is + // just an atomic action, do not care the performance + select { + case <-c.rotate: + c.old = c.now + c.now = &sync.Map{} + c.barrier = &sync.Map{} + default: + } + + // First: load from cache + cacheKey := fmt.Sprintf("%T%v", dst, key) + if val, ok := c.now.Load(cacheKey); ok { + return deepcopier.Copy(val).To(dst) + } + + // Second: load from old cache, or waitting the sigle groutine getting data + ch := make(chan struct{}) + if chVal, ok := c.barrier.LoadOrStore(cacheKey, ch); ok { + close(ch) // the ch is not used + + if val, ok := c.old.Load(cacheKey); ok { + return deepcopier.Copy(val).To(dst) + } + + // type chan: wait the sigle groutine getting data + // type error: already failed + if ch, ok = chVal.(chan struct{}); ok { + <-ch + if val, ok := c.now.Load(cacheKey); ok { + return deepcopier.Copy(val).To(dst) + } + } + + val, _ := c.barrier.Load(cacheKey) + if err, ok := val.(error); ok { + return err + } + + panic("new value lost, please report a bug") + } + + // Third: getting data from CacheType, maybe from db + err := dst.Fulfill(key) + if err != nil { + c.barrier.Store(cacheKey, err) + return err + } + + c.now.Store(cacheKey, dst) + close(ch) // broadcast, wakeup all waiting groutine + + return nil +} + +// Delete immediately specified the cached content to expire +func (c *Cache) Delete(dst Data, key interface{}) { + c.rwmutex.Lock() + defer c.rwmutex.Unlock() + + cacheKey := fmt.Sprintf("%T%v", dst, key) + c.old.Delete(cacheKey) + c.now.Delete(cacheKey) + c.barrier.Store(cacheKey, errors.New(cacheKey+"is deleted")) +} + +// Rotate force refresh cached data +func (c *Cache) Rotate(reset bool) { + c.rwmutex.Lock() + defer c.rwmutex.Unlock() + + if reset { + c.old = &sync.Map{} + } else { + c.old = c.now + } + c.now = &sync.Map{} + c.barrier = &sync.Map{} +} diff --git a/pkg/teeconn/tee_conn.go b/pkg/teeconn/tee_conn.go new file mode 100644 index 0000000..c302949 --- /dev/null +++ b/pkg/teeconn/tee_conn.go @@ -0,0 +1,56 @@ +package teeconn + +import ( + "io" + "net" +) + +type Conn struct { + net.Conn + + buf []byte + offset int + stop bool // read + err error +} + +func New(c net.Conn) *Conn { + return &Conn{Conn: c} +} + +func (t *Conn) Reread() { + t.offset = 0 +} +func (t *Conn) Reset() { + t.buf = []byte{} + t.offset = 0 +} +func (t *Conn) Stop() *Conn { + t.stop = true + return t +} + +func (t *Conn) Read(b []byte) (n int, err error) { + length := len(t.buf) - t.offset + if length > 0 { + n = copy(b, t.buf[t.offset:]) + t.offset += n + return n, t.err + } + + n, t.err = t.Conn.Read(b) + if !t.stop { + t.buf = append(t.buf, b[:n]...) + t.offset += n + } + + return n, t.err +} + +func (t *Conn) Write(b []byte) (n int, err error) { + if t.stop { + return t.Conn.Write(b) + } + + return 0, io.ErrShortWrite +} diff --git a/proxy/dns.go b/proxy/dns.go deleted file mode 100644 index 96379f6..0000000 --- a/proxy/dns.go +++ /dev/null @@ -1,118 +0,0 @@ -package proxy - -import ( - "net" - "strings" - "sync" - "time" - - "github.com/miekg/dns" - "github.com/wweir/sower/dhcp" - "github.com/wweir/util-go/log" -) - -type msgCache struct { - *dns.Msg - time.Time -} - -var cache sync.Map - -func StartDNS(relayServer string, shouldProxy func(string) (bool, bool)) { - var err error - if relayServer, err = pickRelayAddr(relayServer); err != nil { - log.Fatalw("pick upstream dns server", "err", err) - } - log.Infow("upstream dns", "addr", relayServer) - - dns.HandleFunc(".", func(w dns.ResponseWriter, r *dns.Msg) { - // *Msg r has an TSIG record and it was validated - if r.IsTsig() != nil && w.TsigStatus() == nil { - lastTsig := r.Extra[len(r.Extra)-1].(*dns.TSIG) - r.SetTsig(lastTsig.Hdr.Name, dns.HmacMD5, 300, time.Now().Unix()) - } - - //https://stackoverflow.com/questions/4082081/requesting-a-and-aaaa-records-in-single-dns-query/4083071#4083071 - if len(r.Question) == 0 { - return - } - - domain := r.Question[0].Name - if idx := strings.IndexByte(domain, ':'); idx > 0 { - domain = domain[:idx] // trim port - } - - if isblock, isproxy := shouldProxy(domain); isblock { - m := new(dns.Msg) - m.SetReply(r) - w.WriteMsg(m) - } else if isproxy { - host, _, _ := net.SplitHostPort(w.LocalAddr().String()) - w.WriteMsg(localA(r, domain, net.ParseIP(host))) - - } else if val, ok := cache.Load(domain); ok && val.(*msgCache).After(time.Now()) { - msg := val.(*msgCache) - for _, rr := range msg.Answer { - rr.Header().Ttl = uint32(time.Until(msg.Time).Seconds()) - } - msg.SetReply(r) - w.WriteMsg(msg.Msg) - - } else if msg, err := dns.Exchange(r, relayServer); err != nil || msg == nil { - cache.Delete(domain) - if server, err := pickRelayAddr(relayServer); err != nil { - log.Errorw("detect upstream dns", "err", err) - } else if relayServer != server { - relayServer = server - log.Infow("detect upstream dns", "addr", relayServer) - } - - } else { - cache.Delete(domain) - if len(msg.Answer) != 0 { - deadline := time.Now().Add( - time.Duration(msg.Answer[0].Header().Ttl) * time.Second) - cache.Store(domain, &msgCache{ - Msg: msg, - Time: deadline, - }) - } - - w.WriteMsg(msg) - } - }) - - server := &dns.Server{Addr: ":53", Net: "udp"} - log.Infow("start dns", "addr", server.Addr) - log.Fatalw("dns serve fail", "err", server.ListenAndServe()) -} - -func pickRelayAddr(relayServer string) (_ string, err error) { - if relayServer == "" { - if relayServer, err = dhcp.GetDefaultDNSServer(); err != nil { - return "", err - } - } - - if _, _, err := net.SplitHostPort(relayServer); err != nil { - return net.JoinHostPort(relayServer, "53"), nil - } - return relayServer, nil -} - -func localA(r *dns.Msg, domain string, localIP net.IP) *dns.Msg { - m := new(dns.Msg) - m.SetReply(r) - if localIP.To4() != nil { - m.Answer = []dns.RR{&dns.A{ - Hdr: dns.RR_Header{Name: domain, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 20}, - A: localIP, - }} - } else { - m.Answer = []dns.RR{&dns.AAAA{ - Hdr: dns.RR_Header{Name: domain, Rrtype: dns.TypeAAAA, Class: dns.ClassINET, Ttl: 20}, - AAAA: localIP, - }} - } - return m -} diff --git a/proxy/http_proxy.go b/proxy/http_proxy.go deleted file mode 100644 index a431541..0000000 --- a/proxy/http_proxy.go +++ /dev/null @@ -1,84 +0,0 @@ -package proxy - -import ( - "context" - "crypto/tls" - "net" - "net/http" - "net/http/httputil" - "time" - - "github.com/wweir/sower/transport" - "github.com/wweir/sower/util" - "github.com/wweir/util-go/log" -) - -// StartHTTPProxy start http reverse proxy. -// The httputil.ReverseProxy do not supply enough support for https request. -func StartHTTPProxy(httpProxyAddr, serverAddr string, password []byte, - shouldProxy func(string) (bool, bool)) { - - proxy := httputil.ReverseProxy{ - Director: func(r *http.Request) {}, - Transport: &http.Transport{ - DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - return transport.Dial(serverAddr, func(host string) (string, []byte) { - if _, ok := shouldProxy(host); ok { - return httpProxyAddr, password - } - return "", nil - }) - }, - }, - } - - srv := &http.Server{ - Addr: httpProxyAddr, - Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method == http.MethodConnect { - httpsProxy(w, r, serverAddr, password, shouldProxy) - } else { - proxy.ServeHTTP(w, r) - } - }), - // Disable HTTP/2. - TLSNextProto: map[string]func(*http.Server, *tls.Conn, http.Handler){}, - IdleTimeout: 90 * time.Second, - } - - log.Infow("start sower http proxy", "http_proxy", httpProxyAddr) - go log.Fatalw("serve http proxy", "addr", httpProxyAddr, "err", srv.ListenAndServe()) -} - -func httpsProxy(w http.ResponseWriter, r *http.Request, - serverAddr string, password []byte, shouldProxy func(string) (bool, bool)) { - - conn, _, err := w.(http.Hijacker).Hijack() - if err != nil { - http.Error(w, err.Error(), http.StatusServiceUnavailable) - return - } - conn.(*net.TCPConn).SetKeepAlive(true) - defer conn.Close() - - if _, err := conn.Write([]byte(r.Proto + " 200 Connection established\r\n\r\n")); err != nil { - http.Error(w, err.Error(), http.StatusServiceUnavailable) - return - } - - target, _ := util.WithDefaultPort(r.Host, "443") - rc, err := transport.Dial(target, func(host string) (string, []byte) { - if _, ok := shouldProxy(host); ok { - return serverAddr, password - } - return "", nil - }) - if err != nil { - conn.Write([]byte("sower dial " + serverAddr + " fail: " + err.Error())) - conn.Close() - return - } - defer rc.Close() - - relay(conn, rc) -} diff --git a/proxy/proxy.go b/proxy/proxy.go deleted file mode 100644 index bf3d4f2..0000000 --- a/proxy/proxy.go +++ /dev/null @@ -1,140 +0,0 @@ -package proxy - -import ( - "crypto/tls" - "net" - "net/http" - - "github.com/wweir/sower/transport" - "github.com/wweir/util-go/log" - "golang.org/x/crypto/acme/autocert" -) - -func StartClient(serverAddr, password string, - forwards map[string]string, shouldProxy func(string) (bool, bool)) { - - passwordData := []byte(password) - relayToRemote := func(lnAddr, target string, - parseFn func(net.Conn) (net.Conn, string, error), - shouldProxy func(string) (bool, bool)) { - - ln, err := net.Listen("tcp", lnAddr) - if err != nil { - log.Fatalw("tcp listen", "port", lnAddr, "err", err) - } - - for { - conn, err := ln.Accept() - if err != nil { - log.Errorw("tcp accept", "port", lnAddr, "err", err) - continue - } - - go func(conn net.Conn) { - defer conn.Close() - - if parseFn != nil { - if conn, target, err = parseFn(conn); err != nil { - log.Warnw("parse target", "err", err) - return - } - } - - rc, err := transport.Dial(target, func(domain string) (string, []byte) { - if _, ok := shouldProxy(domain); ok { - return serverAddr, passwordData - } - return "", nil - }) - if err != nil { - log.Warnw("dial", "addr", target, "err", err) - return - } - defer rc.Close() - - relay(conn, rc) - }(conn) - } - } - - for from, to := range forwards { - go relayToRemote(from, to, nil, func(string) (bool, bool) { return false, true }) - } - - go relayToRemote(":http", "", ParseHTTP, shouldProxy) - go relayToRemote(":https", "", ParseHTTPS, shouldProxy) - - log.Infow("start sower client", "forwards", forwards) - - select {} -} - -func StartServer(relayTarget, password, cacheDir, certFile, keyFile, email string) { - certManager := autocert.Manager{ - Prompt: autocert.AcceptTOS, - Email: email, - Cache: autocert.DirCache(cacheDir), - } - - tlsConf := &tls.Config{ - GetCertificate: certManager.GetCertificate, - MinVersion: tls.VersionTLS12, - NextProtos: []string{"http/1.1", "h2"}, - } - if certFile != "" && keyFile != "" { - if cert, err := tls.LoadX509KeyPair(certFile, keyFile); err != nil { - log.Fatalw("load certificate", "cert", certFile, "key", keyFile, "err", err) - } else { - tlsConf.GetCertificate = nil - tlsConf.Certificates = []tls.Certificate{cert} - } - } - - // Try to redirect 80 to 443 - go http.ListenAndServe(":80", certManager.HTTPHandler(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - r.URL.Scheme = "https" - if host, _, err := net.SplitHostPort(r.Host); err != nil { - r.URL.Host = r.Host - } else { - r.URL.Host = host - } - - http.Redirect(w, r, r.URL.String(), 301) - }))) - - log.Infow("start sower server", "relay_to", relayTarget) - ln, err := tls.Listen("tcp", ":443", tlsConf) - if err != nil { - log.Fatalw("tcp listen", "err", err) - } - - passwordData := []byte(password) - for { - conn, err := ln.Accept() - if err != nil { - log.Errorw("tcp accept", "err", err) - continue - } - - go func(conn net.Conn) { - defer conn.Close() - - target := relayTarget - - conn, t := transport.ParseTrojanConn(conn, passwordData) - if t != nil { - target = t.Addr() - } - - rc, err := net.Dial("tcp", target) - if err != nil { - log.Errorw("tcp dial", "addr", target, "err", err) - return - } - defer rc.Close() - - relay(conn, rc) - }(conn) - } -} diff --git a/proxy/socks5.go b/proxy/socks5.go deleted file mode 100644 index 5060a33..0000000 --- a/proxy/socks5.go +++ /dev/null @@ -1,44 +0,0 @@ -package proxy - -import ( - "net" - - "github.com/wweir/sower/transport" - "github.com/wweir/util-go/log" -) - -func StartSocks5Proxy(listenAddr, serverAddr string, password []byte) { - ln, err := net.Listen("tcp", listenAddr) - if err != nil { - log.Fatalw("socks5 proxy", "addr", listenAddr, "err", err) - } - - log.Infow("start socks5 proxy", "endpoint", listenAddr) - serveSocks5(ln, serverAddr, password) -} - -func serveSocks5(ln net.Listener, serverAddr string, password []byte) { - conn, err := ln.Accept() - if err != nil { - log.Fatalw("socks5 proxy", "err", err) - } - go serveSocks5(ln, serverAddr, password) - - tgtAddr, err := transport.ParseSocks5(conn) - if err != nil { - log.Errorw("socks5 proxy", "err", err) - return - } - - rc, err := transport.Dial(tgtAddr, func(host string) (string, []byte) { - return serverAddr, password - }) - if err != nil { - log.Errorw("socks5 proxy", "err", err) - return - } - - relay(conn, rc) - conn.Close() - rc.Close() -} diff --git a/proxy/util.go b/proxy/util.go deleted file mode 100644 index 84caadd..0000000 --- a/proxy/util.go +++ /dev/null @@ -1,64 +0,0 @@ -package proxy - -import ( - "bufio" - "crypto/tls" - "io" - "net" - "net/http" - "sync" - "sync/atomic" - "time" - - "github.com/wweir/sower/util" -) - -func ParseHTTP(conn net.Conn) (net.Conn, string, error) { - teeConn := &util.TeeConn{Conn: conn} - defer teeConn.Stop() - - resp, err := http.ReadRequest(bufio.NewReader(teeConn)) - if err != nil { - return teeConn, "", err - } - - resp.Host, _ = util.WithDefaultPort(resp.Host, "80") - return teeConn, resp.Host, nil -} - -func ParseHTTPS(conn net.Conn) (net.Conn, string, error) { - teeConn := &util.TeeConn{Conn: conn} - defer teeConn.Stop() - - var domain string - tls.Server(teeConn, &tls.Config{ - GetConfigForClient: func(hello *tls.ClientHelloInfo) (*tls.Config, error) { - domain = hello.ServerName - return nil, nil - }, - }).Handshake() - - domain, _ = util.WithDefaultPort(domain, "443") - return teeConn, domain, nil -} - -func relay(conn1, conn2 net.Conn) { - wg := &sync.WaitGroup{} - exitFlag := new(int32) - wg.Add(2) - go redirect(conn2, conn1, wg, exitFlag) - redirect(conn1, conn2, wg, exitFlag) - wg.Wait() -} -func redirect(dst, src net.Conn, wg *sync.WaitGroup, exitFlag *int32) { - io.Copy(dst, src) - - if atomic.CompareAndSwapInt32(exitFlag, 0, 1) { - // wakeup blocked goroutine - now := time.Now() - src.SetDeadline(now) - dst.SetDeadline(now) - } - - wg.Done() -} diff --git a/router/dns.go b/router/dns.go new file mode 100644 index 0000000..773c4d8 --- /dev/null +++ b/router/dns.go @@ -0,0 +1,82 @@ +package router + +import ( + "net" + "time" + + "github.com/miekg/dns" + "github.com/rs/zerolog/log" +) + +func (r *Router) ServeDNS(w dns.ResponseWriter, req *dns.Msg) { + // *Msg r has an TSIG record and it was validated + if req.IsTsig() != nil && w.TsigStatus() == nil { + lastTsig := req.Extra[len(req.Extra)-1].(*dns.TSIG) + req.SetTsig(lastTsig.Hdr.Name, dns.HmacMD5, 300, time.Now().Unix()) + } + + // https://stackoverflow.com/questions/4082081/requesting-a-and-aaaa-records-in-single-dns-query/4083071#4083071 + if len(req.Question) == 0 { + w.WriteMsg(r.dnsFail(req, dns.RcodeFormatError)) + return + } + + domain := req.Question[0].Name + switch { + case r.blockRule.Match(domain): + w.WriteMsg(r.dnsFail(req, dns.RcodeNameError)) + return + + case r.directRule.Match(domain): + + case r.proxyRule.Match(domain): + host, _, _ := net.SplitHostPort(w.LocalAddr().String()) + w.WriteMsg(r.dnsProxyA(domain, net.ParseIP(host), req)) + return + } + + conn := <-r.dns.connCh + resp, rtt, err := r.dns.ExchangeWithConn(req, conn) + if err != nil { + log.Error().Err(err). + Dur("rtt", rtt). + Str("domain", domain). + Msg("exchange dns record") + + conn.Close() + w.WriteMsg(r.dnsFail(req, dns.RcodeServerFailure)) + return + } + + select { + case r.dns.connCh <- conn: + default: + conn.Close() + } + w.WriteMsg(resp) +} + +func (r *Router) dnsFail(req *dns.Msg, rcode int) *dns.Msg { + m := new(dns.Msg) + m.SetRcode(req, dns.RcodeServerFailure) + return m +} + +func (r *Router) dnsProxyA(domain string, localIP net.IP, req *dns.Msg) *dns.Msg { + m := new(dns.Msg) + m.SetReply(req) + + if localIP.To4() != nil { + m.Answer = []dns.RR{&dns.A{ + Hdr: dns.RR_Header{Name: domain, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 20}, + A: localIP, + }} + + } else { + m.Answer = []dns.RR{&dns.AAAA{ + Hdr: dns.RR_Header{Name: domain, Rrtype: dns.TypeAAAA, Class: dns.ClassINET, Ttl: 20}, + AAAA: localIP, + }} + } + return m +} diff --git a/router/http_ping.go b/router/http_ping.go deleted file mode 100644 index 01ad5d8..0000000 --- a/router/http_ping.go +++ /dev/null @@ -1,124 +0,0 @@ -package router - -import ( - "bytes" - "crypto/tls" - "encoding/binary" - "io" - "net" - "strconv" - "time" -) - -// Port ========================== -type Port uint16 - -const ( - HTTP Port = 80 - HTTPS Port = 443 -) - -// Ping try connect to a http(s) server with domain though the http addr -func (p Port) Ping(domain string, dial func(string) (net.Conn, error)) error { - conn, err := dial(net.JoinHostPort(domain, p.String())) - if err != nil { - return err - } - defer conn.Close() - - conn.SetDeadline(time.Now().Add(5 * time.Second)) - if _, err := conn.Write(p.PingMsg(domain)); err != nil { - return err - } - - // err -> nil: read something succ - // err -> io.EOF: no such domain or connection refused - // err -> timeout: tcp package has been dropped - if _, err = conn.Read(make([]byte, 10)); err == io.EOF { - err = nil - } - return err -} - -func (p Port) String() string { - return strconv.Itoa(int(p)) -} - -func (p Port) PingMsg(domain string) []byte { - switch p { - case HTTP: - return []byte("TRACE / HTTP/1.1\r\nHost: " + domain + "\r\n\r\n") - case HTTPS: - return NewClientHelloSNIMsg(domain) - default: - panic("invalid port") - } -} - -// SNI ========================== -type clientHelloSNI struct { - ContentType uint8 - Version uint16 - Length uint16 - handshakeProtocol -} -type handshakeProtocol struct { - HandshakeType uint8 - LengthExpand uint8 - Length uint16 - Version uint16 - Random [32]byte - SessionIDLength uint8 - CipherSuitesLength uint16 - CipherSuite uint16 - CompressionMethodsLength uint8 - CompressionMethod uint8 - ExtensionsLength uint16 - extensionServerName -} -type extensionServerName struct { - Type uint16 - Length uint16 - serverNameIndicationExtension -} -type serverNameIndicationExtension struct { - ServerNameListlength uint16 - ServerNameType uint8 - ServerNamelength uint16 - // ServerName []byte // Disable for fix length -} - -func NewClientHelloSNIMsg(domain string) []byte { - length := uint16(len(domain)) - msg := &clientHelloSNI{ - ContentType: 0x16, // Content Type: Handshake (22) - Version: 0x0301, // Version: TLS 1.0 - Length: length + 56, - handshakeProtocol: handshakeProtocol{ - HandshakeType: 0x01, // Handshake Type: Client Hello (1) - Length: length + 52, - Version: 0x0303, // Version: TLS 1.2 (0x0303) - Random: [32]byte{}, // [32]byte{}, - SessionIDLength: 0x0, // Session ID Length: 0 - CipherSuitesLength: 2, // Cipher Suites Length: 84 - CipherSuite: tls.TLS_RSA_WITH_AES_128_GCM_SHA256, - CompressionMethodsLength: 1, // Compression Methods Length: 1 - CompressionMethod: 0x00, // Compression null - ExtensionsLength: length + 9, - extensionServerName: extensionServerName{ - Type: 0x0000, // Type: server_name (0) - Length: length + 5, - serverNameIndicationExtension: serverNameIndicationExtension{ - ServerNameListlength: length + 3, - ServerNameType: 0x00, // Server Name Type: host_name (0) - ServerNamelength: length, - }, - }, - }, - } - - buf := bytes.NewBuffer(make([]byte, 0, length+61)) - binary.Write(buf, binary.BigEndian, msg) - buf.WriteString(domain) - return buf.Bytes() -} diff --git a/router/http_ping_test.go b/router/http_ping_test.go deleted file mode 100644 index 6ffb0df..0000000 --- a/router/http_ping_test.go +++ /dev/null @@ -1,72 +0,0 @@ -// +build debug -// go test -v -tags debug - -package router_test - -import ( - "crypto/tls" - "net" - "strconv" - "testing" - "time" - - "github.com/wweir/sower/router" - "github.com/wweir/sower/transport" -) - -var ( - proxyAddr string = "socks5://127.0.0.1:1080" - password []byte = nil -) - -func TestPort_Ping(t *testing.T) { - direct := func(addr string) (net.Conn, error) { - return net.DialTimeout("tcp", addr, 5*time.Second) - } - proxy := func(addr string) (net.Conn, error) { - conn, err := tls.Dial("tcp", proxyAddr, &tls.Config{}) - if err != nil { - return nil, err - } - host, port, err := net.SplitHostPort(addr) - if err != nil { - return nil, err - } - p, err := strconv.Atoi(port) - if err != nil { - return nil, err - } - return transport.ToTrojanConn(conn, host, uint16(p), password) - } - - type args struct { - domain string - dial func(string) (net.Conn, error) - } - tests := []struct { - name string - p router.Port - args args - wantErr bool - }{ - {"", router.HTTP, args{"baidu.com", direct}, false}, - {"", router.HTTP, args{"baidu.com", proxy}, false}, - {"", router.HTTPS, args{"baidu.com", direct}, false}, - {"", router.HTTPS, args{"baidu.com", proxy}, false}, - {"", router.HTTP, args{"google.com", direct}, true}, - {"", router.HTTP, args{"google.com", proxy}, false}, - {"", router.HTTPS, args{"google.com", direct}, true}, - {"", router.HTTPS, args{"google.com", proxy}, false}, - {"", router.HTTP, args{"smtp.163.com", direct}, true}, - {"", router.HTTP, args{"smtp.163.com", proxy}, true}, - {"", router.HTTPS, args{"smtp.163.com", direct}, true}, - {"", router.HTTPS, args{"smtp.163.com", proxy}, true}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if err := tt.p.Ping(tt.args.domain, tt.args.dial); (err != nil) != tt.wantErr { - t.Errorf("Port.Ping() error = %v, wantErr %v", err, tt.wantErr) - } - }) - } -} diff --git a/router/mmdb.go b/router/mmdb.go new file mode 100644 index 0000000..e6b2923 --- /dev/null +++ b/router/mmdb.go @@ -0,0 +1,52 @@ +package router + +import ( + "context" + "net" + "time" + + "github.com/rs/zerolog/log" +) + +func (r *Router) localSite(domain string) bool { + + ip := net.ParseIP(domain) + if ip == nil { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + ips, err := r.mmdb.Resolver.LookupIP(ctx, "ip", domain) + if err != nil || len(ips) == 0 { + log.Warn().Err(err). + Str("domain", domain). + Int("ips", len(ips)). + Msg("resolve domain") + return false + } + + ip = ips[0] + } + + for _, cidr := range r.mmdb.cidrs { + if cidr.Contains(ip) { + return true + } + } + + if r.mmdb.Reader != nil { + city, err := r.mmdb.City(ip) + if err != nil { + log.Warn().Err(err). + Str("domain", domain). + IPAddr("ip", ip). + Msg("mmdb search") + return false + } + + if city.Country.IsoCode == "CN" { + return true + } + } + + return false +} diff --git a/router/ping.go b/router/ping.go new file mode 100644 index 0000000..ee803a9 --- /dev/null +++ b/router/ping.go @@ -0,0 +1,27 @@ +package router + +import ( + "net" + "net/http" + "time" +) + +var pingClient = http.Client{ + Timeout: 2 * time.Second, +} + +func (r *Router) isAccess(domain string) bool { + p := &ping{} + r.cache.Remember(p, domain) + return p.isAccess +} + +type ping struct { + isAccess bool +} + +func (p *ping) Fulfill(key interface{}) error { + _, err := http.Head(net.JoinHostPort(key.(string), "80")) + p.isAccess = (err == nil) + return nil +} diff --git a/router/router.go b/router/router.go index e197a6d..ed01572 100644 --- a/router/router.go +++ b/router/router.go @@ -2,152 +2,117 @@ package router import ( "net" - "strings" - "sync" - "sync/atomic" + "strconv" "time" - "github.com/wweir/sower/transport" + "github.com/miekg/dns" + geoip2 "github.com/oschwald/geoip2-golang" + "github.com/pkg/errors" + "github.com/rs/zerolog/log" + "github.com/wweir/sower/pkg/dhcp" + "github.com/wweir/sower/pkg/mem" "github.com/wweir/sower/util" - "github.com/wweir/util-go/log" - "github.com/wweir/util-go/mem" ) -// Route implement a router for each request -type Route struct { - once sync.Once - cache *mem.Cache +type Router struct { + blockRule *util.Node + directRule *util.Node + proxyRule *util.Node - ProxyAddress string - password []byte + ProxyDial func(network, host string, port uint16) (net.Conn, error) + cache *mem.Cache - DetectLevel int // dynamic detect proxy level - blockRule *util.Node - proxyRule *util.Node - directRule *util.Node - PersistFn func(string) -} + dns struct { + dns.Client + connCh chan *dns.Conn + } -func NewRoute(address, password string, detectLevel int, - blocklist, proxylist, directlist []string, persist func(string)) *Route { - return &Route{ - cache: mem.New(4 * time.Hour), - ProxyAddress: address, - password: []byte(password), + mmdb struct { + *geoip2.Reader + *net.Resolver - DetectLevel: detectLevel, - blockRule: util.NewNodeFromRules(blocklist...), - proxyRule: util.NewNodeFromRules(proxylist...), - directRule: util.NewNodeFromRules(directlist...), - PersistFn: persist, + cidrs []*net.IPNet } } -// ShouldProxy check if the domain shoule request though proxy -func (r *Route) GenProxyCheck(sync bool) func(string) (bool, bool) { - detect := func(domain string) bool { - go r.cache.Remember(r, domain) - return true +func NewRouter(proxyDial func(network, host string, port uint16) (net.Conn, error)) *Router { + r := Router{ + blockRule: util.NewNodeFromRules(), + directRule: util.NewNodeFromRules(), + proxyRule: util.NewNodeFromRules("google.*"), + ProxyDial: proxyDial, + cache: mem.New(time.Hour), // TODO: config } - if sync { - detect = func(domain string) bool { - r.cache.Remember(r, domain) - if r.proxyRule.Match(domain) { - return true + + r.dns.connCh = make(chan *dns.Conn, 1) + go r.dialDNSConn() + + return &r +} + +func (r *Router) dialDNSConn() { + for { + server, err := dhcp.GetDNSServer() + if err != nil { + time.Sleep(time.Second) + continue + } + + log.Info(). + Str("ip", server). + Msg("get upstream dns server") + + for { + conn, err := dns.DialTimeout("udp", net.JoinHostPort(server, "53"), time.Second) + if err != nil { + log.Error().Err(err).Str("ip", server).Msg("dial dns server") + break } - return !r.directRule.Match(domain) - } - } - return func(domain string) (bool, bool) { - domain = strings.TrimSuffix(domain, ".") - // break deadlook, for wildcard - if sepCount := strings.Count(domain, "."); sepCount == 0 || sepCount >= 5 { - return false, false + r.dns.connCh <- conn } - - if r.blockRule.Match(domain) { - return true, false - } - - if r.proxyRule.Match(domain) { - return false, true - } - if r.directRule.Match(domain) { - return false, false - } - - return false, detect(domain) } } -// Get implement for cache -func (r *Route) Get(key interface{}) (err error) { - domain := key.(string) +func (r *Router) RouteHandle(conn net.Conn, domain string, port uint16) error { + addr := net.JoinHostPort(domain, strconv.FormatUint(uint64(port), 10)) - httpScore, httpsScore := r.detect(domain) - log.Infow("detect", "domain", domain, "http", httpScore, "https", httpsScore) + switch { + case r.blockRule.Match(domain): + return nil - if httpScore+httpsScore >= r.DetectLevel { - r.directRule.Add(domain) - if r.PersistFn != nil { - r.PersistFn(domain) - } + case r.directRule.Match(domain): + return r.DirectHandle(conn, addr) + + case r.proxyRule.Match(domain): + return r.ProxyHandle(conn, domain, port) + + case r.localSite(domain): + return r.DirectHandle(conn, addr) + case r.isAccess(domain): + return r.DirectHandle(conn, addr) + case port == 80: + return r.DirectHandle(conn, addr) + case port == 443: + return r.DirectHandle(conn, addr) + default: + return r.ProxyHandle(conn, domain, port) } +} + +func (r *Router) ProxyHandle(conn net.Conn, domain string, port uint16) error { + start := time.Now() + rc, err := r.ProxyDial("tcp", domain, port) + if err != nil { + return errors.Wrapf(err, "dial proxy to (%s:%d), spend (%s)", domain, port, time.Since(start)) + } + defer rc.Close() + + util.Relay(conn, rc) return nil } -// detect and caculate direct connection and proxy connection score -func (r *Route) detect(domain string) (http, https int) { - wg := sync.WaitGroup{} - httpScore, httpsScore := new(int32), new(int32) - for _, ping := range [...]struct { - shouldProxy bool - port Port - }{ - {shouldProxy: true, port: HTTP}, - {shouldProxy: true, port: HTTPS}, - {shouldProxy: false, port: HTTP}, - {shouldProxy: false, port: HTTPS}, - } { - wg.Add(1) - go func(shouldProxy bool, port Port) { - defer wg.Done() - - if err := port.Ping(domain, func(domain string) (net.Conn, error) { - return transport.Dial(domain, - func(domain string) (proxyAddr string, password []byte) { - if shouldProxy { - return r.ProxyAddress, r.password - } - return "", nil - }) - }); err != nil { - log.Warnw("sower dial", "proxy", shouldProxy, "host", domain, "port", port, "err", err) - return - } - - switch { - case shouldProxy && port == HTTP: - if !atomic.CompareAndSwapInt32(httpScore, 0, -2) { - atomic.AddInt32(httpScore, -1) - } - case shouldProxy && port == HTTPS: - if !atomic.CompareAndSwapInt32(httpsScore, 0, -2) { - atomic.AddInt32(httpsScore, -1) - } - case !shouldProxy && port == HTTP: - if !atomic.CompareAndSwapInt32(httpScore, 0, 2) { - atomic.AddInt32(httpScore, 1) - } - case !shouldProxy && port == HTTPS: - if !atomic.CompareAndSwapInt32(httpsScore, 0, 2) { - atomic.AddInt32(httpsScore, 1) - } - } - }(ping.shouldProxy, ping.port) - } - - wg.Wait() - return int(*httpScore), int(*httpsScore) +func (r *Router) DirectHandle(conn net.Conn, addr string) error { + dur, err := util.RelayTo(conn, addr) + return errors.Wrapf(err, "direct relay to (%s), spend (%s)", addr, dur) } diff --git a/transport/proxy_conn.go b/transport/proxy_conn.go deleted file mode 100644 index 9848300..0000000 --- a/transport/proxy_conn.go +++ /dev/null @@ -1,68 +0,0 @@ -package transport - -import ( - "crypto/md5" - "encoding/binary" - "io" - "net" - "strconv" - - "github.com/wweir/sower/util" -) - -// checksum(>=0x80) + port + target_length + target + data -// data(HTTP, first byte < 0x7F) -type head struct { - Checksum byte - Port uint16 - AddrLen uint8 -} - -func ParseProxyConn(conn net.Conn, password []byte) (net.Conn, string) { - teeConn := &util.TeeConn{Conn: conn} - defer teeConn.Stop() - - h := &head{} - if err := binary.Read(teeConn, binary.BigEndian, h); err != nil || h.Checksum < 0x80 { - return teeConn, "" - } - - buf := make([]byte, int(h.AddrLen)) - if _, err := io.ReadFull(teeConn, buf); err != nil { - return teeConn, "" - } - - if h.Checksum != sumChecksum(buf, password) { - return teeConn, "" - } - - teeConn.Reset() - return teeConn, net.JoinHostPort(string(buf), strconv.Itoa(int(h.Port))) -} - -func ToProxyConn(conn net.Conn, tgtHost string, tgtPort uint16, password []byte) (net.Conn, error) { - h := &head{ - Checksum: sumChecksum([]byte(tgtHost), password), - Port: tgtPort, - AddrLen: uint8(len(tgtHost)), - } - if err := binary.Write(conn, binary.BigEndian, h); err != nil { - conn.Close() - return nil, err - } - - var data = []byte(tgtHost) - var err error - for n, nn := 0, 0; nn < len(data); nn += n { - if n, err = conn.Write(data[nn:]); err != nil { - conn.Close() - return nil, err - } - } - - return conn, nil -} - -func sumChecksum(target, password []byte) byte { - return md5.Sum(append(target, password...))[0] | 0x80 -} diff --git a/transport/socks5_client.go b/transport/socks5_client.go deleted file mode 100644 index 079a9b7..0000000 --- a/transport/socks5_client.go +++ /dev/null @@ -1,124 +0,0 @@ -package transport - -import ( - "encoding/binary" - "net" - "strings" - - "golang.org/x/xerrors" -) - -func IsSocks5Schema(addr string) (string, bool) { - if strings.HasPrefix(addr, "socks5://") { - return strings.TrimPrefix(addr, "socks5://"), true - } - - if strings.HasPrefix(addr, "socks5h://") { - return strings.TrimPrefix(addr, "socks5h://"), true - } - - return addr, false -} - -func ToSocks5(c net.Conn, host string, port uint16) (net.Conn, error) { - return &conn{ - init: make(chan struct{}), - Conn: c, - domain: host, - port: port, - }, nil -} - -type conn struct { - init chan struct{} - domain string - port uint16 - net.Conn -} - -func (c *conn) Read(b []byte) (n int, err error) { - select { - case <-c.init: - return c.Conn.Read(b) - default: - return 0, c.clientHandshake() - } -} - -func (c *conn) Write(b []byte) (n int, err error) { - select { - case <-c.init: - return c.Conn.Write(b) - default: - return 0, c.clientHandshake() - } -} - -func (c *conn) clientHandshake() error { - { - req := &authReq{ - VER: 5, - NMETHODS: 1, - METHODS: [1]byte{0}, // NO AUTHENTICATION REQUIRED - } - if err := binary.Write(c.Conn, binary.BigEndian, req); err != nil { - return xerrors.New(err.Error()) - } - } - { - resp := &authResp{} - if err := binary.Read(c.Conn, binary.BigEndian, resp); err != nil { - return xerrors.New(err.Error()) - } - } - { - reqHead := requestHead{ - VER: 5, // socks5 - CMD: 1, // CONNECT - RSV: 0, // RESERVED - ATYP: 3, // DOMAINNAME - } - - if err := binary.Write(c.Conn, binary.BigEndian, reqHead); err != nil { - return xerrors.New(err.Error()) - } - - buf := make([]byte, 0, 1 /*LEN*/ +len(c.domain)+2 /*PORT*/) - buf = append(buf, byte(len(c.domain))) - buf = append(buf, []byte(c.domain)...) - buf = append(buf, byte(c.port>>8), byte(c.port)) - if _, err := c.Conn.Write(buf); err != nil { - return xerrors.New(err.Error()) - } - } - { - head := replyHead{} - if err := binary.Read(c.Conn, binary.BigEndian, &head); err != nil { - return xerrors.New(err.Error()) - } - - switch head.REP { - case 0x00: - default: - return xerrors.Errorf("socks5 handshake fail, return code: %d", head.REP) - } - - var bindAddr addrType - switch head.ATYP { - case 0x01: // IPv4 - bindAddr = addrTypeIPv4{} - case 0x03: // domain name - bindAddr = addrTypeDomain{} - case 0x04: // IPv6 - bindAddr = addrTypeIPv6{} - default: - return xerrors.New("invalid connect type") - } - if err := bindAddr.Fullfill(c.Conn); err != nil { - return xerrors.New(err.Error()) - } - } - - close(c.init) - return nil -} diff --git a/transport/socks5_rfc.go b/transport/socks5_rfc.go deleted file mode 100644 index 8ff3151..0000000 --- a/transport/socks5_rfc.go +++ /dev/null @@ -1,115 +0,0 @@ -package transport - -import ( - "encoding/binary" - "io" - "net" - "strconv" - - "golang.org/x/xerrors" -) - -// https://tools.ietf.org/html/rfc1928 - -// 1. client send auth request -type authReq struct { - VER byte - NMETHODS byte - METHODS [1]byte // 1 to 255, fix to no authentication -} - -// 2. server response auth request -type authResp struct { - VER byte - METHOD byte -} - -// 3. client request with target address -type requestHead struct { - VER byte - CMD byte - RSV byte - ATYP byte -} - -// 4. server response with the address that assigned to connect to target address -type replyHead struct { - VER byte - REP byte - RSV byte - ATYP byte -} - -type addrType interface { - Fullfill(r io.Reader) error - String() string -} - -// ATYP: -// 0x01 -> net.IPv4len -// 0x03 -> first byte is length -// 0x04 -> net.IPv6len -type addrTypeIPv4 struct { - DST_ADDR [4]byte - DST_PORT uint16 -} - -func (a addrTypeIPv4) Fullfill(r io.Reader) error { - return binary.Read(r, binary.BigEndian, &a) -} -func (a addrTypeIPv4) String() string { - return net.JoinHostPort( - net.IP(a.DST_ADDR[:]).String(), - strconv.FormatUint(uint64(a.DST_PORT), 10), - ) -} - -type addrTypeDomain struct { - DST_ADDR_LEN uint8 - DST_ADDR []byte - DST_PORT uint16 -} - -func (a addrTypeDomain) Fullfill(r io.Reader) error { - buf := make([]byte, 256) - // domain length - if _, err := io.ReadFull(r, buf[:1]); err != nil { - return xerrors.New(err.Error()) - } - a.DST_ADDR_LEN = uint8(buf[0]) - - // domain - if _, err := io.ReadFull(r, buf[:int(buf[0])]); err != nil { - return xerrors.New(err.Error()) - } - a.DST_ADDR = buf[:int(buf[0])] - - // port - if _, err := io.ReadFull(r, buf[:2]); err != nil { - return xerrors.New(err.Error()) - } - a.DST_PORT = binary.BigEndian.Uint16(buf[:2]) - - return nil -} -func (a addrTypeDomain) String() string { - return net.JoinHostPort( - net.IP(a.DST_ADDR[:]).String(), - strconv.FormatUint(uint64(a.DST_PORT), 10), - ) -} - -type addrTypeIPv6 struct { - DST_ADDR [16]byte - DST_PORT uint16 -} - -func (a addrTypeIPv6) Fullfill(r io.Reader) error { - return binary.Read(r, binary.BigEndian, &a) -} -func (a addrTypeIPv6) String() string { - return net.JoinHostPort( - net.IP(a.DST_ADDR[:]).String(), - strconv.FormatUint(uint64(a.DST_PORT), 10), - ) -} diff --git a/transport/socks5_server.go b/transport/socks5_server.go deleted file mode 100644 index 6166e8f..0000000 --- a/transport/socks5_server.go +++ /dev/null @@ -1,74 +0,0 @@ -package transport - -import ( - "encoding/binary" - "net" - - "golang.org/x/xerrors" -) - -func ParseSocks5(conn net.Conn) (tgtaddr string, err error) { - { - authReq := new(authReq) - if err := binary.Read(conn, binary.BigEndian, authReq); err != nil { - return "", xerrors.New(err.Error()) - } - - if authReq.VER != 5 || // socks5 - authReq.NMETHODS != 1 || - authReq.METHODS[0] != 0 { // NO_AUTH - return "", xerrors.New("invalid socks5 auth method") - } - } - { - if err := binary.Write(conn, binary.BigEndian, &authResp{ - VER: 5, - METHOD: 1, - }); err != nil { - return "", xerrors.New(err.Error()) - } - } - { - head := &requestHead{} - if err := binary.Read(conn, binary.BigEndian, head); err != nil { - return "", xerrors.New(err.Error()) - } - if head.VER != 5 || - head.CMD != 1 { - return "", xerrors.New("invalid socks5 connect request") - } - - var addr addrType - switch head.ATYP { - case 0x01: // IPv4 - addr = addrTypeIPv4{} - case 0x03: // domain name - addr = addrTypeDomain{} - case 0x04: // IPv6 - addr = addrTypeIPv6{} - default: - return "", xerrors.New("invalid connect type") - } - if err := addr.Fullfill(conn); err != nil { - return "", xerrors.New(err.Error()) - } - tgtaddr = addr.String() - } - { - if err := binary.Write(conn, binary.BigEndian, &replyHead{ - VER: 5, - REP: 0, - RSV: 0, - ATYP: 1, - }); err != nil { - return "", xerrors.New(err.Error()) - } - - // FIXME: return the real address - if _, err := conn.Write(make([]byte, 6)); err != nil { - return "", xerrors.New(err.Error()) - } - } - - return tgtaddr, nil -} diff --git a/transport/sower/sower.go b/transport/sower/sower.go new file mode 100644 index 0000000..70a911d --- /dev/null +++ b/transport/sower/sower.go @@ -0,0 +1,79 @@ +package sower + +import ( + "bytes" + "crypto/md5" + "encoding/binary" + "net" + "strconv" + + "github.com/wweir/sower/pkg/teeconn" +) + +// https://en.wikipedia.org/wiki/Domain_Name_System +const maxDomainLength = 253 + +var headSize = binary.Size(new(Head)) + +// action(>=0x80) + checksum + port + target + data +// data(HTTP, first byte < 0x7F) +type Head struct { + Cmd byte + Checksum byte + Port uint16 + TgtAddr [maxDomainLength]byte +} + +func (h *Head) Network() string { return "tcp" } +func (h *Head) String() string { + idx := bytes.IndexRune(h.TgtAddr[:], 0) + addr := string(h.TgtAddr[:idx]) + return net.JoinHostPort(addr, strconv.Itoa(int(h.Port))) +} + +type Sower struct { + password []byte +} + +func New(password string) *Sower { + return &Sower{ + password: []byte(password), + } +} + +func (s *Sower) Unwrap(conn *teeconn.Conn) net.Addr { + buf := make([]byte, headSize) + if n, err := conn.Read(buf); err != nil || n != headSize { + return nil + } + + h := &Head{} + binary.Read(bytes.NewReader(buf), binary.BigEndian, h) + switch h.Cmd { + case 0x80: + default: + return nil + } + + if h.Checksum != sumChecksum(h.TgtAddr, s.password) { + return nil + } + + return h +} + +func (s *Sower) Wrap(conn net.Conn, tgtHost string, tgtPort uint16) error { + tgtAddr := [maxDomainLength]byte{} + copy(tgtAddr[:len(tgtHost)], []byte(tgtHost)) + + return binary.Write(conn, binary.BigEndian, &Head{ + Cmd: 0x80, + Checksum: sumChecksum(tgtAddr, s.password), + Port: tgtPort, + TgtAddr: tgtAddr, + }) +} + +func sumChecksum(target [maxDomainLength]byte, password []byte) byte { + return md5.Sum(append(target[:], password...))[0] +} diff --git a/transport/transport.go b/transport/transport.go new file mode 100644 index 0000000..e466041 --- /dev/null +++ b/transport/transport.go @@ -0,0 +1,12 @@ +package transport + +import ( + "net" + + "github.com/wweir/sower/pkg/teeconn" +) + +type Transport interface { + Unwrap(conn *teeconn.Conn) net.Addr + Wrap(conn net.Conn, tgtHost string, tgtPort uint16) error +} diff --git a/transport/transport_test.go b/transport/transport_test.go new file mode 100644 index 0000000..de7206c --- /dev/null +++ b/transport/transport_test.go @@ -0,0 +1,46 @@ +package transport + +import ( + "net" + "strings" + "testing" + + "github.com/rs/zerolog/log" + "github.com/wweir/sower/pkg/teeconn" + "github.com/wweir/sower/transport/sower" + "github.com/wweir/sower/transport/trojan" +) + +func init() { + log.Logger = log.Logger.With().Caller().Logger() +} + +func testPipe(tran Transport) net.Addr { + r, w := net.Pipe() + defer r.Close() + + go func(w net.Conn) { + defer w.Close() + tran.Wrap(w, "sower", 443) + }(w) + + return tran.Unwrap(teeconn.New(r)) +} + +func Test_Transports(t *testing.T) { + if addr := testPipe(newSower()); addr == nil || strings.TrimSpace(addr.String()) != "sower:443" { + t.Errorf("test sower, unexpected address: %s", addr) + } + + if addr := testPipe(newTrojan()); addr == nil || strings.TrimSpace(addr.String()) != "sower:443" { + t.Errorf("test trojan, unexpected address: %s", addr) + } +} + +func newSower() *sower.Sower { + return sower.New("123") +} + +func newTrojan() *trojan.Trojan { + return trojan.New("123") +} diff --git a/transport/trojan.go b/transport/trojan.go deleted file mode 100644 index ed87ffa..0000000 --- a/transport/trojan.go +++ /dev/null @@ -1,122 +0,0 @@ -package transport - -import ( - "crypto/sha256" - "encoding/hex" - "io" - "net" - "strconv" - - "github.com/wweir/sower/util" - "golang.org/x/xerrors" -) - -// +-----------------------+---------+----------------+---------+----------+ -// | hex(SHA224(password)) | CRLF | Trojan Request | CRLF | Payload | -// +-----------------------+---------+----------------+---------+----------+ -// | 56 | X'0D0A' | Variable | X'0D0A' | Variable | -// +-----------------------+---------+----------------+---------+----------+ -// +-----+------+----------+----------+ -// | CMD | ATYP | DST.ADDR | DST.PORT | -// +-----+------+----------+----------+ -// | 1 | 1 | Variable | 2 | -// +-----+------+----------+----------+ -// o CMD -// o CONNECT X'01' -// o UDP :X'03' -// o ATYP -// o IP V4 : X'01' -// o domain: X'03' -// o IP V6 : X'04' - -type Trojan struct { - Password [56]byte - TrojanRequest -} -type TrojanRequest struct { - CMD uint8 - ATYP uint8 - DstAddr []byte - DstPort uint16 -} - -func (t *TrojanRequest) Addr() string { - switch t.ATYP { - case 0x01: - fallthrough - case 0x04: - return net.JoinHostPort(net.IP(t.DstAddr).String(), strconv.Itoa(int(t.DstPort))) - case 0x03: - return string(t.DstAddr) + ":" + strconv.Itoa(int(t.DstPort)) - default: - panic("invalid ATYP") - } -} - -func ParseTrojanConn(conn net.Conn, password []byte) (net.Conn, *Trojan) { - passData := sha256.Sum224(password) - passHead := hex.EncodeToString(passData[:]) - - teeConn := &util.TeeConn{Conn: conn} - defer teeConn.Stop() - - buf := make([]byte, 56+2+1+1+1) - if _, err := io.ReadFull(teeConn, buf); err != nil { - return teeConn, nil - } - if string(buf[:56]) != passHead { - return teeConn, nil - } - - t := Trojan{} - t.CMD, t.ATYP = buf[58], buf[59] - addrLen := buf[60] - switch t.ATYP { - case 0x01: //ipv4 - buf = make([]byte, net.IPv4len+2+2) - buf[0] = addrLen - if _, err := io.ReadFull(teeConn, buf[1:]); err != nil { - return teeConn, nil - } - - case 0x04: //ipv6 - buf = make([]byte, net.IPv6len+2+2) - buf[0] = addrLen - if _, err := io.ReadFull(teeConn, buf[1:]); err != nil { - return teeConn, nil - } - - case 0x03: // domain - buf = make([]byte, addrLen+2+2) - buf[0] = addrLen - if _, err := io.ReadFull(teeConn, buf); err != nil { - return teeConn, nil - } - default: - return teeConn, nil - } - - t.DstAddr = buf[:len(buf)-4] - t.DstPort = uint16(buf[len(buf)-4])<<8 + uint16(buf[len(buf)-3]) - - teeConn.Reset() - return teeConn, &t -} - -func ToTrojanConn(conn net.Conn, tgtHost string, tgtPort uint16, password []byte) (net.Conn, error) { - passData := sha256.Sum224(password) - passHead := hex.EncodeToString(passData[:]) - - buf := make([]byte, 0, 56+2+1+1+1+len(tgtHost)+2+2) - buf = append(buf, []byte(passHead)...) - buf = append(buf, '\r', '\n') - buf = append(buf, 1, 3, uint8(len(tgtHost))) - buf = append(buf, []byte(tgtHost)...) - buf = append(buf, byte(tgtPort>>8), byte(tgtPort)) - buf = append(buf, '\r', '\n') - - if n, err := conn.Write(buf); err != nil || n != len(buf) { - return conn, xerrors.Errorf("n: %d, msg: %s", n, err) - } - return conn, nil -} diff --git a/transport/trojan/trojan.go b/transport/trojan/trojan.go new file mode 100644 index 0000000..517b089 --- /dev/null +++ b/transport/trojan/trojan.go @@ -0,0 +1,182 @@ +package trojan + +import ( + "bytes" + "crypto/sha256" + "encoding/binary" + "encoding/hex" + "io" + "net" + "strconv" + + "github.com/pkg/errors" + "github.com/wweir/sower/pkg/teeconn" +) + +// +-----------------------+---------+----------------+---------+----------+ +// | hex(SHA224(password)) | CRLF | Trojan Request | CRLF | Payload | +// +-----------------------+---------+----------------+---------+----------+ +// | 56 | X'0D0A' | Variable | X'0D0A' | Variable | +// +-----------------------+---------+----------------+---------+----------+ +// +-----+------+----------+----------+ +// | CMD | ATYP | DST.ADDR | DST.PORT | +// +-----+------+----------+----------+ +// | 1 | 1 | Variable | 2 | +// +-----+------+----------+----------+ +// o CMD +// o CONNECT X'01' +// o UDP :X'03' +// o ATYP +// o IP V4 : X'01' +// o domain: X'03' +// o IP V6 : X'04' + +const headLen = 56 + 2 + 1 + 1 + +type staticHead struct { + Passwd [56]byte + CRLF [2]byte + CMD uint8 + ATYP uint8 +} + +type ipv4Addr struct { + ADDR [4]byte + PORT uint16 + CRLF [2]byte +} + +func (*ipv4Addr) Network() string { return "tcp" } +func (a *ipv4Addr) String() string { + return net.JoinHostPort(net.IP(a.ADDR[:]).String(), strconv.Itoa(int(a.PORT))) +} + +type ipv6Addr struct { + ADDR [16]byte + PORT uint16 + CRLF [2]byte +} + +func (*ipv6Addr) Network() string { return "tcp" } +func (a *ipv6Addr) String() string { + return net.JoinHostPort(net.IP(a.ADDR[:]).String(), strconv.Itoa(int(a.PORT))) +} + +type domain struct { + ADDR string + PORT uint16 + CRLF [2]byte +} + +func (*domain) Network() string { return "tcp" } +func (a *domain) String() string { + return net.JoinHostPort(a.ADDR, strconv.Itoa(int(a.PORT))) +} +func (a *domain) Fulfill(r io.Reader) error { + buf := make([]byte, 1) + if n, err := r.Read(buf); err != nil || n != 1 { + return errors.New("read domain length failed") + } + + addrLen := int(buf[0]) + buf = make([]byte, addrLen+4) + if n, err := r.Read(buf); err != nil || n != addrLen+4 { + return errors.Wrap(err, "read doamin failed") + } + + a.ADDR = string(buf[:addrLen]) + a.PORT = uint16(buf[addrLen])<<8 + uint16(buf[addrLen+1]) + return nil +} + +type Trojan struct { + headPasswd []byte + + headIPv4 []byte + headIPv6 []byte + headDomain []byte +} + +func New(password string) *Trojan { + t := &Trojan{ + headPasswd: make([]byte, 56), + } + passSum := sha256.Sum224([]byte(password)) + hex.Encode(t.headPasswd, passSum[:]) + + t.headIPv4 = append(t.headPasswd, 0x0D, 0x0A, 0x01, 0x01) + t.headIPv6 = append(t.headPasswd, 0x0D, 0x0A, 0x01, 0x04) + t.headDomain = append(t.headPasswd, 0x0D, 0x0A, 0x01, 0x03) + return t +} + +func (t *Trojan) Unwrap(conn *teeconn.Conn) net.Addr { + buf := make([]byte, headLen) + // do not use io.ReadFull to avoid hang + if n, err := conn.Read(buf); err != nil || n != headLen { + return nil + } + + head := &staticHead{} + if err := binary.Read(bytes.NewBuffer(buf), binary.BigEndian, head); err != nil { + return nil + } + + if !bytes.Equal(head.Passwd[:], []byte(t.headPasswd)) { + return nil + } + + head.CMD, head.ATYP = buf[58], buf[59] + switch head.ATYP { + case 0x01: //ipv4 + addr := &ipv4Addr{} + if err := binary.Read(conn, binary.BigEndian, addr); err != nil { + return nil + } + return addr + + case 0x04: //ipv6 + addr := &ipv6Addr{} + if err := binary.Read(conn, binary.BigEndian, addr); err != nil { + return nil + } + return addr + + case 0x03: // domain + addr := &domain{} + if err := addr.Fulfill(conn); err != nil { + return nil + } + return addr + + default: + return nil + } +} + +func (t *Trojan) Wrap(conn net.Conn, tgtHost string, tgtPort uint16) error { + var buf []byte + ip := net.ParseIP(tgtHost) + + switch { + case len(ip.To4()) != 0: + buf = make([]byte, headLen+net.IPv4len+4) + buf = append(t.headIPv4, []byte(ip.To4())...) + + case len(ip) != 0: + buf = make([]byte, headLen+net.IPv6len+4) + buf = append(t.headIPv6, []byte(ip)...) + + default: + buf = make([]byte, headLen+1+len(tgtHost)+4) + buf = append(t.headDomain, byte(len(tgtHost))) + buf = append(buf, []byte(tgtHost)...) + } + + buf = append(buf, byte(tgtPort>>8), byte(tgtPort), 0x0D, 0x0A) + + if n, err := conn.Write(buf); err != nil || n != len(buf) { + return errors.Errorf("n: %d, msg: %s", n, err) + } + return nil +} diff --git a/transport/util.go b/transport/util.go deleted file mode 100644 index 9114266..0000000 --- a/transport/util.go +++ /dev/null @@ -1,89 +0,0 @@ -package transport - -import ( - "context" - "crypto/tls" - "net" - "strconv" - - "github.com/wweir/sower/dhcp" - "github.com/wweir/util-go/log" -) - -var ( - persistDNS string - dnsAddr string - resolver = &net.Resolver{ - PreferGo: true, - Dial: func(ctx context.Context, network, address string) (net.Conn, error) { - return (&net.Dialer{}).DialContext(ctx, network, dnsAddr) - }, - } -) - -func SetDNS(err error, dnsIP string) { - if dnsIP != "" { - persistDNS = dnsIP - dnsAddr = net.JoinHostPort(dnsIP, "53") - return - } else if persistDNS != "" { - return - } - - if e, ok := err.(*net.DNSError); !ok /*nil*/ || !e.IsNotFound { - if dnsIP, err = dhcp.GetDefaultDNSServer(); err != nil { - dnsIP, err = dhcp.GetDefaultDNSServer() // retry - } - if err != nil { - log.Errorw("get dns via dhcp", "err", err, "current_dns", dnsAddr) - } else { - dnsAddr = net.JoinHostPort(dnsIP, "53") - } - } -} - -// Dial dial targetAddr with possiable proxy address -func Dial(targetAddr string, dialAddr func(domain string) (proxyAddr string, password []byte)) (net.Conn, error) { - host, port, err := net.SplitHostPort(targetAddr) - if err != nil { - return nil, err - } - - address, password := dialAddr(host) - if address == "" { - ips, err := resolver.LookupIPAddr(context.Background(), host) - if err != nil { //retry - ips, err = resolver.LookupIPAddr(context.Background(), host) - } - if err != nil { - SetDNS(err, "") - return nil, err - } - - return net.Dial("tcp", net.JoinHostPort(ips[0].String(), port)) - } - - p, err := strconv.Atoi(port) - if err != nil { - return nil, err - } - - if addr, ok := IsSocks5Schema(address); ok { - conn, err := net.Dial("tcp", addr) - if err != nil { - return nil, err - } - if conn, err = ToSocks5(conn, host, uint16(p)); err != nil { - conn.Close() - return nil, err - } - return conn, nil - } - - conn, err := tls.DialWithDialer(&net.Dialer{Resolver: resolver}, - "tcp", net.JoinHostPort(address, "443"), &tls.Config{}) - if err != nil { - return nil, err - } - return ToTrojanConn(conn, host, uint16(p), password) -} diff --git a/util/tee_conn.go b/util/tee_conn.go deleted file mode 100644 index 8220929..0000000 --- a/util/tee_conn.go +++ /dev/null @@ -1,50 +0,0 @@ -package util - -import ( - "io" - "net" -) - -type TeeConn struct { - net.Conn - buf []byte - offset int - stop bool // read - EnableWrite bool -} - -func (t *TeeConn) Reread() { - t.offset = 0 -} -func (t *TeeConn) Reset() { - t.buf = []byte{} - t.offset = 0 -} -func (t *TeeConn) Stop() { - t.offset = 0 - t.stop = true -} - -func (t *TeeConn) Read(b []byte) (n int, err error) { - length := len(t.buf) - t.offset - if length > 0 { - n = copy(b, t.buf[t.offset:]) - t.offset += n - return - } - - n, err = t.Conn.Read(b) - if !t.stop { - t.buf = append(t.buf, b[:n]...) - t.offset += n - } - return n, err -} - -func (t *TeeConn) Write(b []byte) (n int, err error) { - if t.stop || t.EnableWrite { - return t.Conn.Write(b) - } - - return 0, io.EOF -} diff --git a/util/util.go b/util/util.go index bc6f4f1..90fd5f5 100644 --- a/util/util.go +++ b/util/util.go @@ -1,11 +1,50 @@ package util -import "net" +import ( + "io" + "net" + "sync" + "sync/atomic" + "time" -func WithDefaultPort(addr string, port string) (address, host string) { - host, _, err := net.SplitHostPort(addr) - if err != nil { - return net.JoinHostPort(addr, port), addr + "github.com/pkg/errors" +) + +func RelayTo(conn net.Conn, addr string) (time.Duration, error) { + if _, _, err := net.SplitHostPort(addr); err != nil { + addr = net.JoinHostPort(addr, "80") } - return addr, host + + start := time.Now() + rc, err := net.DialTimeout("tcp", addr, 5*time.Second) + if err != nil { + return time.Since(start), errors.Wrapf(err, "dial %s", addr) + } + defer rc.Close() + + Relay(conn, rc) + return time.Since(start), nil +} + +func Relay(conn1, conn2 net.Conn) { + wg := &sync.WaitGroup{} + exitFlag := new(int32) + wg.Add(2) + go redirect(conn2, conn1, wg, exitFlag) + redirect(conn1, conn2, wg, exitFlag) + wg.Wait() +} +func redirect(dst, src net.Conn, wg *sync.WaitGroup, exitFlag *int32) { + + // io.Copy(dst, io.TeeReader(src, os.Stdout)) + io.Copy(dst, src) + + if atomic.CompareAndSwapInt32(exitFlag, 0, 1) { + // wakeup blocked goroutine + now := time.Now() + src.SetDeadline(now) + dst.SetDeadline(now) + } + + wg.Done() }