Preserve continue URL on auth error re-login
Failed OIDC, magic-link, and SAML sign-ins sent users to /auth/error without the post-login destination, so Sign in dropped OAuth flows and deep links. Propagate a validated continue query through auth error redirects, recover it from OIDC state when the IdP denies or cancels login, and forward it from AuthErrorPage to /auth/login. Signed-off-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: Bryan FRIMIN <bryan@frimin.fr>
This commit is contained in:
committed by
Bryan Frimin
parent
428d28fade
commit
9abea50507
@@ -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() {
|
||||
<h1 className="text-2xl font-bold">{content.title}</h1>
|
||||
<p className="text-txt-tertiary">{content.description}</p>
|
||||
</div>
|
||||
<Button className="w-full h-10" to="/auth/login">
|
||||
<Button
|
||||
className="w-full h-10"
|
||||
to={{
|
||||
pathname: "/auth/login",
|
||||
search: loginSearch,
|
||||
}}
|
||||
>
|
||||
{t("auth.actions.signIn")}
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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"}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user