diff --git a/cmd/sower/socks5.go b/cmd/sower/socks5.go index 422790c..25df5e4 100644 --- a/cmd/sower/socks5.go +++ b/cmd/sower/socks5.go @@ -4,7 +4,6 @@ import ( "encoding/binary" "io" "net" - "time" "github.com/rs/zerolog/log" "github.com/wweir/sower/router" @@ -18,7 +17,6 @@ func ServeSocks5(ln net.Listener, r *router.Router) { } go ServeSocks5(ln, r) defer conn.Close() - start := time.Now() { auth := new(socks5AuthReq) @@ -68,19 +66,7 @@ func ServeSocks5(ln net.Listener, r *router.Router) { } host, port := addr.Addr() - route, err := r.RouteHandle(conn, host, port) - switch route { - case router.RouteProxy: - case router.RouteDefault: - default: - return - } - log.Err(err). - Str("host", host). - Uint16("port", port). - Str("route", route). - Dur("spend", time.Since(start)). - Msg("serve socsk5") + r.RouteHandle(conn, host, port) } /******************* https://tools.ietf.org/html/rfc1928 *******************/ diff --git a/pkg/deferlog/deferlog.go b/pkg/deferlog/deferlog.go new file mode 100644 index 0000000..abcaa0d --- /dev/null +++ b/pkg/deferlog/deferlog.go @@ -0,0 +1,131 @@ +// Package deferlog is a copy of zerolog/log/log.go for log in defer +package deferlog + +import ( + "context" + "fmt" + "io" + "os" + + "github.com/rs/zerolog" +) + +// Logger is the global logger. +var Logger = zerolog.New(os.Stderr).With().Timestamp().Logger() + +// Output duplicates the global logger and sets w as its output. +func Output(w io.Writer) zerolog.Logger { + return Logger.Output(w) +} + +// With creates a child logger with the field added to its context. +func With() zerolog.Context { + return Logger.With() +} + +// Level creates a child logger with the minimum accepted level set to level. +func Level(level zerolog.Level) zerolog.Logger { + return Logger.Level(level) +} + +// Sample returns a logger with the s sampler. +func Sample(s zerolog.Sampler) zerolog.Logger { + return Logger.Sample(s) +} + +// Hook returns a logger with the h Hook. +func Hook(h zerolog.Hook) zerolog.Logger { + return Logger.Hook(h) +} + +// Err starts a new message with error level with err as a field if not nil or +// with info level if err is nil. +// +// You must call Msg on the returned event in order to send the event. +func Err(err error) *zerolog.Event { + return Logger.Err(err) +} + +// Trace starts a new message with trace level. +// +// You must call Msg on the returned event in order to send the event. +func Trace() *zerolog.Event { + return Logger.Trace() +} + +// Debug starts a new message with debug level. +// +// You must call Msg on the returned event in order to send the event. +func Debug() *zerolog.Event { + return Logger.Debug() +} + +// Info starts a new message with info level. +// +// You must call Msg on the returned event in order to send the event. +func Info() *zerolog.Event { + return Logger.Info() +} + +// Warn starts a new message with warn level. +// +// You must call Msg on the returned event in order to send the event. +func Warn() *zerolog.Event { + return Logger.Warn() +} + +// Error starts a new message with error level. +// +// You must call Msg on the returned event in order to send the event. +func Error() *zerolog.Event { + return Logger.Error() +} + +// Fatal starts a new message with fatal level. The os.Exit(1) function +// is called by the Msg method. +// +// You must call Msg on the returned event in order to send the event. +func Fatal() *zerolog.Event { + return Logger.Fatal() +} + +// Panic starts a new message with panic level. The message is also sent +// to the panic function. +// +// You must call Msg on the returned event in order to send the event. +func Panic() *zerolog.Event { + return Logger.Panic() +} + +// WithLevel starts a new message with level. +// +// You must call Msg on the returned event in order to send the event. +func WithLevel(level zerolog.Level) *zerolog.Event { + return Logger.WithLevel(level) +} + +// Log starts a new message with no level. Setting zerolog.GlobalLevel to +// zerolog.Disabled will still disable events produced by this method. +// +// You must call Msg on the returned event in order to send the event. +func Log() *zerolog.Event { + return Logger.Log() +} + +// Print sends a log event using debug level and no extra field. +// Arguments are handled in the manner of fmt.Print. +func Print(v ...interface{}) { + Logger.Debug().CallerSkipFrame(1).Msg(fmt.Sprint(v...)) +} + +// Printf sends a log event using debug level and no extra field. +// Arguments are handled in the manner of fmt.Printf. +func Printf(format string, v ...interface{}) { + Logger.Debug().CallerSkipFrame(1).Msgf(format, v...) +} + +// Ctx returns the Logger associated with the ctx. If no logger +// is associated, a disabled logger is returned. +func Ctx(ctx context.Context) *zerolog.Logger { + return zerolog.Ctx(ctx) +} diff --git a/pkg/deferlog/enhance.go b/pkg/deferlog/enhance.go new file mode 100644 index 0000000..49d1602 --- /dev/null +++ b/pkg/deferlog/enhance.go @@ -0,0 +1,88 @@ +package deferlog + +import ( + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" +) + +func InfoWarn(err error) *zerolog.Event { + if err != nil { + return Logger.Warn().Err(err) + } + + return Logger.Info() +} +func InfoFatal(err error) *zerolog.Event { + if err != nil { + return Logger.Fatal().Err(err) + } + + return Logger.Info() +} + +func DebugWarn(err error) *zerolog.Event { + if err != nil { + return Logger.Warn().Err(err) + } + + return Logger.Debug() +} +func DebugError(err error) *zerolog.Event { + if err != nil { + return Logger.Error().Err(err) + } + + return Logger.Debug() +} +func DebugFatal(err error) *zerolog.Event { + if err != nil { + return Logger.Fatal().Err(err) + } + + return Logger.Debug() +} + +type StdLogger struct { + *zerolog.Logger +} + +var Std = StdLogger{ + Logger: &log.Logger, +} + +func (std *StdLogger) InfoWarn(err error) *zerolog.Event { + if err != nil { + return Std.Logger.Warn().Err(err) + } + + return Std.Logger.Info() +} +func (std *StdLogger) InfoFatal(err error) *zerolog.Event { + if err != nil { + return Std.Logger.Fatal().Err(err) + } + + return Std.Logger.Info() +} + +func (std *StdLogger) DebugWarn(err error) *zerolog.Event { + if err != nil { + return Std.Logger.Warn().Err(err) + } + + return Std.Logger.Debug() +} +func (std *StdLogger) DebugError(err error) *zerolog.Event { + if err != nil { + return Std.Logger.Error().Err(err) + } + + return Std.Logger.Debug() +} +func (std *StdLogger) DebugFatal(err error) *zerolog.Event { + if err != nil { + return Std.Logger.Fatal().Err(err) + } + + return Std.Logger.Debug() +} diff --git a/util/log.go b/pkg/deferlog/init.go similarity index 57% rename from util/log.go rename to pkg/deferlog/init.go index e11ccbe..2877a0a 100644 --- a/util/log.go +++ b/pkg/deferlog/init.go @@ -1,7 +1,8 @@ -package util +package deferlog import ( "os" + "strconv" "strings" "time" @@ -32,9 +33,19 @@ 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 + + if ok, _ := strconv.ParseBool(os.Getenv("DEBUG")); ok { + SetDefaultLogger(ConsoleLogger, 1, zerolog.DebugLevel) + + } else if fi, _ := os.Stdout.Stat(); (fi.Mode() & os.ModeCharDevice) == 0 { + SetDefaultLogger(StructLogger, 1, zerolog.InfoLevel) + } else { - log.Logger = ConsoleLogger + SetDefaultLogger(ConsoleLogger, 1, zerolog.InfoLevel) } } + +func SetDefaultLogger(logger zerolog.Logger, deferSkip int, logLevel zerolog.Level) { + log.Logger = logger.Level(logLevel) + Logger = logger.With().CallerWithSkipFrameCount(deferSkip + 2).Logger().Level(logLevel) +} diff --git a/router/dns.go b/router/dns.go index 48452d1..3ae47af 100644 --- a/router/dns.go +++ b/router/dns.go @@ -5,7 +5,7 @@ import ( "time" "github.com/miekg/dns" - "github.com/rs/zerolog/log" + "github.com/wweir/sower/pkg/deferlog" ) func (r *Router) ServeDNS(w dns.ResponseWriter, req *dns.Msg) { @@ -36,12 +36,12 @@ func (r *Router) ServeDNS(w dns.ResponseWriter, req *dns.Msg) { 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") + deferlog.Std.DebugWarn(err). + Dur("rtt", rtt). + Str("domain", domain). + Msg("exchange dns record") + if err != nil { conn.Close() w.WriteMsg(r.dnsFail(req, dns.RcodeServerFailure)) return diff --git a/router/router.go b/router/router.go index 5e7e265..f089f5a 100644 --- a/router/router.go +++ b/router/router.go @@ -9,20 +9,12 @@ import ( geoip2 "github.com/oschwald/geoip2-golang" "github.com/pkg/errors" "github.com/rs/zerolog/log" + "github.com/wweir/sower/pkg/deferlog" "github.com/wweir/sower/pkg/dhcp" "github.com/wweir/sower/pkg/mem" "github.com/wweir/sower/util" ) -const ( - RouteBlock = "block" - RouteDirect = "direct" - RouteProxy = "proxy" - RouteLocal = "local" - RouteAccess = "access" - RouteDefault = "default" -) - type ProxyDialFn func(network, host string, port uint16) (net.Conn, error) type Router struct { blockRule *util.Node @@ -101,25 +93,34 @@ func (r *Router) dialDNSConn() { } } -func (r *Router) RouteHandle(conn net.Conn, domain string, port uint16) (string, error) { +func (r *Router) RouteHandle(conn net.Conn, domain string, port uint16) (err error) { + start := time.Now() + defer func() { + deferlog.DebugWarn(err). + Str("domain", domain). + Uint16("port", port). + Dur("spend", time.Since(start)). + Msg("RouteHandle") + }() + addr := net.JoinHostPort(domain, strconv.FormatUint(uint64(port), 10)) switch { case r.blockRule.Match(domain): - return RouteBlock, nil + return nil case r.directRule.Match(domain): - return RouteDirect, r.DirectHandle(conn, addr) + return r.DirectHandle(conn, addr) case r.proxyRule.Match(domain): - return RouteProxy, r.ProxyHandle(conn, domain, port) + return r.ProxyHandle(conn, domain, port) case r.localSite(domain): - return RouteLocal, r.DirectHandle(conn, addr) + return r.DirectHandle(conn, addr) case r.isAccess(domain, port): - return RouteAccess, r.DirectHandle(conn, addr) + return r.DirectHandle(conn, addr) default: - return RouteDefault, r.ProxyHandle(conn, domain, port) + return r.ProxyHandle(conn, domain, port) } } diff --git a/util/relay.go b/util/relay.go index 90fd5f5..4df33e8 100644 --- a/util/relay.go +++ b/util/relay.go @@ -10,7 +10,7 @@ import ( "github.com/pkg/errors" ) -func RelayTo(conn net.Conn, addr string) (time.Duration, error) { +func RelayTo(conn net.Conn, addr string) (dur time.Duration, err error) { if _, _, err := net.SplitHostPort(addr); err != nil { addr = net.JoinHostPort(addr, "80") }