diff --git a/server/server.go b/server/server.go index 11cf7c82d8506a6de45dde5019cfe60214cbec2b..887c89fc8ebf593d9275ce988558ee6b1c9f1141 100644 --- a/server/server.go +++ b/server/server.go @@ -193,8 +193,16 @@ server.router.Use(middleware.Logger) server.router.Use(middleware.Timeout(timeout)) server.router.Use(func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var err error + addr := r.RemoteAddr + if net.ParseIP(addr) == nil { + addr, _, err = net.SplitHostPort(addr) + if err != nil { + panic(fmt.Errorf("Invalid remote address: %s", r.RemoteAddr)) + } + } ctx := context.WithValue(r.Context(), serverCtxKey, server) - ctx = context.WithValue(ctx, remoteAddrCtxKey, r.RemoteAddr) + ctx = context.WithValue(ctx, remoteAddrCtxKey, addr) r = r.WithContext(ctx) next.ServeHTTP(w, r) }) @@ -203,6 +211,8 @@ server.WithQueues(server.email.Queue) return server } +// RemoteAddr returns the remote address for this context. It is guaranteed to +// be valid input for `net.ParseIP()`. func RemoteAddr(ctx context.Context) string { raw, ok := ctx.Value(remoteAddrCtxKey).(string) if !ok {