From e46c1d48946314c8717aa798554befab264c3a72 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=89mile=20R=C3=A9?= Date: Wed, 11 Feb 2026 21:19:15 +0400 Subject: [PATCH] Rename redirect-path to continue + handle redirection on forbidden from wrong org assume MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Émile Ré --- .../iam/AssumePageMutation.graphql.ts | 4 +-- .../components/OrganizationErrorBoundary.tsx | 5 ++- .../src/components/RootErrorBoundary.tsx | 15 +++++--- .../iam/auth/sign-in/PasswordSignInPage.tsx | 7 ++-- .../pages/iam/auth/sign-in/SSOSignInPage.tsx | 2 +- .../pages/iam/organizations/AssumePage.tsx | 12 +++---- e2e/internal/testutil/client.go | 2 +- pkg/gid/gid.go | 2 +- pkg/iam/saml/service.go | 4 +-- pkg/iam/session_service.go | 4 +-- pkg/server/api/connect/v1/saml_handler.go | 36 +++++++++++-------- pkg/server/api/connect/v1/schema.graphql | 2 +- pkg/server/api/connect/v1/schema/schema.go | 14 ++++---- pkg/server/api/connect/v1/types/types.go | 10 +++--- pkg/server/api/connect/v1/v1_resolver.go | 10 +++--- 15 files changed, 74 insertions(+), 55 deletions(-) diff --git a/apps/console/src/__generated__/iam/AssumePageMutation.graphql.ts b/apps/console/src/__generated__/iam/AssumePageMutation.graphql.ts index 7ca78aa77..3a947be8d 100644 --- a/apps/console/src/__generated__/iam/AssumePageMutation.graphql.ts +++ b/apps/console/src/__generated__/iam/AssumePageMutation.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<8717e91a7e1602c06f6c67fa76580ecf>> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -11,8 +11,8 @@ import { ConcreteRequest } from 'relay-runtime'; export type ReauthenticationReason = "POLICY_REQUIREMENT" | "SENSITIVE_ACTION" | "SESSION_EXPIRED"; export type AssumeOrganizationSessionInput = { + continue: string; organizationId: string; - redirectPath: string; }; export type AssumePageMutation$variables = { input: AssumeOrganizationSessionInput; diff --git a/apps/console/src/components/OrganizationErrorBoundary.tsx b/apps/console/src/components/OrganizationErrorBoundary.tsx index ad79cc3b5..9d6e0b6a9 100644 --- a/apps/console/src/components/OrganizationErrorBoundary.tsx +++ b/apps/console/src/components/OrganizationErrorBoundary.tsx @@ -11,9 +11,12 @@ export function OrganizationErrorBoundary() { const search = new URLSearchParams([ ["organization-id", organizationId], - ["redirect-path", window.location.pathname + window.location.search], ]); + if (window.location.href !== window.location.origin) { + search.set("continue", window.location.href); + } + if (error instanceof UnAuthenticatedError) { return ; } diff --git a/apps/console/src/components/RootErrorBoundary.tsx b/apps/console/src/components/RootErrorBoundary.tsx index 77fcbcbde..1f2354559 100644 --- a/apps/console/src/components/RootErrorBoundary.tsx +++ b/apps/console/src/components/RootErrorBoundary.tsx @@ -6,12 +6,19 @@ import { PageError } from "./PageError"; export function RootErrorBoundary() { const error = useRouteError(); - const search = new URLSearchParams([ - ["redirect-path", window.location.pathname + window.location.search], - ]); + const search = new URLSearchParams(); + if (window.location.href !== window.location.origin) { + search.set("continue", window.location.href); + } if (error instanceof UnAuthenticatedError) { - return ; + return ( + + ); } return ; diff --git a/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx index f55d4d83d..0736f845c 100644 --- a/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx +++ b/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx @@ -3,7 +3,7 @@ import { useTranslate } from "@probo/i18n"; import { Button, Field, IconChevronLeft, useToast } from "@probo/ui"; import type { FormEventHandler } from "react"; import { useMutation } from "react-relay"; -import { Link, useLocation, useNavigate, useSearchParams } from "react-router"; +import { Link, useLocation, useSearchParams } from "react-router"; import { graphql } from "relay-runtime"; import type { PasswordSignInPageMutation } from "#/__generated__/iam/PasswordSignInPageMutation.graphql"; @@ -21,7 +21,6 @@ const signInMutation = graphql` export default function PasswordSignInPage() { const location = useLocation(); const [searchParams] = useSearchParams(); - const navigate = useNavigate(); const { __ } = useTranslate(); const { toast } = useToast(); @@ -56,10 +55,11 @@ export default function PasswordSignInPage() { ), variant: "error", }); + window.location.href = "/"; return; } - void navigate(searchParams.get("redirect-path") ?? "/"); + window.location.href = searchParams.get("continue") ?? window.location.origin; }, onError: (e) => { toast({ @@ -67,6 +67,7 @@ export default function PasswordSignInPage() { description: e.message, variant: "error", }); + window.location.href = "/"; }, }); }; diff --git a/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx index 9066ae615..0b868506e 100644 --- a/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx +++ b/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx @@ -130,7 +130,7 @@ function NavigateToSSOLoginURL(props: { } const url = new URL(ssoLoginURL.value); - url.search = "?" + searchParams.toString(); + url.search = searchParams.toString(); window.location.href = url.toString(); }, [__, navigate, ssoLoginURL, toast, searchParams]); diff --git a/apps/console/src/pages/iam/organizations/AssumePage.tsx b/apps/console/src/pages/iam/organizations/AssumePage.tsx index 46a286561..35e772116 100644 --- a/apps/console/src/pages/iam/organizations/AssumePage.tsx +++ b/apps/console/src/pages/iam/organizations/AssumePage.tsx @@ -38,18 +38,18 @@ function AssumePageInner() { const [assumeOrganizationSession] = useMutation(assumeMutation); - const redirectPath = searchParams.get("redirect-path") ?? `/organizations/${organizationId}`; + const continueUrl = searchParams.get("continue") ?? new URL(`/organizations/${organizationId}`, window.location.origin).toString(); useEffect(() => { assumeOrganizationSession({ variables: { - input: { organizationId, redirectPath }, + input: { organizationId, continue: continueUrl }, }, onError: (error) => { if (error instanceof UnAuthenticatedError) { const search = new URLSearchParams([ ["organization-id", organizationId], - ["redirect-path", redirectPath], + ["continue", continueUrl], ]); void navigate({ pathname: "/auth/login", search: "?" + search.toString() }); @@ -68,7 +68,7 @@ function AssumePageInner() { switch (result.__typename) { case "PasswordRequired": search.set("organization-id", organizationId); - search.set("redirect-path", redirectPath); + search.set("continue", continueUrl); void navigate({ pathname: "/auth/password-login", search: "?" + search.toString() }); break; @@ -79,11 +79,11 @@ function AssumePageInner() { window.location.href = samlSSOLoginURL.toString(); break; default: - void navigate(redirectPath); + window.location.href = continueUrl; } }, }); - }, [organizationId, navigate, assumeOrganizationSession, redirectPath, searchParams]); + }, [organizationId, navigate, assumeOrganizationSession, continueUrl, searchParams]); return ( diff --git a/e2e/internal/testutil/client.go b/e2e/internal/testutil/client.go index 324eadcec..d8c5d9f07 100644 --- a/e2e/internal/testutil/client.go +++ b/e2e/internal/testutil/client.go @@ -339,7 +339,7 @@ func (c *Client) assumeOrganizationSession() { err := c.ExecuteConnect(query, map[string]any{ "input": map[string]any{ "organizationId": c.organizationID.String(), - "redirectPath": "/", + "continue": c.baseURL, }, }, nil) require.NoError(c.T, err, "assumeOrganizationSession mutation failed") diff --git a/pkg/gid/gid.go b/pkg/gid/gid.go index 0fd3ecdf6..a9b31a4f9 100644 --- a/pkg/gid/gid.go +++ b/pkg/gid/gid.go @@ -20,7 +20,7 @@ type ( var ( Nil = GID{} - EncodedGIDSize = base64.RawURLEncoding.EncodedLen(GIDSize) + EncodedGIDSize = base64.RawURLEncoding.EncodedLen(len(Nil)) ) // ParseGID parses a string representation of a GID diff --git a/pkg/iam/saml/service.go b/pkg/iam/saml/service.go index f76765d20..500f1b31d 100644 --- a/pkg/iam/saml/service.go +++ b/pkg/iam/saml/service.go @@ -107,7 +107,7 @@ func (s *Service) GenerateSpMetadata() ([]byte, error) { func (s *Service) InitiateLogin( ctx context.Context, configID gid.GID, - redirectPath string, + continuePath string, ) (*url.URL, error) { var ( now = time.Now() @@ -153,7 +153,7 @@ func (s *Service) InitiateLogin( return fmt.Errorf("cannot insert SAML request: %w", err) } - relayState := config.ID.String() + redirectPath + relayState := config.ID.String() + url.QueryEscape(continuePath) redirect, err = req.Redirect(relayState, sp) if err != nil { diff --git a/pkg/iam/session_service.go b/pkg/iam/session_service.go index a556b0bc0..7889a5c6e 100644 --- a/pkg/iam/session_service.go +++ b/pkg/iam/session_service.go @@ -491,7 +491,7 @@ func (s SessionService) AssumeOrganizationSession( ctx context.Context, sessionID gid.GID, organizationID gid.GID, - redirectPath string, + continueURL string, ) (*coredata.Session, *coredata.Membership, error) { var ( now = time.Now() @@ -549,7 +549,7 @@ func (s SessionService) AssumeOrganizationSession( if err == nil && samlConfig.EnforcementPolicy == coredata.SAMLEnforcementPolicyRequired { if rootSession.AuthMethod != coredata.AuthMethodSAML { - redirectURL, err := s.SAMLService.InitiateLogin(ctx, samlConfig.ID, redirectPath) + redirectURL, err := s.SAMLService.InitiateLogin(ctx, samlConfig.ID, continueURL) if err != nil { return fmt.Errorf("cannot initiate SAML login: %w", err) } diff --git a/pkg/server/api/connect/v1/saml_handler.go b/pkg/server/api/connect/v1/saml_handler.go index 32ac3cf27..0b43c2a5d 100644 --- a/pkg/server/api/connect/v1/saml_handler.go +++ b/pkg/server/api/connect/v1/saml_handler.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" "net/http" + "net/url" "github.com/go-chi/chi/v5" "go.gearno.de/kit/httpserver" @@ -11,6 +12,7 @@ import ( "go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/iam" + "go.probo.inc/probo/pkg/saferedirect" "go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/server/api/authn" ) @@ -20,6 +22,7 @@ type SAMLHandler struct { sessionCookie *authn.Cookie baseURL *baseurl.BaseURL logger *log.Logger + safeRedirect *saferedirect.SafeRedirect } func NewSAMLHandler(iam *iam.Service, cookieConfig securecookie.Config, baseURL *baseurl.BaseURL, logger *log.Logger) *SAMLHandler { @@ -28,10 +31,11 @@ func NewSAMLHandler(iam *iam.Service, cookieConfig securecookie.Config, baseURL sessionCookie: authn.NewCookie(&cookieConfig), baseURL: baseURL, logger: logger, + safeRedirect: &saferedirect.SafeRedirect{AllowedHost: baseURL.Host()}, } } -func (h *SAMLHandler) renderInternalServerError(w http.ResponseWriter, r *http.Request) { +func (h *SAMLHandler) renderInternalServerError(w http.ResponseWriter) { httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) } @@ -63,7 +67,7 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { return } - configIDStr := relayState[:gid.EncodedGIDSize+1] + configIDStr := relayState[:gid.EncodedGIDSize] if configIDStr == "" { httpserver.RenderError(w, http.StatusBadRequest, errors.New("missing config ID")) return @@ -75,16 +79,21 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { return } - redirectPath := relayState[gid.EncodedGIDSize+1:] - user, membership, err := h.iam.SAMLService.HandleAssertion(ctx, samlResponse, configID) if err != nil { httpserver.RenderError(w, http.StatusUnauthorized, err) return } - if redirectPath == "" { - redirectPath = "/organizations/" + membership.OrganizationID.String() + continueURL := "/organizations/" + membership.OrganizationID.String() + if len(relayState) > gid.EncodedGIDSize { + unescapedContinueURL, err := url.QueryUnescape(relayState[gid.EncodedGIDSize:]) + + if err != nil { + h.logger.WarnCtx(ctx, "cannot unescape continue URL from RelayState", log.Error(err)) + } else { + continueURL = unescapedContinueURL + } } rootSession := authn.SessionFromContext(ctx) @@ -94,21 +103,21 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { rootSession, err = h.iam.AuthService.OpenSessionWithSAML(ctx, user.ID) if err != nil { h.logger.ErrorCtx(ctx, "cannot open root session", log.Error(err)) - h.renderInternalServerError(w, r) + h.renderInternalServerError(w) return } case rootSession.IdentityID != user.ID: err = h.iam.SessionService.CloseSession(ctx, rootSession.ID) if err != nil { h.logger.ErrorCtx(ctx, "cannot close session", log.Error(err)) - h.renderInternalServerError(w, r) + h.renderInternalServerError(w) return } rootSession, err = h.iam.AuthService.OpenSessionWithSAML(ctx, user.ID) if err != nil { h.logger.ErrorCtx(ctx, "cannot open root session", log.Error(err)) - h.renderInternalServerError(w, r) + h.renderInternalServerError(w) return } } @@ -116,14 +125,13 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { _, _, err = h.iam.SessionService.OpenSAMLChildSessionForOrganization(ctx, rootSession.ID, membership.OrganizationID) if err != nil { h.logger.ErrorCtx(ctx, "cannot open SAML child session", log.Error(err)) - h.renderInternalServerError(w, r) + h.renderInternalServerError(w) return } h.sessionCookie.Set(w, rootSession) - redirectURL := h.baseURL.WithPath(redirectPath).MustString() - http.Redirect(w, r, redirectURL, http.StatusFound) + h.safeRedirect.Redirect(w, r, continueURL, "/organizations/"+membership.OrganizationID.String(), http.StatusFound) } func (h *SAMLHandler) LoginHandler(w http.ResponseWriter, r *http.Request) { @@ -135,7 +143,7 @@ func (h *SAMLHandler) LoginHandler(w http.ResponseWriter, r *http.Request) { return } - redirectPathQueryParam := r.URL.Query().Get("redirect-path") + continueURLQueryParam := r.URL.Query().Get("continue") samlConfigID, err := gid.ParseGID(samlConfigIDParam) if err != nil { @@ -143,7 +151,7 @@ func (h *SAMLHandler) LoginHandler(w http.ResponseWriter, r *http.Request) { return } - url, err := h.iam.SAMLService.InitiateLogin(ctx, samlConfigID, redirectPathQueryParam) + url, err := h.iam.SAMLService.InitiateLogin(ctx, samlConfigID, continueURLQueryParam) if err != nil { panic(fmt.Errorf("cannot initiate SAML login: %w", err)) } diff --git a/pkg/server/api/connect/v1/schema.graphql b/pkg/server/api/connect/v1/schema.graphql index 63a0efd1c..c6bbe80ee 100644 --- a/pkg/server/api/connect/v1/schema.graphql +++ b/pkg/server/api/connect/v1/schema.graphql @@ -669,7 +669,7 @@ input ChangeEmailInput { input AssumeOrganizationSessionInput { organizationId: ID! - redirectPath: String! + continue: String! } input RevokeSessionInput { diff --git a/pkg/server/api/connect/v1/schema/schema.go b/pkg/server/api/connect/v1/schema/schema.go index e4d22d475..03b34798c 100644 --- a/pkg/server/api/connect/v1/schema/schema.go +++ b/pkg/server/api/connect/v1/schema/schema.go @@ -246,7 +246,7 @@ type ComplexityRoot struct { UpdateMembership func(childComplexity int, input types.UpdateMembershipInput) int UpdateOrganization func(childComplexity int, input types.UpdateOrganizationInput) int UpdateSAMLConfiguration func(childComplexity int, input types.UpdateSAMLConfigurationInput) int - UpdateSCIMBridge func(childComplexity int, input types.UpdateSCIMBridgeInput) int + UpdateSCIMBridge func(childComplexity int, input types.UpdateSCIMBridgeInput) int VerifyEmail func(childComplexity int, input types.VerifyEmailInput) int } @@ -380,7 +380,7 @@ type ComplexityRoot struct { SCIMBridge struct { Connector func(childComplexity int) int CreatedAt func(childComplexity int) int - ExcludedUserNames func(childComplexity int) int + ExcludedUserNames func(childComplexity int) int ID func(childComplexity int) int Permission func(childComplexity int, action string) int ScimConfiguration func(childComplexity int) int @@ -3111,7 +3111,7 @@ input ChangeEmailInput { input AssumeOrganizationSessionInput { organizationId: ID! - redirectPath: String! + continue: String! } input RevokeSessionInput { @@ -14931,7 +14931,7 @@ func (ec *executionContext) unmarshalInputAssumeOrganizationSessionInput(ctx con asMap[k] = v } - fieldsInOrder := [...]string{"organizationId", "redirectPath"} + fieldsInOrder := [...]string{"organizationId", "continue"} for _, k := range fieldsInOrder { v, ok := asMap[k] if !ok { @@ -14945,13 +14945,13 @@ func (ec *executionContext) unmarshalInputAssumeOrganizationSessionInput(ctx con return it, err } it.OrganizationID = data - case "redirectPath": - ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("redirectPath")) + case "continue": + ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("continue")) data, err := ec.unmarshalNString2string(ctx, v) if err != nil { return it, err } - it.RedirectPath = data + it.Continue = data } } diff --git a/pkg/server/api/connect/v1/types/types.go b/pkg/server/api/connect/v1/types/types.go index 43ec753d6..a32c3806e 100644 --- a/pkg/server/api/connect/v1/types/types.go +++ b/pkg/server/api/connect/v1/types/types.go @@ -36,7 +36,7 @@ type AcceptInvitationPayload struct { type AssumeOrganizationSessionInput struct { OrganizationID gid.GID `json:"organizationId"` - RedirectPath string `json:"redirectPath"` + Continue string `json:"continue"` } type AssumeOrganizationSessionPayload struct { @@ -416,7 +416,7 @@ type SCIMBridge struct { ScimConfiguration *SCIMConfiguration `json:"scimConfiguration,omitempty"` Connector *Connector `json:"connector,omitempty"` Type coredata.SCIMBridgeType `json:"type"` - ExcludedUserNames []string `json:"excludedUserNames"` + ExcludedUserNames []string `json:"excludedUserNames"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt"` Permission bool `json:"permission"` @@ -561,9 +561,9 @@ type UpdateSAMLConfigurationPayload struct { } type UpdateSCIMBridgeInput struct { - OrganizationID gid.GID `json:"organizationId"` - ScimBridgeID gid.GID `json:"scimBridgeId"` - ExcludedUserNames []string `json:"excludedUserNames"` + OrganizationID gid.GID `json:"organizationId"` + ScimBridgeID gid.GID `json:"scimBridgeId"` + ExcludedUserNames []string `json:"excludedUserNames"` } type UpdateSCIMBridgePayload struct { diff --git a/pkg/server/api/connect/v1/v1_resolver.go b/pkg/server/api/connect/v1/v1_resolver.go index e56e985f7..96a9fa97c 100644 --- a/pkg/server/api/connect/v1/v1_resolver.go +++ b/pkg/server/api/connect/v1/v1_resolver.go @@ -406,6 +406,9 @@ func (r *mutationResolver) SignIn(ctx context.Context, input types.SignInInput) } } + w := gqlutils.HTTPResponseWriterFromContext(ctx) + r.sessionCookie.Set(w, session) + if input.OrganizationID != nil { var err error _, _, err = r.iam.SessionService.OpenPasswordChildSessionForOrganization(ctx, session.ID, *input.OrganizationID) @@ -415,7 +418,7 @@ func (r *mutationResolver) SignIn(ctx context.Context, input types.SignInInput) var errMembershipInactive *iam.ErrMembershipInactive if errors.As(err, &errMembershipNotFound) || errors.As(err, &errMembershipInactive) { - return nil, gqlutils.Forbidden(ctx, err) + return nil, gqlutils.Forbiddenf(ctx, "forbidden") } r.logger.ErrorCtx(ctx, "cannot assume organization", log.Error(err)) @@ -423,9 +426,6 @@ func (r *mutationResolver) SignIn(ctx context.Context, input types.SignInInput) } } - w := gqlutils.HTTPResponseWriterFromContext(ctx) - r.sessionCookie.Set(w, session) - return &types.SignInPayload{ Identity: types.NewIdentity(identity), Session: types.NewSession(session), @@ -680,7 +680,7 @@ func (r *mutationResolver) ChangeEmail(ctx context.Context, input types.ChangeEm func (r *mutationResolver) AssumeOrganizationSession(ctx context.Context, input types.AssumeOrganizationSessionInput) (*types.AssumeOrganizationSessionPayload, error) { rootSession := authn.SessionFromContext(ctx) - childSession, membership, err := r.iam.SessionService.AssumeOrganizationSession(ctx, rootSession.ID, input.OrganizationID, input.RedirectPath) + childSession, membership, err := r.iam.SessionService.AssumeOrganizationSession(ctx, rootSession.ID, input.OrganizationID, input.Continue) if err != nil { var ( errMembershipNotFound *iam.ErrMembershipNotFound