diff --git a/apps/console/src/__generated__/iam/SignInPageMutation.graphql.ts b/apps/console/src/__generated__/iam/PasswordSignInPageMutation.graphql.ts similarity index 71% rename from apps/console/src/__generated__/iam/SignInPageMutation.graphql.ts rename to apps/console/src/__generated__/iam/PasswordSignInPageMutation.graphql.ts index 7d2edfe6b..ae514bff7 100644 --- a/apps/console/src/__generated__/iam/SignInPageMutation.graphql.ts +++ b/apps/console/src/__generated__/iam/PasswordSignInPageMutation.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<02b9c661a564d47f4821ff67a4dae9b0>> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -13,19 +13,19 @@ export type SignInInput = { email: string; password: string; }; -export type SignInPageMutation$variables = { +export type PasswordSignInPageMutation$variables = { input: SignInInput; }; -export type SignInPageMutation$data = { +export type PasswordSignInPageMutation$data = { readonly signIn: { readonly session: { readonly id: string; } | null | undefined; } | null | undefined; }; -export type SignInPageMutation = { - response: SignInPageMutation$data; - variables: SignInPageMutation$variables; +export type PasswordSignInPageMutation = { + response: PasswordSignInPageMutation$data; + variables: PasswordSignInPageMutation$variables; }; const node: ConcreteRequest = (function(){ @@ -78,7 +78,7 @@ return { "argumentDefinitions": (v0/*: any*/), "kind": "Fragment", "metadata": null, - "name": "SignInPageMutation", + "name": "PasswordSignInPageMutation", "selections": (v1/*: any*/), "type": "Mutation", "abstractKey": null @@ -87,20 +87,20 @@ return { "operation": { "argumentDefinitions": (v0/*: any*/), "kind": "Operation", - "name": "SignInPageMutation", + "name": "PasswordSignInPageMutation", "selections": (v1/*: any*/) }, "params": { - "cacheID": "8e13b6ebdfd4ca6bde158dff53f596db", + "cacheID": "f60c1dcde17bef06ff8170d14c077dae", "id": null, "metadata": {}, - "name": "SignInPageMutation", + "name": "PasswordSignInPageMutation", "operationKind": "mutation", - "text": "mutation SignInPageMutation(\n $input: SignInInput!\n) {\n signIn(input: $input) {\n session {\n id\n }\n }\n}\n" + "text": "mutation PasswordSignInPageMutation(\n $input: SignInInput!\n) {\n signIn(input: $input) {\n session {\n id\n }\n }\n}\n" } }; })(); -(node as any).hash = "6424104df106b64a6237d92299239a63"; +(node as any).hash = "9b9b3de59d38ae7c288348c6bfc01afa"; export default node; diff --git a/apps/console/src/__generated__/iam/SSOSignInPageQuery.graphql.ts b/apps/console/src/__generated__/iam/SSOSignInPageQuery.graphql.ts new file mode 100644 index 000000000..f001d4cda --- /dev/null +++ b/apps/console/src/__generated__/iam/SSOSignInPageQuery.graphql.ts @@ -0,0 +1,83 @@ +/** + * @generated SignedSource<<44e9b1516c9cd0ae8419f326791e09e1>> + * @lightSyntaxTransform + * @nogrep + */ + +/* tslint:disable */ +/* eslint-disable */ +// @ts-nocheck + +import { ConcreteRequest } from 'relay-runtime'; +import { Result } from "relay-runtime"; +export type SSOSignInPageQuery$variables = { + email: string; +}; +export type SSOSignInPageQuery$data = { + readonly ssoLoginURL: Result; +}; +export type SSOSignInPageQuery = { + response: SSOSignInPageQuery$data; + variables: SSOSignInPageQuery$variables; +}; + +const node: ConcreteRequest = (function(){ +var v0 = [ + { + "defaultValue": null, + "kind": "LocalArgument", + "name": "email" + } +], +v1 = { + "alias": null, + "args": [ + { + "kind": "Variable", + "name": "email", + "variableName": "email" + } + ], + "kind": "ScalarField", + "name": "ssoLoginURL", + "storageKey": null +}; +return { + "fragment": { + "argumentDefinitions": (v0/*: any*/), + "kind": "Fragment", + "metadata": null, + "name": "SSOSignInPageQuery", + "selections": [ + { + "kind": "CatchField", + "field": (v1/*: any*/), + "to": "RESULT" + } + ], + "type": "Query", + "abstractKey": null + }, + "kind": "Request", + "operation": { + "argumentDefinitions": (v0/*: any*/), + "kind": "Operation", + "name": "SSOSignInPageQuery", + "selections": [ + (v1/*: any*/) + ] + }, + "params": { + "cacheID": "9b7766bbd9f160dc8bbcf34705641486", + "id": null, + "metadata": {}, + "name": "SSOSignInPageQuery", + "operationKind": "query", + "text": "query SSOSignInPageQuery(\n $email: EmailAddr!\n) {\n ssoLoginURL(email: $email)\n}\n" + } +}; +})(); + +(node as any).hash = "5749c48e67881c71bbc99a9e1120767f"; + +export default node; diff --git a/apps/console/src/pages/iam/auth/SignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx similarity index 84% rename from apps/console/src/pages/iam/auth/SignInPage.tsx rename to apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx index 5a97d1c4a..2b13d39d1 100644 --- a/apps/console/src/pages/iam/auth/SignInPage.tsx +++ b/apps/console/src/pages/iam/auth/sign-in/PasswordSignInPage.tsx @@ -1,13 +1,13 @@ import { useTranslate } from "@probo/i18n"; -import { Button, Field, useToast } from "@probo/ui"; +import { Button, Field, IconChevronLeft, useToast } from "@probo/ui"; import type { FormEventHandler } from "react"; import { useMutation } from "react-relay"; import { Link } from "react-router"; import { graphql } from "relay-runtime"; -import type { SignInPageMutation } from "/__generated__/iam/SignInPageMutation.graphql"; +import type { PasswordSignInPageMutation } from "/__generated__/iam/PasswordSignInPageMutation.graphql"; const signInMutation = graphql` - mutation SignInPageMutation($input: SignInInput!) { + mutation PasswordSignInPageMutation($input: SignInInput!) { signIn(input: $input) { session { id @@ -16,13 +16,13 @@ const signInMutation = graphql` } `; -// TODO initial screen + SAML login + after sign in URL state -export default function SignInPage() { +export default function PasswordSignInPage() { const { __ } = useTranslate(); const { toast } = useToast(); - const [signIn, isSigningIn] = useMutation(signInMutation); + const [signIn, isSigningIn] = + useMutation(signInMutation); const handlePasswordLogin: FormEventHandler = (e) => { e.preventDefault(); @@ -54,14 +54,13 @@ export default function SignInPage() { return (
- {/* */} +

{__("Login with Email")} diff --git a/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx new file mode 100644 index 000000000..d631f87ef --- /dev/null +++ b/apps/console/src/pages/iam/auth/sign-in/SSOSignInPage.tsx @@ -0,0 +1,133 @@ +import { useTranslate } from "@probo/i18n"; +import { Button, Field, IconChevronLeft, useToast } from "@probo/ui"; +import { useEffect, useState, type FormEventHandler } from "react"; +import { + usePreloadedQuery, + useQueryLoader, + type PreloadedQuery, +} from "react-relay"; +import { Link, useNavigate } from "react-router"; +import { graphql } from "relay-runtime"; +import type { SSOSignInPageQuery } from "/__generated__/iam/SSOSignInPageQuery.graphql"; + +const ssoAvailabilityQuery = graphql` + query SSOSignInPageQuery($email: EmailAddr!) { + ssoLoginURL(email: $email) @catch(to: RESULT) + } +`; + +export default function SSOSignInPage() { + const { __ } = useTranslate(); + + const [queryRef, loadQuery] = + useQueryLoader(ssoAvailabilityQuery); + const [checking, setChecking] = useState(false); + + const handleSSOCheck: FormEventHandler = (e) => { + e.preventDefault(); + setChecking(true); + const formData = new FormData(e.currentTarget); + const email = formData.get("email")?.toString(); + + if (!email) return; + + loadQuery({ email }, { fetchPolicy: "network-only" }); + }; + + return ( + <> + + + + {__("Back")} + + +

+ {__("Login with SSO")} +

+

+ {__("Enter your work email to continue with SSO")} +

+ + + + + +
+ {__("Don't have an account ?")}{" "} + + {__("Register")} + +
+ + + {queryRef && ( + + )} + + ); +} + +function NavigateToSSOLoginURL(props: { + queryRef: PreloadedQuery; + onSSOAvailabilityCheck: (checking: boolean) => void; +}) { + const { queryRef } = props; + + const { __ } = useTranslate(); + const { toast } = useToast(); + const navigate = useNavigate(); + + const { ssoLoginURL } = usePreloadedQuery( + ssoAvailabilityQuery, + queryRef, + ); + + useEffect(() => { + if (!ssoLoginURL.ok) { + console.log(ssoLoginURL); + toast({ + title: __("Error"), + description: + ssoLoginURL.errors[0] instanceof Error + ? ssoLoginURL.errors[0].message + : __("SSO not available for this email domain"), + variant: "error", + }); + + navigate("/auth/login"); + return; + } + + if (!ssoLoginURL.value) { + toast({ + title: __("Error"), + description: __("SSO not available for this email domain"), + variant: "error", + }); + return; + } + + window.location.href = ssoLoginURL.value; + }, [__, navigate, ssoLoginURL, toast]); + + return null; +} diff --git a/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx new file mode 100644 index 000000000..b68c9e343 --- /dev/null +++ b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx @@ -0,0 +1,57 @@ +import { useTranslate } from "@probo/i18n"; +import { Button } from "@probo/ui"; +import { Link } from "react-router"; + +export default function SignInPage() { + const { __ } = useTranslate(); + + return ( +
+

+ {__("Login to your account")} +

+

+ {__("Choose your login method")} +

+ + + +
+
+
+
+
+ + {__("Or")} + +
+
+ + + +
+ {__("Don't have an account ?")}{" "} + + {__("Register")} + +
+ +
+ {__("Forgot password?")}{" "} + + {__("Reset password")} + +
+
+ ); +} diff --git a/apps/console/src/routes.tsx b/apps/console/src/routes.tsx index 375a47a8e..f5aa07888 100644 --- a/apps/console/src/routes.tsx +++ b/apps/console/src/routes.tsx @@ -65,7 +65,17 @@ const routes = [ children: [ { path: "login", - Component: lazy(() => import("./pages/iam/auth/SignInPage")), + Component: lazy(() => import("./pages/iam/auth/sign-in/SignInPage")), + }, + { + path: "password-login", + Component: lazy( + () => import("./pages/iam/auth/sign-in/PasswordSignInPage"), + ), + }, + { + path: "sso-login", + Component: lazy(() => import("./pages/iam/auth/sign-in/SSOSignInPage")), }, { path: "register", @@ -78,7 +88,7 @@ const routes = [ { path: "signup-from-invitation", Component: lazy( - () => import("./pages/iam/auth/SignUpFromInvitationPage") + () => import("./pages/iam/auth/SignUpFromInvitationPage"), ), }, { @@ -100,13 +110,13 @@ const routes = [ { index: true, Component: lazy( - () => import("./pages/iam/memberships/MembershipsPageLoader") + () => import("./pages/iam/memberships/MembershipsPageLoader"), ), }, { path: "organizations/new", Component: lazy( - () => import("./pages/iam/organizations/NewOrganizationPage") + () => import("./pages/iam/organizations/NewOrganizationPage"), ), }, { @@ -123,7 +133,7 @@ const routes = [ path: "/organizations/:organizationId/employee", Fallback: () => "fallback employee...", Component: lazy( - () => import("./pages/organizations/employee/EmployeeLayoutLoader") + () => import("./pages/organizations/employee/EmployeeLayoutLoader"), ), ErrorBoundary: ErrorBoundary, children: [ @@ -132,7 +142,7 @@ const routes = [ // Component: () => "hello world", Component: lazy( () => - import("./pages/organizations/employee/EmployeeDocumentsPageLoader") + import("./pages/organizations/employee/EmployeeDocumentsPageLoader"), ), }, { @@ -140,7 +150,7 @@ const routes = [ ErrorBoundary: ErrorBoundary, Component: lazy( () => - import("./pages/organizations/employee/EmployeeDocumentSignaturePageLoader") + import("./pages/organizations/employee/EmployeeDocumentSignaturePageLoader"), ), }, ], @@ -148,7 +158,7 @@ const routes = [ { path: "/organizations/:organizationId", Component: lazy( - () => import("./pages/iam/memberships/MembershipLayoutLoader") + () => import("./pages/iam/memberships/MembershipLayoutLoader"), ), ErrorBoundary: ErrorBoundary, children: [ @@ -170,7 +180,7 @@ const routes = [ path: "settings", Fallback: PageSkeleton, Component: lazy( - () => import("./pages/iam/organizations/settings/SettingsLayout") + () => import("./pages/iam/organizations/settings/SettingsLayout"), ), children: [ { @@ -183,35 +193,35 @@ const routes = [ path: "general", Component: lazy( () => - import("./pages/iam/organizations/settings/GeneralSettingsPageLoader") + import("./pages/iam/organizations/settings/GeneralSettingsPageLoader"), ), }, { path: "members", Component: lazy( () => - import("./pages/iam/organizations/settings/MembersPageLoader") + import("./pages/iam/organizations/settings/MembersPageLoader"), ), }, { path: "domain", Component: lazy( () => - import("./pages/organizations/settings/DomainSettingsPageLoader") + import("./pages/organizations/settings/DomainSettingsPageLoader"), ), }, { path: "saml-sso", Component: lazy( () => - import("./pages/iam/organizations/settings/SAMLSettingsPageLoader") + import("./pages/iam/organizations/settings/SAMLSettingsPageLoader"), ), }, { path: "scim", Component: lazy( () => - import("./pages/iam/organizations/settings/SCIMSettingsPageLoader") + import("./pages/iam/organizations/settings/SCIMSettingsPageLoader"), ), }, ], diff --git a/packages/ui/src/Atoms/Link/Link.tsx b/packages/ui/src/Atoms/Link/Link.tsx deleted file mode 100644 index e69de29bb..000000000 diff --git a/pkg/coredata/saml_configuration.go b/pkg/coredata/saml_configuration.go index d1c0dbcf4..f9c2bb612 100644 --- a/pkg/coredata/saml_configuration.go +++ b/pkg/coredata/saml_configuration.go @@ -513,6 +513,73 @@ WHERE return count, nil } +func (s *SAMLConfigurations) LoadVerifiedByEmailDomain(ctx context.Context, conn pg.Conn, emailDomain string) error { + q := ` +SELECT + id, + organization_id, + email_domain, + enforcement_policy, + idp_entity_id, + idp_sso_url, + idp_certificate, + idp_metadata_url, + attribute_email, + attribute_firstname, + attribute_lastname, + attribute_role, + auto_signup_enabled, + domain_verification_token, + domain_verified_at, + created_at, + updated_at +FROM + iam_saml_configurations +WHERE + email_domain = @email_domain + AND domain_verified_at IS NOT NULL +ORDER BY email_domain ASC; +` + + rows, err := conn.Query(ctx, q, pgx.StrictNamedArgs{"email_domain": emailDomain}) + if err != nil { + return fmt.Errorf("cannot query iam_saml_configurations: %w", err) + } + + samlConfigurations, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[SAMLConfiguration]) + if err != nil { + return fmt.Errorf("cannot collect saml_configurations: %w", err) + } + + *s = samlConfigurations + + return nil +} + +func (s *SAMLConfigurations) CountVerifiedByEmailDomain( + ctx context.Context, + conn pg.Conn, + emailDomain string, +) (int, error) { + q := ` +SELECT + COUNT(*) +FROM + iam_saml_configurations +WHERE + email_domain = @email_domain + AND domain_verified_at IS NOT NULL +` + + row := conn.QueryRow(ctx, q, pgx.StrictNamedArgs{"email_domain": emailDomain}) + var count int + if err := row.Scan(&count); err != nil { + return 0, fmt.Errorf("cannot count SAML configurations: %w", err) + } + + return count, nil +} + func (s *SAMLConfigurations) LoadUnverified( ctx context.Context, conn pg.Conn, diff --git a/pkg/iam/account_service.go b/pkg/iam/account_service.go index 836e74190..4e4fe528a 100644 --- a/pkg/iam/account_service.go +++ b/pkg/iam/account_service.go @@ -856,3 +856,54 @@ func (s AccountService) GetProfileForMembership(ctx context.Context, membershipI return profile, nil } + +func (s AccountService) ListSAMLConfigurationsForEmail( + ctx context.Context, + email mail.Addr, +) (coredata.SAMLConfigurations, error) { + samlConfigurations := coredata.SAMLConfigurations{} + + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + err := samlConfigurations.LoadVerifiedByEmailDomain(ctx, conn, email.Domain()) + if err != nil { + return fmt.Errorf("cannot load saml configurations: %w", err) + } + + return nil + }, + ) + if err != nil { + return nil, err + } + + return samlConfigurations, nil +} + +func (s AccountService) CountSAMLConfigurationsForEmail( + ctx context.Context, + email mail.Addr, +) (int, error) { + var ( + count int + samlConfigurations coredata.SAMLConfigurations + ) + + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + count, err = samlConfigurations.CountVerifiedByEmailDomain(ctx, conn, email.Domain()) + if err != nil { + return fmt.Errorf("cannot count saml configurations: %w", err) + } + + return nil + }, + ) + if err != nil { + return 0, err + } + + return count, nil +} diff --git a/pkg/server/api/connect/v1/schema.graphql b/pkg/server/api/connect/v1/schema.graphql index 6bbb65bfe..89e86a38c 100644 --- a/pkg/server/api/connect/v1/schema.graphql +++ b/pkg/server/api/connect/v1/schema.graphql @@ -54,7 +54,8 @@ interface Node { type Query { node(id: ID!): Node @session(required: PRESENT) viewer: Identity @session(required: PRESENT) - checkSSOAvailability(email: String!): SSOAvailability! + ssoLoginURL(email: EmailAddr!): String + @goField(forceResolver: true) @session(required: NONE) } @@ -349,12 +350,6 @@ type SAMLAttributeMappings { role: String! } -type SSOAvailability { - available: Boolean! - samlConfigId: ID - organizationId: ID -} - type SCIMConfiguration implements Node { id: ID! endpointUrl: String! @goField(forceResolver: true) diff --git a/pkg/server/api/connect/v1/schema/schema.go b/pkg/server/api/connect/v1/schema/schema.go index 61bed1ef2..7e064cd30 100644 --- a/pkg/server/api/connect/v1/schema/schema.go +++ b/pkg/server/api/connect/v1/schema/schema.go @@ -294,9 +294,9 @@ type ComplexityRoot struct { } Query struct { - CheckSSOAvailability func(childComplexity int, email string) int - Node func(childComplexity int, id gid.GID) int - Viewer func(childComplexity int) int + Node func(childComplexity int, id gid.GID) int + SsoLoginURL func(childComplexity int, email mail.Addr) int + Viewer func(childComplexity int) int } RegenerateSCIMTokenPayload struct { @@ -399,12 +399,6 @@ type ComplexityRoot struct { Node func(childComplexity int) int } - SSOAvailability struct { - Available func(childComplexity int) int - OrganizationID func(childComplexity int) int - SamlConfigID func(childComplexity int) int - } - Session struct { CreatedAt func(childComplexity int) int ExpiresAt func(childComplexity int) int @@ -541,7 +535,7 @@ type PersonalAPIKeyConnectionResolver interface { type QueryResolver interface { Node(ctx context.Context, id gid.GID) (types.Node, error) Viewer(ctx context.Context) (*types.Identity, error) - CheckSSOAvailability(ctx context.Context, email string) (*types.SSOAvailability, error) + SsoLoginURL(ctx context.Context, email mail.Addr) (*string, error) } type SAMLConfigurationResolver interface { TestLoginURL(ctx context.Context, obj *types.SAMLConfiguration) (string, error) @@ -1593,17 +1587,6 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.PersonalAPIKeyEdge.Node(childComplexity), true - case "Query.checkSSOAvailability": - if e.complexity.Query.CheckSSOAvailability == nil { - break - } - - args, err := ec.field_Query_checkSSOAvailability_args(ctx, rawArgs) - if err != nil { - return 0, false - } - - return e.complexity.Query.CheckSSOAvailability(childComplexity, args["email"].(string)), true case "Query.node": if e.complexity.Query.Node == nil { break @@ -1615,6 +1598,17 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin } return e.complexity.Query.Node(childComplexity, args["id"].(gid.GID)), true + case "Query.ssoLoginURL": + if e.complexity.Query.SsoLoginURL == nil { + break + } + + args, err := ec.field_Query_ssoLoginURL_args(ctx, rawArgs) + if err != nil { + return 0, false + } + + return e.complexity.Query.SsoLoginURL(childComplexity, args["email"].(mail.Addr)), true case "Query.viewer": if e.complexity.Query.Viewer == nil { break @@ -1987,25 +1981,6 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.SCIMEventEdge.Node(childComplexity), true - case "SSOAvailability.available": - if e.complexity.SSOAvailability.Available == nil { - break - } - - return e.complexity.SSOAvailability.Available(childComplexity), true - case "SSOAvailability.organizationId": - if e.complexity.SSOAvailability.OrganizationID == nil { - break - } - - return e.complexity.SSOAvailability.OrganizationID(childComplexity), true - case "SSOAvailability.samlConfigId": - if e.complexity.SSOAvailability.SamlConfigID == nil { - break - } - - return e.complexity.SSOAvailability.SamlConfigID(childComplexity), true - case "Session.createdAt": if e.complexity.Session.CreatedAt == nil { break @@ -2347,7 +2322,8 @@ interface Node { type Query { node(id: ID!): Node @session(required: PRESENT) viewer: Identity @session(required: PRESENT) - checkSSOAvailability(email: String!): SSOAvailability! + ssoLoginURL(email: EmailAddr!): String + @goField(forceResolver: true) @session(required: NONE) } @@ -2642,12 +2618,6 @@ type SAMLAttributeMappings { role: String! } -type SSOAvailability { - available: Boolean! - samlConfigId: ID - organizationId: ID -} - type SCIMConfiguration implements Node { id: ID! endpointUrl: String! @goField(forceResolver: true) @@ -3779,17 +3749,6 @@ func (ec *executionContext) field_Query___type_args(ctx context.Context, rawArgs return args, nil } -func (ec *executionContext) field_Query_checkSSOAvailability_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { - var err error - args := map[string]any{} - arg0, err := graphql.ProcessArgField(ctx, rawArgs, "email", ec.unmarshalNString2string) - if err != nil { - return nil, err - } - args["email"] = arg0 - return args, nil -} - func (ec *executionContext) field_Query_node_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -3801,6 +3760,17 @@ func (ec *executionContext) field_Query_node_args(ctx context.Context, rawArgs m return args, nil } +func (ec *executionContext) field_Query_ssoLoginURL_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { + var err error + args := map[string]any{} + arg0, err := graphql.ProcessArgField(ctx, rawArgs, "email", ec.unmarshalNEmailAddr2goᚗproboᚗincᚋproboᚋpkgᚋmailᚐAddr) + if err != nil { + return nil, err + } + args["email"] = arg0 + return args, nil +} + func (ec *executionContext) field_SAMLConfiguration_permission_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -9531,15 +9501,15 @@ func (ec *executionContext) fieldContext_Query_viewer(_ context.Context, field g return fc, nil } -func (ec *executionContext) _Query_checkSSOAvailability(ctx context.Context, field graphql.CollectedField) (ret graphql.Marshaler) { +func (ec *executionContext) _Query_ssoLoginURL(ctx context.Context, field graphql.CollectedField) (ret graphql.Marshaler) { return graphql.ResolveField( ctx, ec.OperationContext, field, - ec.fieldContext_Query_checkSSOAvailability, + ec.fieldContext_Query_ssoLoginURL, func(ctx context.Context) (any, error) { fc := graphql.GetFieldContext(ctx) - return ec.resolvers.Query().CheckSSOAvailability(ctx, fc.Args["email"].(string)) + return ec.resolvers.Query().SsoLoginURL(ctx, fc.Args["email"].(mail.Addr)) }, func(ctx context.Context, next graphql.Resolver) graphql.Resolver { directive0 := next @@ -9547,11 +9517,11 @@ func (ec *executionContext) _Query_checkSSOAvailability(ctx context.Context, fie directive1 := func(ctx context.Context) (any, error) { required, err := ec.unmarshalNSessionRequirement2goᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSessionRequirement(ctx, "NONE") if err != nil { - var zeroVal *types.SSOAvailability + var zeroVal *string return zeroVal, err } if ec.directives.Session == nil { - var zeroVal *types.SSOAvailability + var zeroVal *string return zeroVal, errors.New("directive session is not implemented") } return ec.directives.Session(ctx, nil, directive0, required) @@ -9560,28 +9530,20 @@ func (ec *executionContext) _Query_checkSSOAvailability(ctx context.Context, fie next = directive1 return next }, - ec.marshalNSSOAvailability2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSSOAvailability, - true, + ec.marshalOString2ᚖstring, true, + false, ) } -func (ec *executionContext) fieldContext_Query_checkSSOAvailability(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { +func (ec *executionContext) fieldContext_Query_ssoLoginURL(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { fc = &graphql.FieldContext{ Object: "Query", Field: field, IsMethod: true, IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - switch field.Name { - case "available": - return ec.fieldContext_SSOAvailability_available(ctx, field) - case "samlConfigId": - return ec.fieldContext_SSOAvailability_samlConfigId(ctx, field) - case "organizationId": - return ec.fieldContext_SSOAvailability_organizationId(ctx, field) - } - return nil, fmt.Errorf("no field named %q was found under type SSOAvailability", field.Name) + return nil, errors.New("field of type String does not have child fields") }, } defer func() { @@ -9591,7 +9553,7 @@ func (ec *executionContext) fieldContext_Query_checkSSOAvailability(ctx context. } }() ctx = graphql.WithFieldContext(ctx, fc) - if fc.Args, err = ec.field_Query_checkSSOAvailability_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { + if fc.Args, err = ec.field_Query_ssoLoginURL_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { ec.Error(ctx, err) return fc, err } @@ -11579,93 +11541,6 @@ func (ec *executionContext) fieldContext_SCIMEventEdge_cursor(_ context.Context, return fc, nil } -func (ec *executionContext) _SSOAvailability_available(ctx context.Context, field graphql.CollectedField, obj *types.SSOAvailability) (ret graphql.Marshaler) { - return graphql.ResolveField( - ctx, - ec.OperationContext, - field, - ec.fieldContext_SSOAvailability_available, - func(ctx context.Context) (any, error) { - return obj.Available, nil - }, - nil, - ec.marshalNBoolean2bool, - true, - true, - ) -} - -func (ec *executionContext) fieldContext_SSOAvailability_available(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { - fc = &graphql.FieldContext{ - Object: "SSOAvailability", - 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 fc, nil -} - -func (ec *executionContext) _SSOAvailability_samlConfigId(ctx context.Context, field graphql.CollectedField, obj *types.SSOAvailability) (ret graphql.Marshaler) { - return graphql.ResolveField( - ctx, - ec.OperationContext, - field, - ec.fieldContext_SSOAvailability_samlConfigId, - func(ctx context.Context) (any, error) { - return obj.SamlConfigID, nil - }, - nil, - ec.marshalOID2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋgidᚐGID, - true, - false, - ) -} - -func (ec *executionContext) fieldContext_SSOAvailability_samlConfigId(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { - fc = &graphql.FieldContext{ - Object: "SSOAvailability", - Field: field, - IsMethod: false, - IsResolver: false, - Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - return nil, errors.New("field of type ID does not have child fields") - }, - } - return fc, nil -} - -func (ec *executionContext) _SSOAvailability_organizationId(ctx context.Context, field graphql.CollectedField, obj *types.SSOAvailability) (ret graphql.Marshaler) { - return graphql.ResolveField( - ctx, - ec.OperationContext, - field, - ec.fieldContext_SSOAvailability_organizationId, - func(ctx context.Context) (any, error) { - return obj.OrganizationID, nil - }, - nil, - ec.marshalOID2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋgidᚐGID, - true, - false, - ) -} - -func (ec *executionContext) fieldContext_SSOAvailability_organizationId(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { - fc = &graphql.FieldContext{ - Object: "SSOAvailability", - Field: field, - IsMethod: false, - IsResolver: false, - Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - return nil, errors.New("field of type ID does not have child fields") - }, - } - return fc, nil -} - func (ec *executionContext) _Session_id(ctx context.Context, field graphql.CollectedField, obj *types.Session) (ret graphql.Marshaler) { return graphql.ResolveField( ctx, @@ -17799,19 +17674,16 @@ func (ec *executionContext) _Query(ctx context.Context, sel ast.SelectionSet) gr } out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return rrm(innerCtx) }) - case "checkSSOAvailability": + case "ssoLoginURL": field := field - innerFunc := func(ctx context.Context, fs *graphql.FieldSet) (res graphql.Marshaler) { + innerFunc := func(ctx context.Context, _ *graphql.FieldSet) (res graphql.Marshaler) { defer func() { if r := recover(); r != nil { ec.Error(ctx, ec.Recover(ctx, r)) } }() - res = ec._Query_checkSSOAvailability(ctx, field) - if res == graphql.Null { - atomic.AddUint32(&fs.Invalids, 1) - } + res = ec._Query_ssoLoginURL(ctx, field) return res } @@ -18917,49 +18789,6 @@ func (ec *executionContext) _SCIMEventEdge(ctx context.Context, sel ast.Selectio return out } -var sSOAvailabilityImplementors = []string{"SSOAvailability"} - -func (ec *executionContext) _SSOAvailability(ctx context.Context, sel ast.SelectionSet, obj *types.SSOAvailability) graphql.Marshaler { - fields := graphql.CollectFields(ec.OperationContext, sel, sSOAvailabilityImplementors) - - out := graphql.NewFieldSet(fields) - deferred := make(map[string]*graphql.FieldSet) - for i, field := range fields { - switch field.Name { - case "__typename": - out.Values[i] = graphql.MarshalString("SSOAvailability") - case "available": - out.Values[i] = ec._SSOAvailability_available(ctx, field, obj) - if out.Values[i] == graphql.Null { - out.Invalids++ - } - case "samlConfigId": - out.Values[i] = ec._SSOAvailability_samlConfigId(ctx, field, obj) - case "organizationId": - out.Values[i] = ec._SSOAvailability_organizationId(ctx, field, obj) - default: - panic("unknown field " + strconv.Quote(field.Name)) - } - } - out.Dispatch(ctx) - if out.Invalids > 0 { - return graphql.Null - } - - atomic.AddInt32(&ec.deferred, int32(len(deferred))) - - for label, dfs := range deferred { - ec.processDeferredGroup(graphql.DeferredGroup{ - Label: label, - Path: graphql.GetPath(ctx), - FieldSet: dfs, - Context: ctx, - }) - } - - return out -} - var sessionImplementors = []string{"Session", "Node"} func (ec *executionContext) _Session(ctx context.Context, sel ast.SelectionSet, obj *types.Session) graphql.Marshaler { @@ -20706,20 +20535,6 @@ var ( } ) -func (ec *executionContext) marshalNSSOAvailability2goᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSSOAvailability(ctx context.Context, sel ast.SelectionSet, v types.SSOAvailability) graphql.Marshaler { - return ec._SSOAvailability(ctx, sel, &v) -} - -func (ec *executionContext) marshalNSSOAvailability2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSSOAvailability(ctx context.Context, sel ast.SelectionSet, v *types.SSOAvailability) graphql.Marshaler { - if v == nil { - if !graphql.HasFieldError(ctx, graphql.GetFieldContext(ctx)) { - graphql.AddErrorf(ctx, "the requested element is null which the schema does not allow") - } - return graphql.Null - } - return ec._SSOAvailability(ctx, sel, v) -} - func (ec *executionContext) marshalNSession2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSession(ctx context.Context, sel ast.SelectionSet, v *types.Session) graphql.Marshaler { if v == nil { if !graphql.HasFieldError(ctx, graphql.GetFieldContext(ctx)) { @@ -21306,24 +21121,6 @@ func (ec *executionContext) marshalOForgotPasswordPayload2ᚖgoᚗproboᚗincᚋ return ec._ForgotPasswordPayload(ctx, sel, v) } -func (ec *executionContext) unmarshalOID2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋgidᚐGID(ctx context.Context, v any) (*gid.GID, error) { - if v == nil { - return nil, nil - } - res, err := gid1.UnmarshalGIDScalar(v) - return &res, graphql.ErrorOnPath(ctx, err) -} - -func (ec *executionContext) marshalOID2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋgidᚐGID(ctx context.Context, sel ast.SelectionSet, v *gid.GID) graphql.Marshaler { - if v == nil { - return graphql.Null - } - _ = sel - _ = ctx - res := gid1.MarshalGIDScalar(*v) - return res -} - func (ec *executionContext) marshalOIdentity2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐIdentity(ctx context.Context, sel ast.SelectionSet, v *types.Identity) graphql.Marshaler { if v == nil { return graphql.Null diff --git a/pkg/server/api/connect/v1/types/types.go b/pkg/server/api/connect/v1/types/types.go index 9d46c7e4b..cf6fe5c19 100644 --- a/pkg/server/api/connect/v1/types/types.go +++ b/pkg/server/api/connect/v1/types/types.go @@ -430,12 +430,6 @@ type SCIMEventEdge struct { Cursor page.CursorKey `json:"cursor"` } -type SSOAvailability struct { - Available bool `json:"available"` - SamlConfigID *gid.GID `json:"samlConfigId,omitempty"` - OrganizationID *gid.GID `json:"organizationId,omitempty"` -} - type Session struct { ID gid.GID `json:"id"` Identity *Identity `json:"identity,omitempty"` diff --git a/pkg/server/api/connect/v1/v1_resolver.go b/pkg/server/api/connect/v1/v1_resolver.go index cfa14c826..fcb722a69 100644 --- a/pkg/server/api/connect/v1/v1_resolver.go +++ b/pkg/server/api/connect/v1/v1_resolver.go @@ -11,11 +11,13 @@ import ( "fmt" "time" + "github.com/99designs/gqlgen/graphql" "github.com/vektah/gqlparser/v2/gqlerror" "go.gearno.de/kit/log" "go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/iam" + "go.probo.inc/probo/pkg/mail" "go.probo.inc/probo/pkg/page" "go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/server/api/connect/v1/schema" @@ -1532,9 +1534,43 @@ func (r *queryResolver) Viewer(ctx context.Context) (*types.Identity, error) { }, nil } -// CheckSSOAvailability is the resolver for the checkSSOAvailability field. -func (r *queryResolver) CheckSSOAvailability(ctx context.Context, email string) (*types.SSOAvailability, error) { - panic(fmt.Errorf("not implemented: CheckSSOAvailability - checkSSOAvailability")) +// SsoLoginURL is the resolver for the ssoLoginURL field. +func (r *queryResolver) SsoLoginURL(ctx context.Context, email mail.Addr) (*string, error) { + count, err := r.iam.AccountService.CountSAMLConfigurationsForEmail(ctx, email) + if err != nil { + r.logger.ErrorCtx(ctx, "cannot count SAML configurations for email", log.Error(err)) + return nil, gqlutils.InternalServerError(ctx) + } + + if count != 1 { + if count == 0 { + graphql.AddError(ctx, graphql.ErrorOnPath( + ctx, + fmt.Errorf("no SAML configuration for email"), + )) + + return nil, nil + } + + graphql.AddError( + ctx, + graphql.ErrorOnPath( + ctx, + fmt.Errorf("multiple SSO configurations found for this domain. Please use your organization-specific SSO login URL"), + ), + ) + return nil, nil + } + + samlConfigs, err := r.iam.AccountService.ListSAMLConfigurationsForEmail(ctx, email) + if err != nil { + r.logger.ErrorCtx(ctx, "cannot list SAML configurations for email", log.Error(err)) + return nil, gqlutils.InternalServerError(ctx) + } + + samlConfig := samlConfigs[0] + loginURL := r.baseURL.WithPath("/api/connect/v1/saml/2.0/" + samlConfig.ID.String()).MustString() + return &loginURL, nil } // TestLoginURL is the resolver for the testLoginUrl field.