From 39207677b52c049cf8db7bc7f4d37d5034e597b3 Mon Sep 17 00:00:00 2001 From: Trong Huu Nguyen Date: Fri, 24 Jan 2025 09:13:59 +0100 Subject: [PATCH] feat(middleware/logentry): add fields for sec-fetch headers --- internal/http/request.go | 43 +++++++++++++++++++++++++++++++++++ pkg/handler/reverseproxy.go | 45 ++++--------------------------------- pkg/middleware/logentry.go | 45 +++++++++++++++++++++---------------- 3 files changed, 73 insertions(+), 60 deletions(-) create mode 100644 internal/http/request.go diff --git a/internal/http/request.go b/internal/http/request.go new file mode 100644 index 0000000..3c56aba --- /dev/null +++ b/internal/http/request.go @@ -0,0 +1,43 @@ +package http + +import ( + "net/http" + "strings" +) + +func IsNavigationRequest(r *http.Request) bool { + // we assume that navigation requests are always GET requests + if r.Method != http.MethodGet { + return false + } + + // check for top-level navigation requests + mode := r.Header.Get("Sec-Fetch-Mode") + dest := r.Header.Get("Sec-Fetch-Dest") + if mode != "" && dest != "" { + return mode == "navigate" && dest == "document" + } + + // fallback if browser doesn't support fetch metadata + return Accepts(r, "text/html") +} + +func Accepts(r *http.Request, accepted ...string) bool { + // iterate over all Accept headers + for _, header := range r.Header.Values("Accept") { + // iterate over all comma-separated values in a single Accept header + for _, v := range strings.Split(header, ",") { + v = strings.ToLower(v) + v = strings.TrimSpace(v) + v = strings.Split(v, ";")[0] + + for _, accept := range accepted { + if v == accept { + return true + } + } + } + } + + return false +} diff --git a/pkg/handler/reverseproxy.go b/pkg/handler/reverseproxy.go index 769494f..9012744 100644 --- a/pkg/handler/reverseproxy.go +++ b/pkg/handler/reverseproxy.go @@ -7,10 +7,10 @@ import ( "net/http" "net/http/httputil" urllib "net/url" - "strings" "github.com/sirupsen/logrus" + httpinternal "github.com/nais/wonderwall/internal/http" "github.com/nais/wonderwall/pkg/handler/acr" "github.com/nais/wonderwall/pkg/handler/autologin" mw "github.com/nais/wonderwall/pkg/middleware" @@ -166,7 +166,7 @@ func handleAutologin(src ReverseProxySource, w http.ResponseWriter, r *http.Requ return loginURL } - if isNavigationRequest(r) { + if httpinternal.IsNavigationRequest(r) { target := r.URL.String() location := loginURL(target, "navigation request detected; redirecting to login...") http.Redirect(w, r, location, http.StatusFound) @@ -183,7 +183,7 @@ func handleAutologin(src ReverseProxySource, w http.ResponseWriter, r *http.Requ w.Header().Set("Location", location) w.WriteHeader(http.StatusUnauthorized) - if accepts(r, "*/*", "application/json") { + if httpinternal.Accepts(r, "*/*", "application/json") { w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{"error": "unauthenticated, please log in"}`)) } else { @@ -194,50 +194,13 @@ func handleAutologin(src ReverseProxySource, w http.ResponseWriter, r *http.Requ func isRelevantAccessLog(r *http.Request) bool { if r.Method == http.MethodGet { // only log GET requests that are navigation requests - return isNavigationRequest(r) + return httpinternal.IsNavigationRequest(r) } // all other methods are relevant return true } -func isNavigationRequest(r *http.Request) bool { - // we assume that navigation requests are always GET requests - if r.Method != http.MethodGet { - return false - } - - // check for top-level navigation requests - mode := r.Header.Get("Sec-Fetch-Mode") - dest := r.Header.Get("Sec-Fetch-Dest") - if mode != "" && dest != "" { - return mode == "navigate" && dest == "document" - } - - // fallback if browser doesn't support fetch metadata - return accepts(r, "text/html") -} - -func accepts(r *http.Request, accepted ...string) bool { - // iterate over all Accept headers - for _, header := range r.Header.Values("Accept") { - // iterate over all comma-separated values in a single Accept header - for _, v := range strings.Split(header, ",") { - v = strings.ToLower(v) - v = strings.TrimSpace(v) - v = strings.Split(v, ";")[0] - - for _, accept := range accepted { - if v == accept { - return true - } - } - } - } - - return false -} - type logrusErrorWriter struct{} func (w logrusErrorWriter) Write(p []byte) (n int, err error) { diff --git a/pkg/middleware/logentry.go b/pkg/middleware/logentry.go index 22bdf13..b5676cf 100644 --- a/pkg/middleware/logentry.go +++ b/pkg/middleware/logentry.go @@ -8,6 +8,7 @@ import ( "time" "github.com/go-chi/chi/v5/middleware" + httpinternal "github.com/nais/wonderwall/internal/http" log "github.com/sirupsen/logrus" "github.com/nais/wonderwall/pkg/cookie" @@ -60,27 +61,21 @@ type requestLogger struct { } func (l *requestLogger) NewLogEntry(r *http.Request) *requestLoggerEntry { - referer := r.Referer() - refererUrl, err := url.Parse(referer) - if err == nil { - refererUrl.RawQuery = "" - refererUrl.RawFragment = "" - referer = refererUrl.String() - } - entry := &requestLoggerEntry{} - correlationID := middleware.GetReqID(r.Context()) - fields := log.Fields{ - "correlation_id": correlationID, - "provider": l.Provider, - "request_cookies": nonEmptyRequestCookies(r), - "request_host": r.Host, - "request_method": r.Method, - "request_path": r.URL.Path, - "request_protocol": r.Proto, - "request_referer": referer, - "request_user_agent": r.UserAgent(), + "correlation_id": middleware.GetReqID(r.Context()), + "provider": l.Provider, + "request_cookies": nonEmptyRequestCookies(r), + "request_host": r.Host, + "request_is_navigational": httpinternal.IsNavigationRequest(r), + "request_method": r.Method, + "request_path": r.URL.Path, + "request_protocol": r.Proto, + "request_referer": refererStripped(r), + "request_sec_fetch_dest": r.Header.Get("Sec-Fetch-Dest"), + "request_sec_fetch_mode": r.Header.Get("Sec-Fetch-Mode"), + "request_sec_fetch_site": r.Header.Get("Sec-Fetch-Site"), + "request_user_agent": r.UserAgent(), } entry.Logger = l.Logger.WithFields(fields) @@ -153,3 +148,15 @@ func isRelevantCookie(name string) bool { return false } + +func refererStripped(r *http.Request) string { + referer := r.Referer() + refererUrl, err := url.Parse(referer) + if err == nil { + refererUrl.RawQuery = "" + refererUrl.RawFragment = "" + referer = refererUrl.String() + } + + return referer +}