package site import ( "context" "errors" "fmt" "net/http" "time" "ruben/inventory2/internal/consts" "ruben/inventory2/internal/domains/authentication" ) // just keep this around long enough for testing auth middleware.. func (s *Server) testAuthEndpoint(w http.ResponseWriter, r *http.Request) { fmt.Println("SUCCESS:", getAccessTokenClaims(r.Context())) http.Redirect(w, r, "/", http.StatusTemporaryRedirect) } type customClaimsKey struct{} // auth middleware to verify access_token cookie and set custom claims in the request context func (s *Server) authenticate(h http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ck, err := r.Cookie("access_token") if err != nil { http.Redirect(w, r, "/", http.StatusTemporaryRedirect) return } ctx := r.Context() claims, expiration, err := s.auth.GetAccessTokenClaimsAndExpiration(ctx, ck.Value) if err != nil { if errors.Is(err, consts.ErrNotFound) { http.Redirect(w, r, "/", http.StatusTemporaryRedirect) return } http.Error(w, fmt.Sprintf("failed to authenticate: %v", err), http.StatusInternalServerError) return } if expiration.Before(time.Now()) { deleteCookieInResponse(w, "access_token") http.Redirect(w, r, "/", http.StatusTemporaryRedirect) return } h.ServeHTTP(w, r.WithContext(setAccessTokenClaims(ctx, claims))) }) } func (s *Server) getAccessTokenClaims(r *http.Request) (authentication.AccessTokenClaims, bool) { ck, err := r.Cookie("access_token") if err != nil { return authentication.AccessTokenClaims{}, false } ctx := r.Context() claims, expiration, err := s.auth.GetAccessTokenClaimsAndExpiration(ctx, ck.Value) if err != nil { return authentication.AccessTokenClaims{}, false } if expiration.Before(time.Now()) { return authentication.AccessTokenClaims{}, false } return claims, true } // stores custom claims in request context func setAccessTokenClaims(ctx context.Context, claims authentication.AccessTokenClaims) context.Context { return context.WithValue(ctx, customClaimsKey{}, claims) } // get custom claims from request context func getAccessTokenClaims(ctx context.Context) authentication.AccessTokenClaims { c, _ := ctx.Value(customClaimsKey{}).(authentication.AccessTokenClaims) return c }