diff --git a/apps/trust/src/pages/auth/ConnectPage.tsx b/apps/trust/src/pages/auth/ConnectPage.tsx index 488023bcc..88ee2b69b 100644 --- a/apps/trust/src/pages/auth/ConnectPage.tsx +++ b/apps/trust/src/pages/auth/ConnectPage.tsx @@ -8,13 +8,14 @@ import { useMutation, usePreloadedQuery, } from "react-relay"; +import { useSearchParams } from "react-router"; import { graphql } from "relay-runtime"; import { z } from "zod"; import { useFormWithSchema } from "#/hooks/useFormWithSchema"; import { getPathPrefix } from "#/utils/pathPrefix"; -import type { ConnectPageMutation } from "./__generated__/ConnectPageMutation.graphql"; +import type { ConnectPageMutation, SendMagicLinkInput } from "./__generated__/ConnectPageMutation.graphql"; import type { ConnectPageQuery } from "./__generated__/ConnectPageQuery.graphql"; export const connectPageQuery = graphql` @@ -53,11 +54,22 @@ export function ConnectPage(props: { const [magicLinkSent, setMagicLinkSent] = useState(false); const interval = useRef(undefined); const [timer, setTimer] = useState(timerDurationSeconds); + const [searchParams] = useSearchParams(); const { currentTrustCenter: { organization }, } = usePreloadedQuery(connectPageQuery, queryRef); + const continueUrlParam = searchParams.get("continue"); + let safeContinueUrl: string; + if (continueUrlParam) { + const continueUrl = new URL(continueUrlParam); + safeContinueUrl = window.location.origin + continueUrl.pathname + continueUrl.search; + } else { + const pathPrefix = getPathPrefix(); + safeContinueUrl = window.location.origin + pathPrefix ? getPathPrefix() : "/"; + } + useEffect(() => { if (!magicLinkSent && interval.current) { clearInterval(interval.current); @@ -92,14 +104,18 @@ export function ConnectPage(props: { ); const handleSubmit = handleSubmitWrapper(({ email }: FormData) => { + const input: SendMagicLinkInput = { email }; + if (safeContinueUrl) { + input.continue = safeContinueUrl; + } sendMagicLink({ variables: { input: { email, + continue: safeContinueUrl, }, }, - onCompleted: (data, errors: GraphQLError[] | null) => { - console.log(data, errors); + onCompleted: (_, errors: GraphQLError[] | null) => { if (errors) { for (const err of errors) { if (err.extensions?.code === "ALREADY_AUTHENTICATED") { diff --git a/apps/trust/src/pages/auth/VerifyMagicLinkPage.tsx b/apps/trust/src/pages/auth/VerifyMagicLinkPage.tsx index acc084823..7de9ab5d2 100644 --- a/apps/trust/src/pages/auth/VerifyMagicLinkPage.tsx +++ b/apps/trust/src/pages/auth/VerifyMagicLinkPage.tsx @@ -14,7 +14,7 @@ import type { VerifyMagicLinkPageMutation } from "./__generated__/VerifyMagicLin const verifyMagicLinkMutation = graphql` mutation VerifyMagicLinkPageMutation($input: VerifyMagicLinkInput!) { verifyMagicLink(input: $input) { - success + continue } } `; @@ -38,7 +38,7 @@ export default function VerifyMagicLinkPagePageMutation() { variables: { input: { token }, }, - onCompleted: (_, errors: GraphQLError[] | null) => { + onCompleted: (response, errors: GraphQLError[] | null) => { if (errors) { for (const err of errors) { if (err.extensions?.code === "ALREADY_AUTHENTICATED") { @@ -55,13 +55,21 @@ export default function VerifyMagicLinkPagePageMutation() { return; } + const { verifyMagicLink } = response; + toast({ title: __("Success"), description: __("Your have successfully signed in"), variant: "success", }); - const pathPrefix = getPathPrefix(); - window.location.href = pathPrefix ? getPathPrefix() : "/"; + + if (verifyMagicLink?.continue) { + const continueUrl = new URL(verifyMagicLink.continue); + window.location.href = window.location.origin + continueUrl.pathname + continueUrl.search; + } else { + const pathPrefix = getPathPrefix(); + window.location.href = pathPrefix ? getPathPrefix() : "/"; + } }, onError: (err) => { toast({ diff --git a/apps/trust/src/pages/auth/__generated__/ConnectPageMutation.graphql.ts b/apps/trust/src/pages/auth/__generated__/ConnectPageMutation.graphql.ts index b420a2bc8..d358edd83 100644 --- a/apps/trust/src/pages/auth/__generated__/ConnectPageMutation.graphql.ts +++ b/apps/trust/src/pages/auth/__generated__/ConnectPageMutation.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<> + * @generated SignedSource<<711ecaa392c23004a3bd1dfb24a5751f>> * @lightSyntaxTransform * @nogrep */ @@ -10,6 +10,7 @@ import { ConcreteRequest } from 'relay-runtime'; export type SendMagicLinkInput = { + continue?: string | null | undefined; email: any; }; export type ConnectPageMutation$variables = { diff --git a/apps/trust/src/pages/auth/__generated__/VerifyMagicLinkPageMutation.graphql.ts b/apps/trust/src/pages/auth/__generated__/VerifyMagicLinkPageMutation.graphql.ts index 47e328c65..83bb04e3d 100644 --- a/apps/trust/src/pages/auth/__generated__/VerifyMagicLinkPageMutation.graphql.ts +++ b/apps/trust/src/pages/auth/__generated__/VerifyMagicLinkPageMutation.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<5037d9722f3c6cc15c2a323ef954e3d6>> + * @generated SignedSource<<5570f368d6cc3c3be80a32390e12c8f9>> * @lightSyntaxTransform * @nogrep */ @@ -17,7 +17,7 @@ export type VerifyMagicLinkPageMutation$variables = { }; export type VerifyMagicLinkPageMutation$data = { readonly verifyMagicLink: { - readonly success: boolean; + readonly continue: string | null | undefined; } | null | undefined; }; export type VerifyMagicLinkPageMutation = { @@ -52,7 +52,7 @@ v1 = [ "alias": null, "args": null, "kind": "ScalarField", - "name": "success", + "name": "continue", "storageKey": null } ], @@ -77,16 +77,16 @@ return { "selections": (v1/*: any*/) }, "params": { - "cacheID": "07cf89de3f37725d847cda46557467f5", + "cacheID": "05d0e504b6f11ad7dd11ed84059a96ac", "id": null, "metadata": {}, "name": "VerifyMagicLinkPageMutation", "operationKind": "mutation", - "text": "mutation VerifyMagicLinkPageMutation(\n $input: VerifyMagicLinkInput!\n) {\n verifyMagicLink(input: $input) {\n success\n }\n}\n" + "text": "mutation VerifyMagicLinkPageMutation(\n $input: VerifyMagicLinkInput!\n) {\n verifyMagicLink(input: $input) {\n continue\n }\n}\n" } }; })(); -(node as any).hash = "074415601c4d50f50d177c06dfab64ef"; +(node as any).hash = "cc9e7f7886d9d61b95f13c5ba66c41ec"; export default node; diff --git a/pkg/iam/auth_service.go b/pkg/iam/auth_service.go index 728ad84aa..6ac8c551d 100644 --- a/pkg/iam/auth_service.go +++ b/pkg/iam/auth_service.go @@ -64,6 +64,7 @@ type ( Email mail.Addr URLPath string OrganizationID gid.GID + Continue *string // If users tries to connect to compliance page, we must brand the emails accordingly CompliancePageID *gid.GID } @@ -73,7 +74,8 @@ type ( } MagicLinkData struct { - Email mail.Addr `json:"email"` + Email mail.Addr `json:"email"` + Continue *string `json:"continue"` } ) @@ -593,7 +595,8 @@ func (s AuthService) SendMagicLink(ctx context.Context, req *SendMagicLinkReques TokenTypeMagicLink, s.magicLinkTokenValidity, MagicLinkData{ - Email: req.Email, + Email: req.Email, + Continue: req.Continue, }, ) if err != nil { @@ -670,16 +673,16 @@ func (s AuthService) SendMagicLink(ctx context.Context, req *SendMagicLinkReques ) } -func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString string) (*coredata.Identity, *coredata.Session, error) { +func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString string) (*coredata.Identity, *coredata.Session, *string, error) { var ( now = time.Now() - identity = &coredata.Identity{} session = &coredata.Session{} + identity = &coredata.Identity{} ) payload, err := statelesstoken.ValidateToken[MagicLinkData](s.tokenSecret, TokenTypeMagicLink, tokenString) if err != nil { - return nil, nil, NewInvalidTokenError() + return nil, nil, nil, NewInvalidTokenError() } if err := s.pg.WithTx( @@ -737,10 +740,10 @@ func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString s return nil }, ); err != nil { - return nil, nil, err + return nil, nil, nil, err } - return identity, session, nil + return identity, session, payload.Data.Continue, nil } func (s *AuthService) UpdateIdentity(ctx context.Context, identityID gid.GID, fullName string) (*coredata.Identity, error) { diff --git a/pkg/server/api/trust/v1/graphql_handler.go b/pkg/server/api/trust/v1/graphql_handler.go index 0ea0ad76f..f07bf01a4 100644 --- a/pkg/server/api/trust/v1/graphql_handler.go +++ b/pkg/server/api/trust/v1/graphql_handler.go @@ -21,6 +21,7 @@ import ( "go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/esign" "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" "go.probo.inc/probo/pkg/server/api/trust/v1/schema" @@ -38,6 +39,7 @@ func NewGraphQLHandler(iamSvc *iam.Service, trustSvc *trust.Service, esignSvc *e logger: logger, baseURL: baseURL, sessionCookie: authn.NewCookie(&cookieConfig), + safeRedirect: &saferedirect.SafeRedirect{AllowedHost: baseURL.Host()}, }, Directives: schema.DirectiveRoot{ Session: session.Directive, diff --git a/pkg/server/api/trust/v1/resolver.go b/pkg/server/api/trust/v1/resolver.go index bb0b384cf..dcdd058b6 100644 --- a/pkg/server/api/trust/v1/resolver.go +++ b/pkg/server/api/trust/v1/resolver.go @@ -26,6 +26,7 @@ import ( "go.probo.inc/probo/pkg/esign" "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" "go.probo.inc/probo/pkg/server/api/compliancepage" @@ -52,6 +53,7 @@ type ( iam *iam.Service sessionCookie *authn.Cookie baseURL *baseurl.BaseURL + safeRedirect *saferedirect.SafeRedirect } ) diff --git a/pkg/server/api/trust/v1/schema.graphql b/pkg/server/api/trust/v1/schema.graphql index 56d996d8e..4e488a370 100644 --- a/pkg/server/api/trust/v1/schema.graphql +++ b/pkg/server/api/trust/v1/schema.graphql @@ -553,6 +553,7 @@ type TrustCenterAccess implements Node { input SendMagicLinkInput { email: EmailAddr! + continue: String } type SendMagicLinkPayload { @@ -564,7 +565,7 @@ input VerifyMagicLinkInput { } type VerifyMagicLinkPayload { - success: Boolean! + continue: String } type RequestAccessesPayload { diff --git a/pkg/server/api/trust/v1/schema/schema.go b/pkg/server/api/trust/v1/schema/schema.go index 1d149dad7..21e20dab5 100644 --- a/pkg/server/api/trust/v1/schema/schema.go +++ b/pkg/server/api/trust/v1/schema/schema.go @@ -277,7 +277,7 @@ type ComplexityRoot struct { } VerifyMagicLinkPayload struct { - Success func(childComplexity int) int + Continue func(childComplexity int) int } } @@ -1205,12 +1205,12 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.ComplexityRoot.VendorEdge.Node(childComplexity), true - case "VerifyMagicLinkPayload.success": - if e.ComplexityRoot.VerifyMagicLinkPayload.Success == nil { + case "VerifyMagicLinkPayload.continue": + if e.ComplexityRoot.VerifyMagicLinkPayload.Continue == nil { break } - return e.ComplexityRoot.VerifyMagicLinkPayload.Success(childComplexity), true + return e.ComplexityRoot.VerifyMagicLinkPayload.Continue(childComplexity), true } return 0, false @@ -1860,6 +1860,7 @@ type TrustCenterAccess implements Node { input SendMagicLinkInput { email: EmailAddr! + continue: String } type SendMagicLinkPayload { @@ -1871,7 +1872,7 @@ input VerifyMagicLinkInput { } type VerifyMagicLinkPayload { - success: Boolean! + continue: String } type RequestAccessesPayload { @@ -3796,8 +3797,8 @@ func (ec *executionContext) fieldContext_Mutation_verifyMagicLink(ctx context.Co IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { - case "success": - return ec.fieldContext_VerifyMagicLinkPayload_success(ctx, field) + case "continue": + return ec.fieldContext_VerifyMagicLinkPayload_continue(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type VerifyMagicLinkPayload", field.Name) }, @@ -6855,30 +6856,30 @@ func (ec *executionContext) fieldContext_VendorEdge_node(_ context.Context, fiel return fc, nil } -func (ec *executionContext) _VerifyMagicLinkPayload_success(ctx context.Context, field graphql.CollectedField, obj *types.VerifyMagicLinkPayload) (ret graphql.Marshaler) { +func (ec *executionContext) _VerifyMagicLinkPayload_continue(ctx context.Context, field graphql.CollectedField, obj *types.VerifyMagicLinkPayload) (ret graphql.Marshaler) { return graphql.ResolveField( ctx, ec.OperationContext, field, - ec.fieldContext_VerifyMagicLinkPayload_success, + ec.fieldContext_VerifyMagicLinkPayload_continue, func(ctx context.Context) (any, error) { - return obj.Success, nil + return obj.Continue, nil }, nil, - ec.marshalNBoolean2bool, - true, + ec.marshalOString2áš–string, true, + false, ) } -func (ec *executionContext) fieldContext_VerifyMagicLinkPayload_success(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { +func (ec *executionContext) fieldContext_VerifyMagicLinkPayload_continue(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { fc = &graphql.FieldContext{ Object: "VerifyMagicLinkPayload", Field: field, IsMethod: false, IsResolver: false, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - return nil, errors.New("field of type Boolean does not have child fields") + return nil, errors.New("field of type String does not have child fields") }, } return fc, nil @@ -8559,7 +8560,7 @@ func (ec *executionContext) unmarshalInputSendMagicLinkInput(ctx context.Context asMap[k] = v } - fieldsInOrder := [...]string{"email"} + fieldsInOrder := [...]string{"email", "continue"} for _, k := range fieldsInOrder { v, ok := asMap[k] if !ok { @@ -8573,6 +8574,13 @@ func (ec *executionContext) unmarshalInputSendMagicLinkInput(ctx context.Context return it, err } it.Email = data + case "continue": + ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("continue")) + data, err := ec.unmarshalOString2áš–string(ctx, v) + if err != nil { + return it, err + } + it.Continue = data } } return it, nil @@ -11243,11 +11251,8 @@ func (ec *executionContext) _VerifyMagicLinkPayload(ctx context.Context, sel ast switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("VerifyMagicLinkPayload") - case "success": - out.Values[i] = ec._VerifyMagicLinkPayload_success(ctx, field, obj) - if out.Values[i] == graphql.Null { - out.Invalids++ - } + case "continue": + out.Values[i] = ec._VerifyMagicLinkPayload_continue(ctx, field, obj) default: panic("unknown field " + strconv.Quote(field.Name)) } diff --git a/pkg/server/api/trust/v1/types/types.go b/pkg/server/api/trust/v1/types/types.go index 688536539..a8685f6f9 100644 --- a/pkg/server/api/trust/v1/types/types.go +++ b/pkg/server/api/trust/v1/types/types.go @@ -194,7 +194,8 @@ type RequestTrustCenterFileAccessInput struct { } type SendMagicLinkInput struct { - Email mail.Addr `json:"email"` + Email mail.Addr `json:"email"` + Continue *string `json:"continue,omitempty"` } type SendMagicLinkPayload struct { @@ -296,5 +297,5 @@ type VerifyMagicLinkInput struct { } type VerifyMagicLinkPayload struct { - Success bool `json:"success"` + Continue *string `json:"continue,omitempty"` } diff --git a/pkg/server/api/trust/v1/v1_resolver.go b/pkg/server/api/trust/v1/v1_resolver.go index 645ca3348..48834649d 100644 --- a/pkg/server/api/trust/v1/v1_resolver.go +++ b/pkg/server/api/trust/v1/v1_resolver.go @@ -157,11 +157,22 @@ func (r *frameworkResolver) DarkLogoURL(ctx context.Context, obj *types.Framewor func (r *mutationResolver) SendMagicLink(ctx context.Context, input types.SendMagicLinkInput) (*types.SendMagicLinkPayload, error) { trustCenter := compliancepage.CompliancePageFromContext(ctx) + var continueURLString *string + if input.Continue != nil { + safeURL, ok := r.safeRedirect.Validate(*input.Continue) + if !ok { + return nil, gqlutils.Invalidf(ctx, "invalid continue URL") + } + + continueURLString = &safeURL + } + req := &iam.SendMagicLinkRequest{ Email: input.Email, CompliancePageID: &trustCenter.ID, OrganizationID: trustCenter.OrganizationID, URLPath: "verify-magic-link", + Continue: continueURLString, } if err := r.iam.AuthService.SendMagicLink(ctx, req); err != nil { @@ -174,7 +185,7 @@ func (r *mutationResolver) SendMagicLink(ctx context.Context, input types.SendMa // VerifyMagicLink is the resolver for the verifyMagicLink field. func (r *mutationResolver) VerifyMagicLink(ctx context.Context, input types.VerifyMagicLinkInput) (*types.VerifyMagicLinkPayload, error) { - identity, session, err := r.iam.AuthService.OpenSessionWithMagicLink(ctx, input.Token) + identity, session, continueURL, err := r.iam.AuthService.OpenSessionWithMagicLink(ctx, input.Token) if err != nil { var errInvalidToken *iam.ErrInvalidToken if errors.As(err, &errInvalidToken) { @@ -196,7 +207,7 @@ func (r *mutationResolver) VerifyMagicLink(ctx context.Context, input types.Veri r.sessionCookie.Set(w, session) return &types.VerifyMagicLinkPayload{ - Success: true, + Continue: continueURL, }, nil }