From 7d024f3bfefc5187e70bcc31df11701d83ea52af Mon Sep 17 00:00:00 2001 From: wweir Date: Sun, 18 Jul 2021 01:10:01 +0800 Subject: [PATCH] support rule based config --- .gitignore | 3 + cmd/client/main.go | 280 ++++++++++++++++------------------ cmd/client/proxy.go | 126 +++++++++++++++ cmd/server/main.go | 22 --- go.mod | 4 + go.sum | 20 ++- router/{mmdb.go => contry.go} | 2 +- router/ping.go | 9 +- router/router.go | 45 +++--- util/log.go | 40 +++++ util/{util.go => relay.go} | 0 11 files changed, 364 insertions(+), 187 deletions(-) create mode 100644 cmd/client/proxy.go rename router/{mmdb.go => contry.go} (92%) create mode 100644 util/log.go rename util/{util.go => relay.go} (100%) diff --git a/.gitignore b/.gitignore index a022dbd..f61a17b 100644 --- a/.gitignore +++ b/.gitignore @@ -17,7 +17,10 @@ *.swo .vscode .idea +**.env + +**.mmdb /cmd/client/client /client /cmd/server/server diff --git a/cmd/client/main.go b/cmd/client/main.go index 569c71f..32331ef 100644 --- a/cmd/client/main.go +++ b/cmd/client/main.go @@ -2,73 +2,83 @@ package main import ( "bufio" - "crypto/tls" + "io" "net" "net/http" + "net/url" "os" + "strconv" "strings" - "time" "github.com/cristalhq/aconfig" + "github.com/cristalhq/aconfig/aconfigdotenv" + "github.com/cristalhq/aconfig/aconfighcl" + "github.com/cristalhq/aconfig/aconfigtoml" + "github.com/cristalhq/aconfig/aconfigyaml" "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 { + Remote 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"` + DNS struct { + Enable bool `default:"true" usage:"enable DNS proxy"` + Serve string `usage:"dns server ip, default all, eg: 127.0.0.1"` + Fallback string `default:"223.5.5.5" usage:"fallback dns server"` + } + Socks5 struct { + Enable bool `default:"true" usage:"enable sock5 proxy"` + Addr string `default:":1080" usage:"socks5 listen address"` + } `flag:"socks5"` Router struct { - ProxyList []string - ProxyRefs []string - DirectList []string - DirectRefs []string - BlockList []string - BlockRefs []string + Block struct { + File string `usage:"block list file, parsed as '**.line_text'"` + Rules []string `usage:"block list rules"` + } + Direct struct { + File string `usage:"direct list file, parsed as '**.line_text'"` + Rules []string `usage:"direct list rules"` + } + Proxy struct { + File string `usage:"proxy list file, parsed as '**.line_text'"` + Rules []string `usage:"proxy list rules"` + } + + Country struct { + MMDB string `usage:"mmdb file"` + File string `usage:"CIDR block list file"` + Rules []string `usage:"CIDR list rules"` + } } }{} ) 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 + if err := aconfig.LoaderFor(&conf, aconfig.Config{ + FileFlag: "conf", + Files: []string{".env", + "config.yml", "config.yaml", "config.json", "config.toml", "config.hcl"}, + FileDecoders: map[string]aconfig.FileDecoder{ + ".env": aconfigdotenv.New(), + ".yml": aconfigyaml.New(), + ".yaml": aconfigyaml.New(), + ".toml": aconfigtoml.New(), + ".hcl": aconfighcl.New(), + ".tf": aconfighcl.New(), }, - }).With().Timestamp().Caller().Logger() - - if err := aconfig.LoaderFor(&conf, aconfig.Config{}).Load(); err != nil { + }).Load(); err != nil { log.Fatal().Err(err). + Interface("conf", conf). Msg("Load config") } @@ -80,140 +90,120 @@ func init() { } func main() { - - r := router.NewRouter(genProxyDial()) + proxtDial := GenProxyDial(conf.Remote.Type, conf.Remote.Addr, conf.Remote.Password) + r := router.NewRouter(conf.DNS.Fallback, conf.Router.Country.MMDB, proxtDial, + append(conf.Router.Block.Rules, parseRuleLines(proxtDial, conf.Router.Block.File, "**.")...), + append(conf.Router.Direct.Rules, parseRuleLines(proxtDial, conf.Router.Direct.File, "**.")...), + append(conf.Router.Proxy.Rules, parseRuleLines(proxtDial, conf.Router.Proxy.File, "**.")...), + append(conf.Router.Country.Rules, parseRuleLines(proxtDial, conf.Router.Country.File, "")...), + ) go func() { - lnHTTP, err := net.Listen("tcp", net.JoinHostPort(conf.DNSServeIP, "80")) + if !conf.DNS.Enable { + log.Info().Msg("DNS proxy disabled") + return + } + + lnHTTP, err := net.Listen("tcp", net.JoinHostPort(conf.DNS.Serve, "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")) + lnHTTPS, err := net.Listen("tcp", net.JoinHostPort(conf.DNS.Serve, "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.Info().Msg("DNS proxy started") + if err := dns.ListenAndServe(conf.DNS.Serve, "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) + go func() { + if !conf.Socks5.Enable { + log.Info().Msg("SOCKS5 proxy disabled") + return + } + + ln, err := net.Listen("tcp", conf.Socks5.Addr) + if err != nil { + log.Fatal().Err(err).Msg("listen port") + } + log.Info().Msgf("SOCKS5 proxy listening on %s", conf.Socks5.Addr) + 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) +func parseRuleLines(proxyDial router.ProxyDialFn, file, linePrefix string) []string { + var r io.Reader + if _, err := url.Parse(file); err == nil { + client := &http.Client{ + Transport: &http.Transport{ + Dial: func(network, addr string) (net.Conn, error) { + domain, port, _ := net.SplitHostPort(addr) + p, _ := strconv.Atoi(port) + return proxyDial("tcp", domain, uint16(p)) + }, + }, } - c, err := tls.Dial("tcp", proxyAddr, tlsCfg) + resp, err := client.Get(file) if err != nil { - return nil, err + log.Error().Err(err). + Str("file", file). + Msg("proxy read response") + return nil + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + log.Error(). + Int("status", resp.StatusCode). + Str("file", file). + Msg("proxy response status") + return nil } - if err := proxy.Wrap(c, host, port); err != nil { - return nil, err + r = resp.Body + + } else { + f, err := os.Open(file) + if err != nil { + log.Error().Err(err). + Str("file", file). + Msg("open file") + return nil + } + defer f.Close() + + r = f + } + + var lines []string + br := bufio.NewReader(r) + for { + line, _, err := br.ReadLine() + if err == io.EOF { + break + } else if err != nil { + log.Error().Err(err). + Str("file", file). + Msg("read line") + return nil } - return c, nil + if strings.TrimSpace(string(line)) == "" { + continue + } + + // use line content as suffix + lines = append(lines, linePrefix+string(line)) } -} -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") + return lines } diff --git a/cmd/client/proxy.go b/cmd/client/proxy.go new file mode 100644 index 0000000..2c96f6b --- /dev/null +++ b/cmd/client/proxy.go @@ -0,0 +1,126 @@ +package main + +import ( + "bufio" + "crypto/tls" + "net" + "net/http" + "time" + + "github.com/pkg/errors" + "github.com/rs/zerolog/log" + "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" +) + +func GenProxyDial(proxyType, proxyHost, proxyPassword string) router.ProxyDialFn { + var ( + proxyAddr = net.JoinHostPort(proxyHost, "443") + tlsCfg = &tls.Config{} + proxy transport.Transport + ) + switch conf.Remote.Type { + case "sower": + proxy = sower.New(conf.Remote.Password) + case "trojan": + proxy = trojan.New(conf.Remote.Password) + default: + log.Fatal(). + Str("type", conf.Remote.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/server/main.go b/cmd/server/main.go index adedd26..d6b9da4 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -6,13 +6,9 @@ import ( "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" @@ -37,24 +33,6 @@ var ( ) 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") } diff --git a/go.mod b/go.mod index ca6353b..64bc400 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,10 @@ go 1.16 require ( github.com/cristalhq/aconfig v0.16.2 + github.com/cristalhq/aconfig/aconfigdotenv v0.16.1 + github.com/cristalhq/aconfig/aconfighcl v0.16.1 + github.com/cristalhq/aconfig/aconfigtoml v0.16.1 + github.com/cristalhq/aconfig/aconfigyaml v0.16.1 github.com/krolaw/dhcp4 v0.0.0-20190909130307-a50d88189771 github.com/libp2p/go-reuseport v0.0.2 github.com/miekg/dns v1.1.43 diff --git a/go.sum b/go.sum index 1592f95..42147a0 100644 --- a/go.sum +++ b/go.sum @@ -1,9 +1,25 @@ +github.com/BurntSushi/toml v0.3.1 h1:WXkYYl6Yr3qBf1K79EBnL4mak0OimBfB0XUf9Vl28OQ= +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/go.mod h1:NXaRp+1e6bkO4dJn+wZ71xyaihMDYPtCSvEhMTm/H3E= github.com/cristalhq/aconfig v0.16.2 h1:cN+y3rtHyK9ar7NRiY4E+USs07P2qnZnGeZfNzXoV2g= github.com/cristalhq/aconfig v0.16.2/go.mod h1:NXaRp+1e6bkO4dJn+wZ71xyaihMDYPtCSvEhMTm/H3E= -github.com/davecgh/go-spew v1.1.0 h1:ZDRjVQ15GmhC3fiQ8ni8+OwkZQO4DARzQgrnXU1Liz8= +github.com/cristalhq/aconfig/aconfigdotenv v0.16.1 h1:nsybrhIghZ2y0hQAijDlqih7KYgs03UujQ+GcyButLc= +github.com/cristalhq/aconfig/aconfigdotenv v0.16.1/go.mod h1:VQMzt9eS0Zlxwsuf/zkJuk1oZLm72cjpAkEkMrpbOm0= +github.com/cristalhq/aconfig/aconfighcl v0.16.1 h1:RiziovFKOwB6+H1LCD7DXxrFgoKDEJxxYLSsAOsj+vQ= +github.com/cristalhq/aconfig/aconfighcl v0.16.1/go.mod h1:YvZyVGtty/7yIdyOPPl68u7ENaUMq6ybZAph+n3g4II= +github.com/cristalhq/aconfig/aconfigtoml v0.16.1 h1:q9DJfJqdLvZj7vL3bOIR6QM0ySpA0i6X8CZ8FNob6nc= +github.com/cristalhq/aconfig/aconfigtoml v0.16.1/go.mod h1:VfRDnJBq09TZXiQdnnd4CBsNyj4s7AEDJlCxDSDt5DI= +github.com/cristalhq/aconfig/aconfigyaml v0.16.1 h1:ghmzWFolCuja6BzBvHkNAn2Qsewu6xqdAXQyWnXjvZw= +github.com/cristalhq/aconfig/aconfigyaml v0.16.1/go.mod h1:UgM0LkO4TKbC0s/oUP2p5POM6QiQwNvOFXostf6con0= 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/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA= +github.com/hashicorp/hcl v1.0.0 h1:0Anlzjpi4vEasTeNFn2mLJgTSwt0+6sfsiTG8qcWGx4= +github.com/hashicorp/hcl v1.0.0/go.mod h1:E5yfLk+7swimpb2L/Alb/PJmXilQ/rhwaUYs4T20WEQ= +github.com/joho/godotenv v1.3.0 h1:Zjp+RcGpHhGlrMbJzXTrZZPrWj+1vfm90La1wgB6Bhc= +github.com/joho/godotenv v1.3.0/go.mod h1:7hK45KPybAkOC6peb+G5yklZfMxEjkZhHbwpqxOKXbg= 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= github.com/libp2p/go-reuseport v0.0.2 h1:XSG94b1FJfGA01BUrT82imejHQyTxO4jEWqheyCXYvU= @@ -67,5 +83,7 @@ golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8T gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v2 v2.3.0 h1:clyUAQHOM3G0M3f5vQj7LuJrETvjVot3Z5el9nffUtU= +gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= 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/router/mmdb.go b/router/contry.go similarity index 92% rename from router/mmdb.go rename to router/contry.go index e6b2923..0bcc1c3 100644 --- a/router/mmdb.go +++ b/router/contry.go @@ -15,7 +15,7 @@ func (r *Router) localSite(domain string) bool { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() - ips, err := r.mmdb.Resolver.LookupIP(ctx, "ip", domain) + ips, err := net.DefaultResolver.LookupIP(ctx, "ip", domain) if err != nil || len(ips) == 0 { log.Warn().Err(err). Str("domain", domain). diff --git a/router/ping.go b/router/ping.go index ee803a9..9ff0c1c 100644 --- a/router/ping.go +++ b/router/ping.go @@ -10,7 +10,14 @@ var pingClient = http.Client{ Timeout: 2 * time.Second, } -func (r *Router) isAccess(domain string) bool { +func (r *Router) isAccess(domain string, port uint16) bool { + switch port { + case 80: + case 443: + default: + return false + } + p := &ping{} r.cache.Remember(p, domain) return p.isAccess diff --git a/router/router.go b/router/router.go index ed01572..1580897 100644 --- a/router/router.go +++ b/router/router.go @@ -14,53 +14,64 @@ import ( "github.com/wweir/sower/util" ) +type ProxyDialFn func(network, host string, port uint16) (net.Conn, error) type Router struct { blockRule *util.Node directRule *util.Node proxyRule *util.Node - - ProxyDial func(network, host string, port uint16) (net.Conn, error) - cache *mem.Cache + ProxyDial ProxyDialFn + cache *mem.Cache dns struct { + fallbackDNS string dns.Client connCh chan *dns.Conn } mmdb struct { *geoip2.Reader - *net.Resolver - cidrs []*net.IPNet } } -func NewRouter(proxyDial func(network, host string, port uint16) (net.Conn, error)) *Router { +func NewRouter(fallbackDNS, mmdbFile string, proxyDial ProxyDialFn, + blockList, directList, proxyList, directCIDRs []string) *Router { + r := Router{ - blockRule: util.NewNodeFromRules(), - directRule: util.NewNodeFromRules(), - proxyRule: util.NewNodeFromRules("google.*"), + blockRule: util.NewNodeFromRules(blockList...), + directRule: util.NewNodeFromRules(directList...), + proxyRule: util.NewNodeFromRules(proxyList...), ProxyDial: proxyDial, cache: mem.New(time.Hour), // TODO: config } + r.dns.fallbackDNS = fallbackDNS r.dns.connCh = make(chan *dns.Conn, 1) go r.dialDNSConn() + var err error + r.mmdb.Reader, err = geoip2.Open(mmdbFile) + log.Err(err).Str("file", mmdbFile).Msg("open geoip2 db") + r.mmdb.cidrs = make([]*net.IPNet, 0, len(directCIDRs)) + for _, cidr := range directCIDRs { + _, ipnet, err := net.ParseCIDR(cidr) + if err != nil { + log.Error().Err(err).Msg("Failed to parse CIDR") + } + r.mmdb.cidrs = append(r.mmdb.cidrs, ipnet) + } return &r } func (r *Router) dialDNSConn() { for { server, err := dhcp.GetDNSServer() - if err != nil { - time.Sleep(time.Second) - continue + if server == "" { + server = r.dns.fallbackDNS } - - log.Info(). - Str("ip", server). - Msg("get upstream dns server") + log.Err(err). + Str("DNS", server). + Msg("get DNS server") for { conn, err := dns.DialTimeout("udp", net.JoinHostPort(server, "53"), time.Second) @@ -89,7 +100,7 @@ func (r *Router) RouteHandle(conn net.Conn, domain string, port uint16) error { case r.localSite(domain): return r.DirectHandle(conn, addr) - case r.isAccess(domain): + case r.isAccess(domain, port): return r.DirectHandle(conn, addr) case port == 80: return r.DirectHandle(conn, addr) diff --git a/util/log.go b/util/log.go new file mode 100644 index 0000000..e11ccbe --- /dev/null +++ b/util/log.go @@ -0,0 +1,40 @@ +package util + +import ( + "os" + "strings" + "time" + + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/rs/zerolog/pkgerrors" +) + +var StructLogger = zerolog.New(os.Stdout). + With().Caller().Timestamp().Logger() + +var ConsoleLogger = zerolog.New(zerolog.ConsoleWriter{ + Out: os.Stdout, + 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() + +func init() { + zerolog.ErrorStackMarshaler = func(err error) interface{} { + return pkgerrors.MarshalStack(err) + } + if fi, _ := os.Stdout.Stat(); (fi.Mode() & os.ModeCharDevice) == 0 { + log.Logger = StructLogger + } else { + log.Logger = ConsoleLogger + } +} diff --git a/util/util.go b/util/relay.go similarity index 100% rename from util/util.go rename to util/relay.go