diff --git a/apps/console/src/pages/iam/auth/AuthErrorPage.tsx b/apps/console/src/pages/iam/auth/AuthErrorPage.tsx index cd4141a77..58838afc1 100644 --- a/apps/console/src/pages/iam/auth/AuthErrorPage.tsx +++ b/apps/console/src/pages/iam/auth/AuthErrorPage.tsx @@ -75,6 +75,11 @@ export default function AuthErrorPage() { const [searchParams] = useSearchParams(); const content = useAuthErrorContent(searchParams.get("error")); + const continueParam = searchParams.get("continue"); + const loginSearch = continueParam + ? `?${new URLSearchParams({ continue: continueParam }).toString()}` + : ""; + usePageTitle(content.title); return ( @@ -83,7 +88,13 @@ export default function AuthErrorPage() {

{content.title}

{content.description}

- diff --git a/pkg/iam/auth_service.go b/pkg/iam/auth_service.go index 15542bab5..2e1bd0717 100644 --- a/pkg/iam/auth_service.go +++ b/pkg/iam/auth_service.go @@ -670,6 +670,15 @@ func (s AuthService) GetMagicLinkEmail(ctx context.Context, tokenString string) return payload.Data.Email, nil } +func (s AuthService) MagicLinkContinueFromToken(tokenString string) (*string, error) { + payload, err := statelesstoken.ValidateTokenAllowExpired[MagicLinkData](s.tokenSecret, TokenTypeMagicLink, tokenString) + if err != nil { + return nil, err + } + + return payload.Data.Continue, nil +} + func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString string) (*coredata.Identity, *coredata.Session, *string, error) { var ( now = time.Now() diff --git a/pkg/iam/oidc/service.go b/pkg/iam/oidc/service.go index 91fa37e71..358701992 100644 --- a/pkg/iam/oidc/service.go +++ b/pkg/iam/oidc/service.go @@ -381,6 +381,53 @@ func (s *Service) InitiateLogin( return authURL, nil } +// ContinueURLFromCallbackState loads and removes a pending login state and returns +// its continue URL. Used when the IdP callback fails before an authorization code +// is exchanged (user denied consent, cancelled login, etc.). +func (s *Service) ContinueURLFromCallbackState( + ctx context.Context, + provider coredata.OIDCProvider, + stateParam string, +) (string, error) { + if stateParam == "" { + return "", NewInvalidStateError() + } + + var oidcState coredata.OIDCState + + err := s.pg.WithTx( + ctx, + func(ctx context.Context, tx pg.Tx) error { + if err := oidcState.LoadByIDForUpdate(ctx, tx, stateParam); err != nil { + if errors.Is(err, coredata.ErrResourceNotFound) { + return NewInvalidStateError() + } + + return fmt.Errorf("cannot load oidc state: %w", err) + } + + if err := oidcState.Delete(ctx, tx); err != nil { + return fmt.Errorf("cannot delete oidc state: %w", err) + } + + return nil + }, + ) + if err != nil { + return "", err + } + + if time.Now().After(oidcState.ExpiresAt) { + return "", NewInvalidStateError() + } + + if oidcState.Provider != provider { + return "", NewInvalidStateError() + } + + return oidcState.ContinueURL, nil +} + func (s *Service) HandleCallback( ctx context.Context, provider coredata.OIDCProvider, @@ -417,11 +464,11 @@ func (s *Service) HandleCallback( } if time.Now().After(oidcState.ExpiresAt) { - return nil, "", nil, NewInvalidStateError() + return nil, oidcState.ContinueURL, nil, NewInvalidStateError() } if oidcState.Provider != provider { - return nil, "", nil, NewInvalidStateError() + return nil, oidcState.ContinueURL, nil, NewInvalidStateError() } token, err := info.oauth2Config.Exchange( @@ -430,26 +477,26 @@ func (s *Service) HandleCallback( oauth2.SetAuthURLParam("code_verifier", oidcState.CodeVerifier), ) if err != nil { - return nil, "", nil, NewCodeExchangeError(err) + return nil, oidcState.ContinueURL, nil, NewCodeExchangeError(err) } rawIDToken, ok := token.Extra("id_token").(string) if !ok { - return nil, "", nil, NewIDTokenMissingError() + return nil, oidcState.ContinueURL, nil, NewIDTokenMissingError() } claims, err := s.verifyAndParseIDToken(ctx, info, rawIDToken, oidcState.Nonce) if err != nil { - return nil, "", nil, fmt.Errorf("cannot verify id token: %w", err) + return nil, oidcState.ContinueURL, nil, fmt.Errorf("cannot verify id token: %w", err) } if err := validateIDTokenClaims(info, claims); err != nil { - return nil, "", nil, err + return nil, oidcState.ContinueURL, nil, err } email, err := mail.ParseAddr(claims.Email) if err != nil { - return nil, "", nil, fmt.Errorf("cannot parse email from id token: %w", err) + return nil, oidcState.ContinueURL, nil, fmt.Errorf("cannot parse email from id token: %w", err) } var identity *coredata.Identity @@ -496,7 +543,7 @@ func (s *Service) HandleCallback( }, ) if err != nil { - return nil, "", nil, err + return nil, oidcState.ContinueURL, nil, err } return identity, oidcState.ContinueURL, oidcState.OrganizationID, nil diff --git a/pkg/server/api/connect/v1/auth_error.go b/pkg/server/api/connect/v1/auth_error.go index 581607a84..f14935573 100644 --- a/pkg/server/api/connect/v1/auth_error.go +++ b/pkg/server/api/connect/v1/auth_error.go @@ -35,10 +35,14 @@ const ( authErrorMagicLinkInvalid = "magic_link_invalid" ) -func redirectAuthError(w http.ResponseWriter, r *http.Request, code string) { +func redirectAuthError(w http.ResponseWriter, r *http.Request, code string, continueURL string) { q := url.Values{} q.Set("error", code) + if continueURL != "" { + q.Set("continue", continueURL) + } + redirectURL := url.URL{ Path: "/auth/error", RawQuery: q.Encode(), diff --git a/pkg/server/api/connect/v1/auth_error_test.go b/pkg/server/api/connect/v1/auth_error_test.go index 3329f286b..21b4d03f6 100644 --- a/pkg/server/api/connect/v1/auth_error_test.go +++ b/pkg/server/api/connect/v1/auth_error_test.go @@ -35,11 +35,29 @@ func TestRedirectAuthError(t *testing.T) { req := httptest.NewRequest(http.MethodGet, "/api/connect/v1/oidc/google/callback", nil) rec := httptest.NewRecorder() - redirectAuthError(rec, req, authErrorPersonalAccountNotAllowed) + redirectAuthError(rec, req, authErrorPersonalAccountNotAllowed, "") assert.Equal(t, http.StatusFound, rec.Code) location, err := rec.Result().Location() require.NoError(t, err) assert.Equal(t, "/auth/error", location.Path) assert.Equal(t, authErrorPersonalAccountNotAllowed, location.Query().Get("error")) + assert.Empty(t, location.Query().Get("continue")) +} + +func TestRedirectAuthErrorWithContinue(t *testing.T) { + t.Parallel() + + req := httptest.NewRequest(http.MethodGet, "/api/connect/v1/oidc/google/callback", nil) + rec := httptest.NewRecorder() + + continueURL := "/overview" + redirectAuthError(rec, req, authErrorAuthenticationFailed, continueURL) + + assert.Equal(t, http.StatusFound, rec.Code) + location, err := rec.Result().Location() + require.NoError(t, err) + assert.Equal(t, "/auth/error", location.Path) + assert.Equal(t, authErrorAuthenticationFailed, location.Query().Get("error")) + assert.Equal(t, continueURL, location.Query().Get("continue")) } diff --git a/pkg/server/api/connect/v1/oidc_handler.go b/pkg/server/api/connect/v1/oidc_handler.go index a5d2857ce..cb02f15b7 100644 --- a/pkg/server/api/connect/v1/oidc_handler.go +++ b/pkg/server/api/connect/v1/oidc_handler.go @@ -21,6 +21,7 @@ package connect_v1 import ( + "context" "errors" "net/http" "strings" @@ -60,6 +61,31 @@ func NewOIDCHandler( } } +func (h *OIDCHandler) redirectAuthError(w http.ResponseWriter, r *http.Request, code string, continueURL string) { + safeContinue := "" + + if continueURL != "" { + if validated, ok := h.safeRedirect.Validate(r.Context(), continueURL); ok { + safeContinue = validated + } + } + + redirectAuthError(w, r, code, safeContinue) +} + +func (h *OIDCHandler) continueURLFromCallbackState( + ctx context.Context, + provider coredata.OIDCProvider, + stateParam string, +) string { + continueURL, err := h.iam.OIDCService.ContinueURLFromCallbackState(ctx, provider, stateParam) + if err != nil { + return "" + } + + return continueURL +} + func (h *OIDCHandler) LoginHandler(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -116,7 +142,8 @@ func (h *OIDCHandler) CallbackHandler(w http.ResponseWriter, r *http.Request) { log.String("error", errParam), log.String("error_description", r.URL.Query().Get("error_description")), ) - redirectAuthError(w, r, authErrorAuthenticationFailed) + continueURL := h.continueURLFromCallbackState(ctx, provider, r.URL.Query().Get("state")) + h.redirectAuthError(w, r, authErrorAuthenticationFailed, continueURL) return } @@ -125,7 +152,12 @@ func (h *OIDCHandler) CallbackHandler(w http.ResponseWriter, r *http.Request) { code := r.URL.Query().Get("code") if stateParam == "" || code == "" { - redirectAuthError(w, r, authErrorInvalidState) + continueURL := "" + if stateParam != "" { + continueURL = h.continueURLFromCallbackState(ctx, provider, stateParam) + } + + h.redirectAuthError(w, r, authErrorInvalidState, continueURL) return } @@ -133,25 +165,25 @@ func (h *OIDCHandler) CallbackHandler(w http.ResponseWriter, r *http.Request) { identity, continueURL, organizationID, err := h.iam.OIDCService.HandleCallback(ctx, provider, stateParam, code) if err != nil { if _, ok := errors.AsType[*oidc.ErrPersonalAccountNotAllowed](err); ok { - redirectAuthError(w, r, authErrorPersonalAccountNotAllowed) + h.redirectAuthError(w, r, authErrorPersonalAccountNotAllowed, continueURL) return } if _, ok := errors.AsType[*oidc.ErrEmailNotVerified](err); ok { - redirectAuthError(w, r, authErrorEmailNotVerified) + h.redirectAuthError(w, r, authErrorEmailNotVerified, continueURL) return } if _, ok := errors.AsType[*oidc.ErrInvalidState](err); ok { - redirectAuthError(w, r, authErrorInvalidState) + h.redirectAuthError(w, r, authErrorInvalidState, continueURL) return } h.logger.ErrorCtx(ctx, "cannot handle OIDC callback", log.Error(err)) - redirectAuthError(w, r, authErrorAuthenticationFailed) + h.redirectAuthError(w, r, authErrorAuthenticationFailed, continueURL) return } @@ -259,6 +291,21 @@ func NewMagicLinkHandler( } } +func (h *MagicLinkHandler) redirectAuthError(w http.ResponseWriter, r *http.Request, code string, token string) { + safeContinue := "" + + if token != "" { + continueURL, err := h.iam.AuthService.MagicLinkContinueFromToken(token) + if err == nil && continueURL != nil && *continueURL != "" { + if validated, ok := h.safeRedirect.Validate(r.Context(), *continueURL); ok { + safeContinue = validated + } + } + } + + redirectAuthError(w, r, code, safeContinue) +} + func (h *MagicLinkHandler) SendHandler(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -313,29 +360,33 @@ func (h *MagicLinkHandler) VerifyHandler(w http.ResponseWriter, r *http.Request) token := r.URL.Query().Get("token") if token == "" { - redirectAuthError(w, r, authErrorMagicLinkInvalid) + h.redirectAuthError(w, r, authErrorMagicLinkInvalid, "") + return } identity, session, continueURL, err := h.iam.AuthService.OpenSessionWithMagicLink(ctx, token) if err != nil { if _, ok := errors.AsType[*iam.ErrExpiredToken](err); ok { - redirectAuthError(w, r, authErrorMagicLinkExpired) + h.redirectAuthError(w, r, authErrorMagicLinkExpired, token) + return } if _, ok := errors.AsType[*iam.ErrTokenAlreadyUsed](err); ok { - redirectAuthError(w, r, authErrorMagicLinkAlreadyUsed) + h.redirectAuthError(w, r, authErrorMagicLinkAlreadyUsed, token) + return } if _, ok := errors.AsType[*iam.ErrInvalidToken](err); ok { - redirectAuthError(w, r, authErrorMagicLinkInvalid) + h.redirectAuthError(w, r, authErrorMagicLinkInvalid, token) + return } h.logger.ErrorCtx(ctx, "cannot open session with magic link", log.Error(err)) - redirectAuthError(w, r, authErrorAuthenticationFailed) + h.redirectAuthError(w, r, authErrorAuthenticationFailed, token) return } diff --git a/pkg/server/api/connect/v1/saml_handler.go b/pkg/server/api/connect/v1/saml_handler.go index 5fef228ff..a6353991b 100644 --- a/pkg/server/api/connect/v1/saml_handler.go +++ b/pkg/server/api/connect/v1/saml_handler.go @@ -21,6 +21,7 @@ package connect_v1 import ( + "context" "errors" "fmt" "net/http" @@ -59,9 +60,28 @@ func (h *SAMLHandler) renderInternalServerError(w http.ResponseWriter) { httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) } -func (h *SAMLHandler) renderAssertionError(w http.ResponseWriter, r *http.Request, err error) { +func (h *SAMLHandler) continueURLFromRelayState(ctx context.Context, relayState string) string { + if len(relayState) <= gid.EncodedGIDSize { + return "" + } + + unescapedContinueURL, err := url.QueryUnescape(relayState[gid.EncodedGIDSize:]) + if err != nil { + return "" + } + + safeContinue, ok := h.safeRedirect.Validate(ctx, unescapedContinueURL) + if !ok { + return "" + } + + return safeContinue +} + +func (h *SAMLHandler) renderAssertionError(w http.ResponseWriter, r *http.Request, relayState string, err error) { h.logger.ErrorCtx(r.Context(), "cannot handle SAML assertion", log.Error(err)) - redirectAuthError(w, r, authErrorAuthenticationFailed) + continueURL := h.continueURLFromRelayState(r.Context(), relayState) + redirectAuthError(w, r, authErrorAuthenticationFailed, continueURL) } func (h *SAMLHandler) MetadataHandler(w http.ResponseWriter, r *http.Request) { @@ -106,7 +126,8 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { user, membership, err := h.iam.SAMLService.HandleAssertion(ctx, samlResponse, configID) if err != nil { - h.renderAssertionError(w, r, err) + h.renderAssertionError(w, r, relayState, err) + return } diff --git a/pkg/statelesstoken/statelesstoken.go b/pkg/statelesstoken/statelesstoken.go index ce0f9f43e..3d0dbb88c 100644 --- a/pkg/statelesstoken/statelesstoken.go +++ b/pkg/statelesstoken/statelesstoken.go @@ -128,6 +128,25 @@ func DecodePayload[T any](tokenString string) (*Payload[T], error) { // ValidateToken validates a token and unmarshals the payload // It returns an error if the token is invalid or expired func ValidateToken[T any](secret string, tokenType string, tokenString string) (*Payload[T], error) { + payload, err := parseSignedToken[T](secret, tokenType, tokenString) + if err != nil { + return nil, err + } + + if time.Now().After(payload.ExpiresAt) { + return nil, &ErrExpiredToken{message: "token has expired"} + } + + return payload, nil +} + +// ValidateTokenAllowExpired validates signature and type but ignores expiration. +// Use only when recovering non-sensitive metadata (e.g. post-auth redirect URLs). +func ValidateTokenAllowExpired[T any](secret string, tokenType string, tokenString string) (*Payload[T], error) { + return parseSignedToken[T](secret, tokenType, tokenString) +} + +func parseSignedToken[T any](secret string, tokenType string, tokenString string) (*Payload[T], error) { parts := strings.Split(tokenString, ".") if len(parts) != 2 { return nil, &ErrInvalidToken{message: "invalid token format"} @@ -154,10 +173,6 @@ func ValidateToken[T any](secret string, tokenType string, tokenString string) ( return nil, fmt.Errorf("cannot unmarshal token payload: %w", err) } - if time.Now().After(payload.ExpiresAt) { - return nil, &ErrExpiredToken{message: "token has expired"} - } - if payload.Type != tokenType { return nil, &ErrInvalidToken{message: "invalid token type"} }