From 1356245aff08138bcd00a8b6d352e50b9c0ab0b7 Mon Sep 17 00:00:00 2001 From: iacore Date: Wed, 28 Aug 2024 10:20:17 +0000 Subject: [PATCH] separate out error handler --- handler/error_handler.go | 66 ++++++++++++++++++++++++++++++++++++++++ main.go | 63 +++----------------------------------- session/aux.go | 2 +- 3 files changed, 71 insertions(+), 60 deletions(-) create mode 100644 handler/error_handler.go diff --git a/handler/error_handler.go b/handler/error_handler.go new file mode 100644 index 0000000..cae6006 --- /dev/null +++ b/handler/error_handler.go @@ -0,0 +1,66 @@ +package handler + +import ( + "bytes" + "log" + "maps" + "net/http" + "net/http/httptest" + "slices" + + "codeberg.org/vnpower/pixivfe/v2/routes" +) + +type UserContext struct { + Err error + StatusCode int +} + +type userContextKey struct{} + +var UserContextKey = userContextKey{} + +func GetUserContext(r *http.Request) *UserContext { + return r.Context().Value(UserContextKey).(*UserContext) +} + +func CatchError(handler func(w http.ResponseWriter, r *http.Request) error) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + header_backup := http.Header{} + for k, v := range w.Header() { + header_backup[k] = slices.Clone(v) + } + recorder := httptest.ResponseRecorder{ + HeaderMap: w.Header(), + Body: new(bytes.Buffer), + Code: 200, + } + err := handler(&recorder, r) + if err != nil { + clear(header_backup) + maps.Copy(w.Header(), header_backup) + GetUserContext(r).Err = err + } else { + _, _ = recorder.Body.WriteTo(w) + w.WriteHeader(recorder.Code) + } + } +} + +func ErrorHandler(w http.ResponseWriter, r *http.Request) { // error handler + err := GetUserContext(r).Err + + if err != nil { + log.Printf("Internal Server Error: %s", err) + code := GetUserContext(r).StatusCode + if code == 0 { + code = http.StatusInternalServerError + } + w.WriteHeader(code) + // Send custom error page + err = routes.ErrorPage(w, r, err) + if err != nil { + log.Printf("Error rendering error route: %s", err) + } + } +} diff --git a/main.go b/main.go index e30be32..189072b 100644 --- a/main.go +++ b/main.go @@ -1,19 +1,15 @@ package main import ( - "bytes" "context" "errors" "fmt" "log" - "maps" "net" "net/http" - "net/http/httptest" "os" "os/exec" "runtime" - "slices" "strings" "sync" "syscall" @@ -24,6 +20,7 @@ import ( "codeberg.org/vnpower/pixivfe/v2/config" "codeberg.org/vnpower/pixivfe/v2/core" + . "codeberg.org/vnpower/pixivfe/v2/handler" "codeberg.org/vnpower/pixivfe/v2/routes" "codeberg.org/vnpower/pixivfe/v2/template" ) @@ -46,19 +43,6 @@ func CanRequestSkipLogger(r *http.Request) bool { strings.HasPrefix(path, "/proxy/i.pximg.net/") } -type UserContext struct { - err error - statusCode int -} - -type userContextKey struct{} - -var UserContextKey = userContextKey{} - -func GetUserContext(r *http.Request) *UserContext { - return r.Context().Value(UserContextKey).(*UserContext) -} - // Todo: Should we put middlewares in a separate file? // IPRateLimiter represents an IP rate limiter. type IPRateLimiter struct { @@ -108,7 +92,7 @@ func MiddlewareChain(handler http.Handler) http.Handler { if !limiter.Allow(ip) { CatchError(func(w http.ResponseWriter, r *http.Request) error { - GetUserContext(r).statusCode = http.StatusTooManyRequests + GetUserContext(r).StatusCode = http.StatusTooManyRequests return errors.New("Too many requests") })(w, r) } else { @@ -152,23 +136,7 @@ func main() { router.ServeHTTP(w, r) } - { // error handler - err := GetUserContext(r).err - - if err != nil { - log.Printf("Internal Server Error: %s", err) - code := GetUserContext(r).statusCode - if code == 0 { - code = http.StatusInternalServerError - } - w.WriteHeader(code) - // Send custom error page - err = routes.ErrorPage(w, r, err) - if err != nil { - log.Printf("Error rendering error route: %s", err) - } - } - } + ErrorHandler(w, r) end_time := time.Now() @@ -179,7 +147,7 @@ func main() { method := r.Method path := r.URL.Path status := w.statusCode - err := GetUserContext(r).err + err := GetUserContext(r).Err log.Printf("%v +%v %v %v %v %v %v", time, latency, ip, method, path, status, err) } @@ -292,29 +260,6 @@ func defineRoutes() *mux.Router { return router } -func CatchError(handler func(w http.ResponseWriter, r *http.Request) error) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - header_backup := http.Header{} - for k, v := range w.Header() { - header_backup[k] = slices.Clone(v) - } - recorder := httptest.ResponseRecorder{ - HeaderMap: w.Header(), - Body: new(bytes.Buffer), - Code: 200, - } - err := handler(&recorder, r) - if err != nil { - clear(header_backup) - maps.Copy(w.Header(), header_backup) - GetUserContext(r).err = err - } else { - _, _ = recorder.Body.WriteTo(w) - w.WriteHeader(recorder.Code) - } - } -} - type ResponseWriterInterceptStatus struct { statusCode int http.ResponseWriter diff --git a/session/aux.go b/session/aux.go index 1d6f439..f666d39 100644 --- a/session/aux.go +++ b/session/aux.go @@ -6,7 +6,7 @@ import ( "net/url" "strings" - config "codeberg.org/vnpower/pixivfe/v2/config" + "codeberg.org/vnpower/pixivfe/v2/config" ) func GetPixivToken(r *http.Request) string {