diff --git a/apps/console/src/components/organizations/InviteUserDialog.tsx b/apps/console/src/components/organizations/InviteUserDialog.tsx index 146cfce6c..914fbcfdf 100644 --- a/apps/console/src/components/organizations/InviteUserDialog.tsx +++ b/apps/console/src/components/organizations/InviteUserDialog.tsx @@ -45,14 +45,15 @@ const schema = z.object({ type Props = PropsWithChildren & { connectionId?: string; + onRefetch: () => void; }; -export function InviteUserDialog({ children, connectionId }: Props) { +export function InviteUserDialog({ children, connectionId, onRefetch }: Props) { const { __ } = useTranslate(); const organizationId = useOrganizationId(); const [inviteUser, isInviting] = useMutationWithToasts(inviteMutation, { - successMessage: __("User invited successfully"), - errorMessage: __("Failed to invite user"), + successMessage: __("Invitation sent successfully"), + errorMessage: __("Failed to send invitation"), }); const { register, handleSubmit, formState, reset, control } = useFormWithSchema( schema, @@ -72,9 +73,10 @@ export function InviteUserDialog({ children, connectionId }: Props) { }, connections: connectionId ? [connectionId] : ["SettingsPageInvitations_invitations"], }, - onSuccess: () => { + onCompleted: () => { reset(); dialogRef.current?.close(); + onRefetch(); }, }); }); diff --git a/apps/console/src/hooks/graph/VendorGraph.ts b/apps/console/src/hooks/graph/VendorGraph.ts index 46e89990b..31e879039 100644 --- a/apps/console/src/hooks/graph/VendorGraph.ts +++ b/apps/console/src/hooks/graph/VendorGraph.ts @@ -147,7 +147,7 @@ export const paginatedVendorsFragment = graphql` `; export const vendorNodeQuery = graphql` - query VendorGraphNodeQuery($vendorId: ID!, $organizationId: ID!) { + query VendorGraphNodeQuery($vendorId: ID!) { node(id: $vendorId) { ... on Vendor { id @@ -166,9 +166,6 @@ export const vendorNodeQuery = graphql` viewer { user { id - people(organizationId: $organizationId) { - id - } } } } diff --git a/apps/console/src/hooks/graph/__generated__/VendorGraphNodeQuery.graphql.ts b/apps/console/src/hooks/graph/__generated__/VendorGraphNodeQuery.graphql.ts index c132831e8..a2b604aa8 100644 --- a/apps/console/src/hooks/graph/__generated__/VendorGraphNodeQuery.graphql.ts +++ b/apps/console/src/hooks/graph/__generated__/VendorGraphNodeQuery.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -11,7 +11,6 @@ import { ConcreteRequest } from 'relay-runtime'; import { FragmentRefs } from "relay-runtime"; export type VendorGraphNodeQuery$variables = { - organizationId: string; vendorId: string; }; export type VendorGraphNodeQuery$data = { @@ -25,9 +24,6 @@ export type VendorGraphNodeQuery$data = { readonly viewer: { readonly user: { readonly id: string; - readonly people: { - readonly id: string; - } | null | undefined; }; }; }; @@ -37,153 +33,132 @@ export type VendorGraphNodeQuery = { }; const node: ConcreteRequest = (function(){ -var v0 = { - "defaultValue": null, - "kind": "LocalArgument", - "name": "organizationId" -}, -v1 = { - "defaultValue": null, - "kind": "LocalArgument", - "name": "vendorId" -}, -v2 = [ +var v0 = [ + { + "defaultValue": null, + "kind": "LocalArgument", + "name": "vendorId" + } +], +v1 = [ { "kind": "Variable", "name": "id", "variableName": "vendorId" } ], -v3 = { +v2 = { "alias": null, "args": null, "kind": "ScalarField", "name": "id", "storageKey": null }, -v4 = { +v3 = { "alias": null, "args": null, "kind": "ScalarField", "name": "snapshotId", "storageKey": null }, -v5 = { +v4 = { "alias": null, "args": null, "kind": "ScalarField", "name": "name", "storageKey": null }, -v6 = { +v5 = { "alias": null, "args": null, "kind": "ScalarField", "name": "websiteUrl", "storageKey": null }, -v7 = [ - (v3/*: any*/) +v6 = [ + (v2/*: any*/) ], -v8 = { +v7 = { "alias": null, "args": null, "concreteType": "User", "kind": "LinkedField", "name": "user", "plural": false, - "selections": [ - (v3/*: any*/), - { - "alias": null, - "args": [ - { - "kind": "Variable", - "name": "organizationId", - "variableName": "organizationId" - } - ], - "concreteType": "People", - "kind": "LinkedField", - "name": "people", - "plural": false, - "selections": (v7/*: any*/), - "storageKey": null - } - ], + "selections": (v6/*: any*/), "storageKey": null }, -v9 = { +v8 = { "alias": null, "args": null, "kind": "ScalarField", "name": "__typename", "storageKey": null }, -v10 = { +v9 = { "alias": null, "args": null, "kind": "ScalarField", "name": "description", "storageKey": null }, -v11 = [ +v10 = [ { "kind": "Literal", "name": "first", "value": 50 } ], -v12 = { +v11 = { "alias": null, "args": null, "kind": "ScalarField", "name": "validUntil", "storageKey": null }, -v13 = { +v12 = { "alias": null, "args": null, "kind": "ScalarField", "name": "fileName", "storageKey": null }, -v14 = { +v13 = { "alias": null, "args": null, "kind": "ScalarField", "name": "cursor", "storageKey": null }, -v15 = { +v14 = { "alias": null, "args": null, "kind": "ScalarField", "name": "endCursor", "storageKey": null }, -v16 = { +v15 = { "alias": null, "args": null, "kind": "ScalarField", "name": "hasNextPage", "storageKey": null }, -v17 = { +v16 = { "alias": null, "args": null, "kind": "ScalarField", "name": "hasPreviousPage", "storageKey": null }, -v18 = { +v17 = { "alias": null, "args": null, "kind": "ScalarField", "name": "startCursor", "storageKey": null }, -v19 = { +v18 = { "alias": null, "args": null, "concreteType": "PageInfo", @@ -191,14 +166,14 @@ v19 = { "name": "pageInfo", "plural": false, "selections": [ + (v14/*: any*/), (v15/*: any*/), (v16/*: any*/), - (v17/*: any*/), - (v18/*: any*/) + (v17/*: any*/) ], "storageKey": null }, -v20 = { +v19 = { "kind": "ClientExtension", "selections": [ { @@ -210,26 +185,26 @@ v20 = { } ] }, -v21 = [ +v20 = [ "orderBy" ], -v22 = { +v21 = { "alias": null, "args": null, "kind": "ScalarField", "name": "createdAt", "storageKey": null }, -v23 = { +v22 = { "alias": null, "args": null, "kind": "ScalarField", "name": "updatedAt", "storageKey": null }, -v24 = [ - (v3/*: any*/), - (v13/*: any*/), +v23 = [ + (v2/*: any*/), + (v12/*: any*/), { "alias": null, "args": null, @@ -244,22 +219,19 @@ v24 = [ "name": "validFrom", "storageKey": null }, - (v12/*: any*/), - (v22/*: any*/) + (v11/*: any*/), + (v21/*: any*/) ]; return { "fragment": { - "argumentDefinitions": [ - (v0/*: any*/), - (v1/*: any*/) - ], + "argumentDefinitions": (v0/*: any*/), "kind": "Fragment", "metadata": null, "name": "VendorGraphNodeQuery", "selections": [ { "alias": null, - "args": (v2/*: any*/), + "args": (v1/*: any*/), "concreteType": null, "kind": "LinkedField", "name": "node", @@ -268,10 +240,10 @@ return { { "kind": "InlineFragment", "selections": [ + (v2/*: any*/), (v3/*: any*/), (v4/*: any*/), (v5/*: any*/), - (v6/*: any*/), { "args": null, "kind": "FragmentSpread", @@ -322,7 +294,7 @@ return { "name": "viewer", "plural": false, "selections": [ - (v8/*: any*/) + (v7/*: any*/) ], "storageKey": null } @@ -332,30 +304,27 @@ return { }, "kind": "Request", "operation": { - "argumentDefinitions": [ - (v1/*: any*/), - (v0/*: any*/) - ], + "argumentDefinitions": (v0/*: any*/), "kind": "Operation", "name": "VendorGraphNodeQuery", "selections": [ { "alias": null, - "args": (v2/*: any*/), + "args": (v1/*: any*/), "concreteType": null, "kind": "LinkedField", "name": "node", "plural": false, "selections": [ - (v9/*: any*/), - (v3/*: any*/), + (v8/*: any*/), + (v2/*: any*/), { "kind": "InlineFragment", "selections": [ + (v3/*: any*/), (v4/*: any*/), (v5/*: any*/), - (v6/*: any*/), - (v10/*: any*/), + (v9/*: any*/), { "alias": null, "args": null, @@ -447,7 +416,7 @@ return { "kind": "LinkedField", "name": "businessOwner", "plural": false, - "selections": (v7/*: any*/), + "selections": (v6/*: any*/), "storageKey": null }, { @@ -457,12 +426,12 @@ return { "kind": "LinkedField", "name": "securityOwner", "plural": false, - "selections": (v7/*: any*/), + "selections": (v6/*: any*/), "storageKey": null }, { "alias": null, - "args": (v11/*: any*/), + "args": (v10/*: any*/), "concreteType": "VendorComplianceReportConnection", "kind": "LinkedField", "name": "complianceReports", @@ -484,7 +453,7 @@ return { "name": "node", "plural": false, "selections": [ - (v3/*: any*/), + (v2/*: any*/), { "alias": null, "args": null, @@ -492,7 +461,7 @@ return { "name": "reportDate", "storageKey": null }, - (v12/*: any*/), + (v11/*: any*/), { "alias": null, "args": null, @@ -508,7 +477,7 @@ return { "name": "file", "plural": false, "selections": [ - (v13/*: any*/), + (v12/*: any*/), { "alias": null, "args": null, @@ -523,27 +492,27 @@ return { "name": "size", "storageKey": null }, - (v3/*: any*/) + (v2/*: any*/) ], "storageKey": null }, - (v9/*: any*/) + (v8/*: any*/) ], "storageKey": null }, - (v14/*: any*/) + (v13/*: any*/) ], "storageKey": null }, - (v19/*: any*/), - (v20/*: any*/) + (v18/*: any*/), + (v19/*: any*/) ], "storageKey": "complianceReports(first:50)" }, { "alias": null, - "args": (v11/*: any*/), - "filters": (v21/*: any*/), + "args": (v10/*: any*/), + "filters": (v20/*: any*/), "handle": "connection", "key": "VendorComplianceTabFragment_complianceReports", "kind": "LinkedHandle", @@ -551,7 +520,7 @@ return { }, { "alias": null, - "args": (v11/*: any*/), + "args": (v10/*: any*/), "concreteType": "VendorContactConnection", "kind": "LinkedField", "name": "contacts", @@ -573,7 +542,7 @@ return { "name": "node", "plural": false, "selections": [ - (v3/*: any*/), + (v2/*: any*/), { "alias": null, "args": null, @@ -602,25 +571,25 @@ return { "name": "role", "storageKey": null }, + (v21/*: any*/), (v22/*: any*/), - (v23/*: any*/), - (v9/*: any*/) + (v8/*: any*/) ], "storageKey": null }, - (v14/*: any*/) + (v13/*: any*/) ], "storageKey": null }, - (v19/*: any*/), - (v20/*: any*/) + (v18/*: any*/), + (v19/*: any*/) ], "storageKey": "contacts(first:50)" }, { "alias": null, - "args": (v11/*: any*/), - "filters": (v21/*: any*/), + "args": (v10/*: any*/), + "filters": (v20/*: any*/), "handle": "connection", "key": "VendorContactsTabFragment_contacts", "kind": "LinkedHandle", @@ -628,7 +597,7 @@ return { }, { "alias": null, - "args": (v11/*: any*/), + "args": (v10/*: any*/), "concreteType": "VendorServiceConnection", "kind": "LinkedField", "name": "services", @@ -650,28 +619,28 @@ return { "name": "node", "plural": false, "selections": [ - (v3/*: any*/), - (v5/*: any*/), - (v10/*: any*/), + (v2/*: any*/), + (v4/*: any*/), + (v9/*: any*/), + (v21/*: any*/), (v22/*: any*/), - (v23/*: any*/), - (v9/*: any*/) + (v8/*: any*/) ], "storageKey": null }, - (v14/*: any*/) + (v13/*: any*/) ], "storageKey": null }, - (v19/*: any*/), - (v20/*: any*/) + (v18/*: any*/), + (v19/*: any*/) ], "storageKey": "services(first:50)" }, { "alias": null, - "args": (v11/*: any*/), - "filters": (v21/*: any*/), + "args": (v10/*: any*/), + "filters": (v20/*: any*/), "handle": "connection", "key": "VendorServicesTabFragment_services", "kind": "LinkedHandle", @@ -679,7 +648,7 @@ return { }, { "alias": null, - "args": (v11/*: any*/), + "args": (v10/*: any*/), "concreteType": "VendorRiskAssessmentConnection", "kind": "LinkedField", "name": "riskAssessments", @@ -701,8 +670,8 @@ return { "name": "node", "plural": false, "selections": [ - (v3/*: any*/), - (v22/*: any*/), + (v2/*: any*/), + (v21/*: any*/), { "alias": null, "args": null, @@ -731,11 +700,11 @@ return { "name": "notes", "storageKey": null }, - (v9/*: any*/) + (v8/*: any*/) ], "storageKey": null }, - (v14/*: any*/) + (v13/*: any*/) ], "storageKey": null }, @@ -747,21 +716,21 @@ return { "name": "pageInfo", "plural": false, "selections": [ - (v16/*: any*/), (v15/*: any*/), - (v17/*: any*/), - (v18/*: any*/) + (v14/*: any*/), + (v16/*: any*/), + (v17/*: any*/) ], "storageKey": null }, - (v20/*: any*/) + (v19/*: any*/) ], "storageKey": "riskAssessments(first:50)" }, { "alias": null, - "args": (v11/*: any*/), - "filters": (v21/*: any*/), + "args": (v10/*: any*/), + "filters": (v20/*: any*/), "handle": "connection", "key": "VendorRiskAssessmentTabFragment_riskAssessments", "kind": "LinkedHandle", @@ -774,7 +743,7 @@ return { "kind": "LinkedField", "name": "businessAssociateAgreement", "plural": false, - "selections": (v24/*: any*/), + "selections": (v23/*: any*/), "storageKey": null }, { @@ -784,7 +753,7 @@ return { "kind": "LinkedField", "name": "dataPrivacyAgreement", "plural": false, - "selections": (v24/*: any*/), + "selections": (v23/*: any*/), "storageKey": null } ], @@ -802,24 +771,24 @@ return { "name": "viewer", "plural": false, "selections": [ - (v8/*: any*/), - (v3/*: any*/) + (v7/*: any*/), + (v2/*: any*/) ], "storageKey": null } ] }, "params": { - "cacheID": "faee4f059ef9baab46d5f2420e36bb60", + "cacheID": "b4e4f090951f025be6f359df3d2dc41a", "id": null, "metadata": {}, "name": "VendorGraphNodeQuery", "operationKind": "query", - "text": "query VendorGraphNodeQuery(\n $vendorId: ID!\n $organizationId: ID!\n) {\n node(id: $vendorId) {\n __typename\n ... on Vendor {\n id\n snapshotId\n name\n websiteUrl\n ...useVendorFormFragment\n ...VendorComplianceTabFragment\n ...VendorContactsTabFragment\n ...VendorServicesTabFragment\n ...VendorRiskAssessmentTabFragment\n ...VendorOverviewTabBusinessAssociateAgreementFragment\n ...VendorOverviewTabDataPrivacyAgreementFragment\n }\n id\n }\n viewer {\n user {\n id\n people(organizationId: $organizationId) {\n id\n }\n }\n id\n }\n}\n\nfragment VendorComplianceTabFragment on Vendor {\n complianceReports(first: 50) {\n edges {\n node {\n id\n ...VendorComplianceTabFragment_report\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorComplianceTabFragment_report on VendorComplianceReport {\n id\n reportDate\n validUntil\n reportName\n file {\n fileName\n mimeType\n size\n id\n }\n}\n\nfragment VendorContactsTabFragment on Vendor {\n contacts(first: 50) {\n edges {\n node {\n id\n ...VendorContactsTabFragment_contact\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorContactsTabFragment_contact on VendorContact {\n id\n fullName\n email\n phone\n role\n createdAt\n updatedAt\n}\n\nfragment VendorOverviewTabBusinessAssociateAgreementFragment on Vendor {\n businessAssociateAgreement {\n id\n fileName\n fileUrl\n validFrom\n validUntil\n createdAt\n }\n}\n\nfragment VendorOverviewTabDataPrivacyAgreementFragment on Vendor {\n dataPrivacyAgreement {\n id\n fileName\n fileUrl\n validFrom\n validUntil\n createdAt\n }\n}\n\nfragment VendorRiskAssessmentTabFragment on Vendor {\n id\n riskAssessments(first: 50) {\n edges {\n node {\n id\n ...VendorRiskAssessmentTabFragment_assessment\n __typename\n }\n cursor\n }\n pageInfo {\n hasNextPage\n endCursor\n hasPreviousPage\n startCursor\n }\n }\n}\n\nfragment VendorRiskAssessmentTabFragment_assessment on VendorRiskAssessment {\n id\n createdAt\n expiresAt\n dataSensitivity\n businessImpact\n notes\n}\n\nfragment VendorServicesTabFragment on Vendor {\n services(first: 50) {\n edges {\n node {\n id\n ...VendorServicesTabFragment_service\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorServicesTabFragment_service on VendorService {\n id\n name\n description\n createdAt\n updatedAt\n}\n\nfragment useVendorFormFragment on Vendor {\n id\n name\n description\n category\n statusPageUrl\n termsOfServiceUrl\n privacyPolicyUrl\n serviceLevelAgreementUrl\n dataProcessingAgreementUrl\n websiteUrl\n legalName\n headquarterAddress\n certifications\n countries\n securityPageUrl\n trustPageUrl\n businessOwner {\n id\n }\n securityOwner {\n id\n }\n}\n" + "text": "query VendorGraphNodeQuery(\n $vendorId: ID!\n) {\n node(id: $vendorId) {\n __typename\n ... on Vendor {\n id\n snapshotId\n name\n websiteUrl\n ...useVendorFormFragment\n ...VendorComplianceTabFragment\n ...VendorContactsTabFragment\n ...VendorServicesTabFragment\n ...VendorRiskAssessmentTabFragment\n ...VendorOverviewTabBusinessAssociateAgreementFragment\n ...VendorOverviewTabDataPrivacyAgreementFragment\n }\n id\n }\n viewer {\n user {\n id\n }\n id\n }\n}\n\nfragment VendorComplianceTabFragment on Vendor {\n complianceReports(first: 50) {\n edges {\n node {\n id\n ...VendorComplianceTabFragment_report\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorComplianceTabFragment_report on VendorComplianceReport {\n id\n reportDate\n validUntil\n reportName\n file {\n fileName\n mimeType\n size\n id\n }\n}\n\nfragment VendorContactsTabFragment on Vendor {\n contacts(first: 50) {\n edges {\n node {\n id\n ...VendorContactsTabFragment_contact\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorContactsTabFragment_contact on VendorContact {\n id\n fullName\n email\n phone\n role\n createdAt\n updatedAt\n}\n\nfragment VendorOverviewTabBusinessAssociateAgreementFragment on Vendor {\n businessAssociateAgreement {\n id\n fileName\n fileUrl\n validFrom\n validUntil\n createdAt\n }\n}\n\nfragment VendorOverviewTabDataPrivacyAgreementFragment on Vendor {\n dataPrivacyAgreement {\n id\n fileName\n fileUrl\n validFrom\n validUntil\n createdAt\n }\n}\n\nfragment VendorRiskAssessmentTabFragment on Vendor {\n id\n riskAssessments(first: 50) {\n edges {\n node {\n id\n ...VendorRiskAssessmentTabFragment_assessment\n __typename\n }\n cursor\n }\n pageInfo {\n hasNextPage\n endCursor\n hasPreviousPage\n startCursor\n }\n }\n}\n\nfragment VendorRiskAssessmentTabFragment_assessment on VendorRiskAssessment {\n id\n createdAt\n expiresAt\n dataSensitivity\n businessImpact\n notes\n}\n\nfragment VendorServicesTabFragment on Vendor {\n services(first: 50) {\n edges {\n node {\n id\n ...VendorServicesTabFragment_service\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n hasPreviousPage\n startCursor\n }\n }\n id\n}\n\nfragment VendorServicesTabFragment_service on VendorService {\n id\n name\n description\n createdAt\n updatedAt\n}\n\nfragment useVendorFormFragment on Vendor {\n id\n name\n description\n category\n statusPageUrl\n termsOfServiceUrl\n privacyPolicyUrl\n serviceLevelAgreementUrl\n dataProcessingAgreementUrl\n websiteUrl\n legalName\n headquarterAddress\n certifications\n countries\n securityPageUrl\n trustPageUrl\n businessOwner {\n id\n }\n securityOwner {\n id\n }\n}\n" } }; })(); -(node as any).hash = "5b907e5cc3446745ff2414ee06839aed"; +(node as any).hash = "0a13dd059626880c4a51f8a47b9c2423"; export default node; diff --git a/apps/console/src/layouts/MainLayout.tsx b/apps/console/src/layouts/MainLayout.tsx index fc6af76ce..205e58466 100644 --- a/apps/console/src/layouts/MainLayout.tsx +++ b/apps/console/src/layouts/MainLayout.tsx @@ -32,6 +32,8 @@ import { IconPlusLarge, IconChevronDown, Avatar, + IconPeopleAdd, + Badge, } from "@probo/ui"; import { useTranslate } from "@probo/i18n"; import { graphql } from "relay-runtime"; @@ -85,6 +87,9 @@ const OrganizationSelectorFragment = graphql` endCursor } } + invitations(first: 1, filter: {onlyPending: true}) { + totalCount + } } `; @@ -287,6 +292,7 @@ function OrganizationSelector({ ); const organizations = data.organizations.edges.map((edge) => edge.node); + const pendingInvitationsCount = data.invitations.totalCount; const handleLoadMore = (e?: React.MouseEvent) => { e?.preventDefault(); @@ -301,54 +307,79 @@ function OrganizationSelector({ }; return ( - - {currentOrganization?.name || ""} - - } - > -
- {organizations.map((organization) => ( - + - - - {organization.name} + {currentOrganization?.name || ""} + + } + > +
+ {organizations.map((organization) => ( + + + + {organization.name} + + + ))} + {hasNext && ( +
+ +
+ )} +
+ + {pendingInvitationsCount > 0 && ( + + + + {__("Invitations")} + + {pendingInvitationsCount} + - ))} - {hasNext && ( -
- -
)} -
- - - - - {__("Add organization")} + + + + {__("Add organization")} + + +
+ {pendingInvitationsCount > 0 && ( + + + + + ); +} + type OrganizationCardProps = { organization: { id: string; diff --git a/apps/console/src/pages/__generated__/OrganizationsPageQuery.graphql.ts b/apps/console/src/pages/__generated__/OrganizationsPageQuery.graphql.ts index 69aaee473..5b3acc6f3 100644 --- a/apps/console/src/pages/__generated__/OrganizationsPageQuery.graphql.ts +++ b/apps/console/src/pages/__generated__/OrganizationsPageQuery.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<84d0e43eced78862d1c82603e9c332e8>> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -12,6 +12,24 @@ import { ConcreteRequest } from 'relay-runtime'; export type OrganizationsPageQuery$variables = Record; export type OrganizationsPageQuery$data = { readonly viewer: { + readonly invitations: { + readonly __id: string; + readonly edges: ReadonlyArray<{ + readonly node: { + readonly acceptedAt: any | null | undefined; + readonly createdAt: any; + readonly email: string; + readonly expiresAt: any; + readonly fullName: string; + readonly id: string; + readonly organization: { + readonly id: string; + readonly name: string; + }; + readonly role: string; + }; + }>; + }; readonly organizations: { readonly __id: string; readonly edges: ReadonlyArray<{ @@ -45,7 +63,65 @@ v1 = { "name": "id", "storageKey": null }, -v2 = [ +v2 = { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "name", + "storageKey": null +}, +v3 = { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "__typename", + "storageKey": null +}, +v4 = { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "cursor", + "storageKey": null +}, +v5 = { + "alias": null, + "args": null, + "concreteType": "PageInfo", + "kind": "LinkedField", + "name": "pageInfo", + "plural": false, + "selections": [ + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "endCursor", + "storageKey": null + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "hasNextPage", + "storageKey": null + } + ], + "storageKey": null +}, +v6 = { + "kind": "ClientExtension", + "selections": [ + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "__id", + "storageKey": null + } + ] +}, +v7 = [ { "alias": null, "args": null, @@ -63,13 +139,7 @@ v2 = [ "plural": false, "selections": [ (v1/*: any*/), - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "name", - "storageKey": null - }, + (v2/*: any*/), { "alias": null, "args": null, @@ -77,71 +147,129 @@ v2 = [ "name": "logoUrl", "storageKey": null }, + (v3/*: any*/) + ], + "storageKey": null + }, + (v4/*: any*/) + ], + "storageKey": null + }, + (v5/*: any*/), + (v6/*: any*/) +], +v8 = { + "kind": "Literal", + "name": "filter", + "value": { + "onlyPending": true + } +}, +v9 = { + "kind": "Literal", + "name": "orderBy", + "value": { + "direction": "DESC", + "field": "CREATED_AT" + } +}, +v10 = [ + { + "alias": null, + "args": null, + "concreteType": "InvitationEdge", + "kind": "LinkedField", + "name": "edges", + "plural": true, + "selections": [ + { + "alias": null, + "args": null, + "concreteType": "Invitation", + "kind": "LinkedField", + "name": "node", + "plural": false, + "selections": [ + (v1/*: any*/), { "alias": null, "args": null, "kind": "ScalarField", - "name": "__typename", + "name": "email", "storageKey": null - } + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "fullName", + "storageKey": null + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "role", + "storageKey": null + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "expiresAt", + "storageKey": null + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "acceptedAt", + "storageKey": null + }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "createdAt", + "storageKey": null + }, + { + "alias": null, + "args": null, + "concreteType": "Organization", + "kind": "LinkedField", + "name": "organization", + "plural": false, + "selections": [ + (v1/*: any*/), + (v2/*: any*/) + ], + "storageKey": null + }, + (v3/*: any*/) ], "storageKey": null }, - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "cursor", - "storageKey": null - } + (v4/*: any*/) ], "storageKey": null }, - { - "alias": null, - "args": null, - "concreteType": "PageInfo", - "kind": "LinkedField", - "name": "pageInfo", - "plural": false, - "selections": [ - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "endCursor", - "storageKey": null - }, - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "hasNextPage", - "storageKey": null - } - ], - "storageKey": null - }, - { - "kind": "ClientExtension", - "selections": [ - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "__id", - "storageKey": null - } - ] - } + (v5/*: any*/), + (v6/*: any*/) ], -v3 = [ - { - "kind": "Literal", - "name": "first", - "value": 1000 - }, +v11 = { + "kind": "Literal", + "name": "first", + "value": 1000 +}, +v12 = [ + (v11/*: any*/), (v0/*: any*/) +], +v13 = [ + (v8/*: any*/), + (v11/*: any*/), + (v9/*: any*/) ]; return { "fragment": { @@ -167,8 +295,21 @@ return { "kind": "LinkedField", "name": "__OrganizationsPage_organizations_connection", "plural": false, - "selections": (v2/*: any*/), + "selections": (v7/*: any*/), "storageKey": "__OrganizationsPage_organizations_connection(orderBy:{\"direction\":\"ASC\",\"field\":\"NAME\"})" + }, + { + "alias": "invitations", + "args": [ + (v8/*: any*/), + (v9/*: any*/) + ], + "concreteType": "InvitationConnection", + "kind": "LinkedField", + "name": "__OrganizationsPage_invitations_connection", + "plural": false, + "selections": (v10/*: any*/), + "storageKey": "__OrganizationsPage_invitations_connection(filter:{\"onlyPending\":true},orderBy:{\"direction\":\"DESC\",\"field\":\"CREATED_AT\"})" } ], "storageKey": null @@ -193,17 +334,17 @@ return { "selections": [ { "alias": null, - "args": (v3/*: any*/), + "args": (v12/*: any*/), "concreteType": "OrganizationConnection", "kind": "LinkedField", "name": "organizations", "plural": false, - "selections": (v2/*: any*/), + "selections": (v7/*: any*/), "storageKey": "organizations(first:1000,orderBy:{\"direction\":\"ASC\",\"field\":\"NAME\"})" }, { "alias": null, - "args": (v3/*: any*/), + "args": (v12/*: any*/), "filters": [ "orderBy" ], @@ -212,6 +353,28 @@ return { "kind": "LinkedHandle", "name": "organizations" }, + { + "alias": null, + "args": (v13/*: any*/), + "concreteType": "InvitationConnection", + "kind": "LinkedField", + "name": "invitations", + "plural": false, + "selections": (v10/*: any*/), + "storageKey": "invitations(filter:{\"onlyPending\":true},first:1000,orderBy:{\"direction\":\"DESC\",\"field\":\"CREATED_AT\"})" + }, + { + "alias": null, + "args": (v13/*: any*/), + "filters": [ + "orderBy", + "filter" + ], + "handle": "connection", + "key": "OrganizationsPage_invitations", + "kind": "LinkedHandle", + "name": "invitations" + }, (v1/*: any*/) ], "storageKey": null @@ -219,7 +382,7 @@ return { ] }, "params": { - "cacheID": "1735764e6816660969c5f96922320ac5", + "cacheID": "5675b3eb7810ef04bf531d7bf988c86c", "id": null, "metadata": { "connection": [ @@ -231,16 +394,25 @@ return { "viewer", "organizations" ] + }, + { + "count": null, + "cursor": null, + "direction": "forward", + "path": [ + "viewer", + "invitations" + ] } ] }, "name": "OrganizationsPageQuery", "operationKind": "query", - "text": "query OrganizationsPageQuery {\n viewer {\n organizations(first: 1000, orderBy: {field: NAME, direction: ASC}) {\n edges {\n node {\n id\n name\n logoUrl\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n }\n }\n id\n }\n}\n" + "text": "query OrganizationsPageQuery {\n viewer {\n organizations(first: 1000, orderBy: {field: NAME, direction: ASC}) {\n edges {\n node {\n id\n name\n logoUrl\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n }\n }\n invitations(first: 1000, orderBy: {field: CREATED_AT, direction: DESC}, filter: {onlyPending: true}) {\n edges {\n node {\n id\n email\n fullName\n role\n expiresAt\n acceptedAt\n createdAt\n organization {\n id\n name\n }\n __typename\n }\n cursor\n }\n pageInfo {\n endCursor\n hasNextPage\n }\n }\n id\n }\n}\n" } }; })(); -(node as any).hash = "6fde39384e4678f88f17c34bbd30e684"; +(node as any).hash = "a548f9c3a434e12c079fb2ff651822d2"; export default node; diff --git a/apps/console/src/pages/__generated__/OrganizationsPage_AcceptInvitationMutation.graphql.ts b/apps/console/src/pages/__generated__/OrganizationsPage_AcceptInvitationMutation.graphql.ts new file mode 100644 index 000000000..60158b26c --- /dev/null +++ b/apps/console/src/pages/__generated__/OrganizationsPage_AcceptInvitationMutation.graphql.ts @@ -0,0 +1,105 @@ +/** + * @generated SignedSource<> + * @lightSyntaxTransform + * @nogrep + */ + +/* tslint:disable */ +/* eslint-disable */ +// @ts-nocheck + +import { ConcreteRequest } from 'relay-runtime'; +export type AcceptInvitationInput = { + invitationId: string; +}; +export type OrganizationsPage_AcceptInvitationMutation$variables = { + input: AcceptInvitationInput; +}; +export type OrganizationsPage_AcceptInvitationMutation$data = { + readonly acceptInvitation: { + readonly invitation: { + readonly id: string; + }; + }; +}; +export type OrganizationsPage_AcceptInvitationMutation = { + response: OrganizationsPage_AcceptInvitationMutation$data; + variables: OrganizationsPage_AcceptInvitationMutation$variables; +}; + +const node: ConcreteRequest = (function(){ +var v0 = [ + { + "defaultValue": null, + "kind": "LocalArgument", + "name": "input" + } +], +v1 = [ + { + "alias": null, + "args": [ + { + "kind": "Variable", + "name": "input", + "variableName": "input" + } + ], + "concreteType": "AcceptInvitationPayload", + "kind": "LinkedField", + "name": "acceptInvitation", + "plural": false, + "selections": [ + { + "alias": null, + "args": null, + "concreteType": "Invitation", + "kind": "LinkedField", + "name": "invitation", + "plural": false, + "selections": [ + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "id", + "storageKey": null + } + ], + "storageKey": null + } + ], + "storageKey": null + } +]; +return { + "fragment": { + "argumentDefinitions": (v0/*: any*/), + "kind": "Fragment", + "metadata": null, + "name": "OrganizationsPage_AcceptInvitationMutation", + "selections": (v1/*: any*/), + "type": "Mutation", + "abstractKey": null + }, + "kind": "Request", + "operation": { + "argumentDefinitions": (v0/*: any*/), + "kind": "Operation", + "name": "OrganizationsPage_AcceptInvitationMutation", + "selections": (v1/*: any*/) + }, + "params": { + "cacheID": "cc4a442037edaa624948b5be0c009823", + "id": null, + "metadata": {}, + "name": "OrganizationsPage_AcceptInvitationMutation", + "operationKind": "mutation", + "text": "mutation OrganizationsPage_AcceptInvitationMutation(\n $input: AcceptInvitationInput!\n) {\n acceptInvitation(input: $input) {\n invitation {\n id\n }\n }\n}\n" + } +}; +})(); + +(node as any).hash = "190213ab6fdc068343a270b4fa94e160"; + +export default node; diff --git a/apps/console/src/pages/auth/ConfirmInvitationPage.tsx b/apps/console/src/pages/auth/SignupFromInvitationPage.tsx similarity index 62% rename from apps/console/src/pages/auth/ConfirmInvitationPage.tsx rename to apps/console/src/pages/auth/SignupFromInvitationPage.tsx index e71ea451a..d03eeeca7 100644 --- a/apps/console/src/pages/auth/ConfirmInvitationPage.tsx +++ b/apps/console/src/pages/auth/SignupFromInvitationPage.tsx @@ -1,35 +1,49 @@ -import { Link, useNavigate } from "react-router"; +import { Link, useNavigate, useSearchParams } from "react-router"; import { Button, Field, useToast } from "@probo/ui"; import { useTranslate } from "@probo/i18n"; import { z } from "zod"; import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { usePageTitle } from "@probo/hooks"; import { buildEndpoint } from "/providers/RelayProviders"; +import { useEffect } from "react"; const schema = z.object({ + fullName: z.string().min(2), password: z.string().min(8), }); -export default function ConfirmInvitationPage() { +export default function SignupFromInvitationPage() { const { __ } = useTranslate(); const navigate = useNavigate(); const { toast } = useToast(); - const { register, handleSubmit, formState } = useFormWithSchema( + const [searchParams] = useSearchParams(); + + const { register, handleSubmit, formState, reset } = useFormWithSchema( schema, { defaultValues: { + fullName: "", password: "", }, } ); + useEffect(() => { + const fullNameFromParams = searchParams.get("fullName") || ""; + if (fullNameFromParams) { + reset({ + fullName: fullNameFromParams, + password: "", + }); + } + }, [searchParams, reset]); + const onSubmit = handleSubmit(async (data) => { - const searchParams = new URLSearchParams(location.search); const token = searchParams.get("token"); if (!token) { toast({ - title: __("Confirmation failed"), + title: __("Signup failed"), description: __("Invalid or missing invitation token"), variant: "error", }); @@ -37,7 +51,7 @@ export default function ConfirmInvitationPage() { } const response = await fetch( - buildEndpoint("/api/console/v1/auth/invitation"), + buildEndpoint("/api/console/v1/auth/signup-from-invitation"), { method: "POST", headers: { @@ -47,16 +61,16 @@ export default function ConfirmInvitationPage() { body: JSON.stringify({ token: token, password: data.password, + fullName: data.fullName, }), } ); - // Registration failed if (!response.ok) { const errorData = await response.json().catch(() => ({})); toast({ - title: __("Confirmation failed"), - description: errorData.message || __("Confirmation failed"), + title: __("Signup failed"), + description: errorData.message || __("Signup failed"), variant: "error", }); return; @@ -64,24 +78,33 @@ export default function ConfirmInvitationPage() { toast({ title: __("Success"), - description: __("Invitation confirmed successfully"), + description: __("Account created successfully. Please accept your invitation to join the organization."), variant: "success", }); navigate("/", { replace: true }); }); - usePageTitle(__("Confirm invitation")); + usePageTitle(__("Create your account")); return (
-

{__("Confirm invitation")}

+

{__("Create your account")}

- {__("Enter your information to confirm your invitation")} + {__("Set your password to join the organization")}

+ + {formState.isLoading - ? __("Confirming invitation...") - : __("Confirm invitation")} + ? __("Creating account...") + : __("Create account")} diff --git a/apps/console/src/pages/organizations/SettingsPage.tsx b/apps/console/src/pages/organizations/SettingsPage.tsx index cfe518c6c..c71b31f03 100644 --- a/apps/console/src/pages/organizations/SettingsPage.tsx +++ b/apps/console/src/pages/organizations/SettingsPage.tsx @@ -50,13 +50,13 @@ import { sprintf } from "@probo/helpers"; import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { z } from "zod"; import type { NodeOf } from "/types"; -import { useMutationWithToasts } from "/hooks/useMutationWithToasts"; import { useOrganizationId } from "/hooks/useOrganizationId"; import { InviteUserDialog } from "/components/organizations/InviteUserDialog"; import { useDeleteOrganizationMutation } from "/hooks/graph/OrganizationGraph"; import { useNavigate } from "react-router"; import { DeleteOrganizationDialog } from "/components/organizations/DeleteOrganizationDialog"; import { CustomDomainManager } from "/components/customDomains/CustomDomainManager"; +import { useMutationWithToasts } from "/hooks/useMutationWithToasts"; const organizationSchema = z.object({ name: z.string().min(1, "Organization name is required"), @@ -220,6 +220,7 @@ const deleteHorizontalLogoMutation = graphql` export default function SettingsPage({ queryRef }: Props) { const { __ } = useTranslate(); const navigate = useNavigate(); + const organizationId = useOrganizationId(); const organizationKey = usePreloadedQuery( organizationViewQuery, queryRef @@ -240,6 +241,14 @@ export default function SettingsPage({ queryRef }: Props) { organizationKey as SettingsPageInvitationsFragment$key ); + const refetchMemberships = () => { + membershipsPagination.refetch({}, { fetchPolicy: 'network-only' }); + }; + + const refetchInvitations = () => { + invitationsPagination.refetch({}, { fetchPolicy: 'network-only' }); + }; + const [updateOrganization] = useMutation(updateOrganizationMutation); const [deleteHorizontalLogo, isDeletingHorizontalLogo] = useMutationWithToasts( deleteHorizontalLogoMutation, @@ -253,24 +262,6 @@ export default function SettingsPage({ queryRef }: Props) { const invitations = invitationsPagination.data.invitations?.edges.map((edge) => edge.node) || []; const [activeTab, setActiveTab] = useState<"memberships" | "invitations">("memberships"); - const refetchMemberships = ({ order }: { order: { direction: string; field: string } }) => { - membershipsPagination.refetch({ - order: { - direction: order.direction as "ASC" | "DESC", - field: order.field as "CREATED_AT" | "FULL_NAME" | "EMAIL_ADDRESS" | "ROLE" - } - }); - }; - - const refetchInvitations = ({ order }: { order: { direction: string; field: string } }) => { - invitationsPagination.refetch({ - order: { - direction: order.direction as "ASC" | "DESC", - field: order.field as "CREATED_AT" | "EXPIRES_AT" | "FULL_NAME" | "EMAIL" | "ROLE" | "STATUS" | "ACCEPTED_AT" - } - }); - }; - const { formState, handleSubmit, register, reset } = useFormWithSchema( organizationSchema, { @@ -572,7 +563,10 @@ export default function SettingsPage({ queryRef }: Props) {

{__("Workspace members")}

- +
@@ -603,7 +597,14 @@ export default function SettingsPage({ queryRef }: Props) { {activeTab === "memberships" && ( { + membershipsPagination.refetch({ + order: { + direction: order.direction as "ASC" | "DESC", + field: order.field as "CREATED_AT" | "FULL_NAME" | "EMAIL_ADDRESS" | "ROLE" + } + }); + }} > @@ -623,7 +624,13 @@ export default function SettingsPage({ queryRef }: Props) { ) : ( memberships.map((membership) => ( - + )) )} @@ -633,7 +640,14 @@ export default function SettingsPage({ queryRef }: Props) { {activeTab === "invitations" && ( { + invitationsPagination.refetch({ + order: { + direction: order.direction as "ASC" | "DESC", + field: order.field as "CREATED_AT" | "EXPIRES_AT" | "FULL_NAME" | "EMAIL" | "ROLE" | "STATUS" | "ACCEPTED_AT" + } + }); + }} > @@ -659,6 +673,8 @@ export default function SettingsPage({ queryRef }: Props) { key={invitation.id} invitation={invitation} connectionId={invitationsPagination.data.invitations?.__id} + organizationId={organizationId} + onRefetch={refetchInvitations} /> )) )} @@ -794,9 +810,12 @@ function Connectors(props: { } const removeMemberMutation = graphql` - mutation SettingsPage_RemoveMemberMutation($input: RemoveMemberInput!) { + mutation SettingsPage_RemoveMemberMutation( + $input: RemoveMemberInput! + $connections: [ID!]! + ) { removeMember(input: $input) { - success + deletedMemberId @deleteEdge(connections: $connections) } } `; @@ -804,16 +823,16 @@ const removeMemberMutation = graphql` function InvitationRow(props: { invitation: NodeOf; connectionId?: string; + organizationId: string; + onRefetch: () => void; }) { const { __ } = useTranslate(); const confirm = useConfirm(); const [deleteInvitation, isDeleting] = useMutationWithToasts( deleteInvitationMutation, { - successMessage: sprintf( - __("Invitation for %s deleted successfully"), - props.invitation.fullName - ), + successMessage: __("Invitation deleted successfully"), + errorMessage: __("Failed to delete invitation"), } ); @@ -830,6 +849,9 @@ function InvitationRow(props: { }, connections: props.connectionId ? [props.connectionId] : [], }, + onCompleted: () => { + props.onRefetch(); + }, }); }, { @@ -885,15 +907,16 @@ function InvitationRow(props: { ); } -function MembershipRow(props: { membership: NodeOf }) { +function MembershipRow(props: { + membership: NodeOf; + connectionId?: string; + organizationId: string; + onRefetch: () => void; +}) { const { __ } = useTranslate(); - const organizationId = useOrganizationId(); const [removeMember, isRemoving] = useMutationWithToasts(removeMemberMutation, { - successMessage: sprintf( - __("Member %s removed successfully"), - props.membership.fullName - ), - errorMessage: sprintf(__("Failed to remove member %s"), props.membership.fullName), + successMessage: __("Member removed successfully"), + errorMessage: __("Failed to remove member"), }); const confirm = useConfirm(); const [isRemoved, setIsRemoved] = useState(false); @@ -909,11 +932,13 @@ function MembershipRow(props: { membership: NodeOf { + onCompleted: () => { setIsRemoved(true); + props.onRefetch(); }, }); }, diff --git a/apps/console/src/pages/organizations/__generated__/SettingsPageInvitationsFragment.graphql.ts b/apps/console/src/pages/organizations/__generated__/SettingsPageInvitationsFragment.graphql.ts index da2a45b73..f4c8a7daa 100644 --- a/apps/console/src/pages/organizations/__generated__/SettingsPageInvitationsFragment.graphql.ts +++ b/apps/console/src/pages/organizations/__generated__/SettingsPageInvitationsFragment.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<4f099d1ee6b4635ca8129eca77d3f7e8>> + * @generated SignedSource<<09561c2c29459fc840773b414af6083a>> * @lightSyntaxTransform * @nogrep */ @@ -9,6 +9,7 @@ // @ts-nocheck import { ReaderFragment } from 'relay-runtime'; +export type InvitationStatus = "ACCEPTED" | "EXPIRED" | "PENDING"; import { FragmentRefs } from "relay-runtime"; export type SettingsPageInvitationsFragment$data = { readonly id: string; @@ -23,6 +24,7 @@ export type SettingsPageInvitationsFragment$data = { readonly fullName: string; readonly id: string; readonly role: string; + readonly status: InvitationStatus; }; }>; readonly totalCount: number; @@ -171,6 +173,13 @@ return { "name": "role", "storageKey": null }, + { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "status", + "storageKey": null + }, { "alias": null, "args": null, @@ -273,6 +282,6 @@ return { }; })(); -(node as any).hash = "d971d93653991efde284ed2b87068698"; +(node as any).hash = "f9a1ec38579cea21312ba0a20bb7394a"; export default node; diff --git a/apps/console/src/pages/organizations/__generated__/SettingsPage_RemoveMemberMutation.graphql.ts b/apps/console/src/pages/organizations/__generated__/SettingsPage_RemoveMemberMutation.graphql.ts index 326c0fe77..cf363eea2 100644 --- a/apps/console/src/pages/organizations/__generated__/SettingsPage_RemoveMemberMutation.graphql.ts +++ b/apps/console/src/pages/organizations/__generated__/SettingsPage_RemoveMemberMutation.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<994eb713978ecef3329fd9873e4c744c>> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -14,11 +14,12 @@ export type RemoveMemberInput = { organizationId: string; }; export type SettingsPage_RemoveMemberMutation$variables = { + connections: ReadonlyArray; input: RemoveMemberInput; }; export type SettingsPage_RemoveMemberMutation$data = { readonly removeMember: { - readonly success: boolean; + readonly deletedMemberId: string; }; }; export type SettingsPage_RemoveMemberMutation = { @@ -27,67 +28,106 @@ export type SettingsPage_RemoveMemberMutation = { }; const node: ConcreteRequest = (function(){ -var v0 = [ +var v0 = { + "defaultValue": null, + "kind": "LocalArgument", + "name": "connections" +}, +v1 = { + "defaultValue": null, + "kind": "LocalArgument", + "name": "input" +}, +v2 = [ { - "defaultValue": null, - "kind": "LocalArgument", - "name": "input" + "kind": "Variable", + "name": "input", + "variableName": "input" } ], -v1 = [ - { - "alias": null, - "args": [ - { - "kind": "Variable", - "name": "input", - "variableName": "input" - } - ], - "concreteType": "RemoveMemberPayload", - "kind": "LinkedField", - "name": "removeMember", - "plural": false, - "selections": [ - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "success", - "storageKey": null - } - ], - "storageKey": null - } -]; +v3 = { + "alias": null, + "args": null, + "kind": "ScalarField", + "name": "deletedMemberId", + "storageKey": null +}; return { "fragment": { - "argumentDefinitions": (v0/*: any*/), + "argumentDefinitions": [ + (v0/*: any*/), + (v1/*: any*/) + ], "kind": "Fragment", "metadata": null, "name": "SettingsPage_RemoveMemberMutation", - "selections": (v1/*: any*/), + "selections": [ + { + "alias": null, + "args": (v2/*: any*/), + "concreteType": "RemoveMemberPayload", + "kind": "LinkedField", + "name": "removeMember", + "plural": false, + "selections": [ + (v3/*: any*/) + ], + "storageKey": null + } + ], "type": "Mutation", "abstractKey": null }, "kind": "Request", "operation": { - "argumentDefinitions": (v0/*: any*/), + "argumentDefinitions": [ + (v1/*: any*/), + (v0/*: any*/) + ], "kind": "Operation", "name": "SettingsPage_RemoveMemberMutation", - "selections": (v1/*: any*/) + "selections": [ + { + "alias": null, + "args": (v2/*: any*/), + "concreteType": "RemoveMemberPayload", + "kind": "LinkedField", + "name": "removeMember", + "plural": false, + "selections": [ + (v3/*: any*/), + { + "alias": null, + "args": null, + "filters": null, + "handle": "deleteEdge", + "key": "", + "kind": "ScalarHandle", + "name": "deletedMemberId", + "handleArgs": [ + { + "kind": "Variable", + "name": "connections", + "variableName": "connections" + } + ] + } + ], + "storageKey": null + } + ] }, "params": { - "cacheID": "97e29046871ce8aab01abf98a62236fc", + "cacheID": "e2dd0f4d7327ce3bc97754c85d3f700d", "id": null, "metadata": {}, "name": "SettingsPage_RemoveMemberMutation", "operationKind": "mutation", - "text": "mutation SettingsPage_RemoveMemberMutation(\n $input: RemoveMemberInput!\n) {\n removeMember(input: $input) {\n success\n }\n}\n" + "text": "mutation SettingsPage_RemoveMemberMutation(\n $input: RemoveMemberInput!\n) {\n removeMember(input: $input) {\n deletedMemberId\n }\n}\n" } }; })(); -(node as any).hash = "f61071a0fb6f6554e56e79b6b04bc135"; +(node as any).hash = "9909a8b95f8d8621ffdf02da34ec8da2"; export default node; diff --git a/apps/console/src/routes.tsx b/apps/console/src/routes.tsx index 2e83a6704..1a37806ee 100644 --- a/apps/console/src/routes.tsx +++ b/apps/console/src/routes.tsx @@ -74,8 +74,8 @@ const routes = [ Component: lazy(() => import("./pages/auth/ConfirmEmailPage")), }, { - path: "confirm-invitation", - Component: lazy(() => import("./pages/auth/ConfirmInvitationPage")), + path: "signup-from-invitation", + Component: lazy(() => import("./pages/auth/SignupFromInvitationPage")), }, { path: "forgot-password", diff --git a/apps/console/src/routes/vendorRoutes.ts b/apps/console/src/routes/vendorRoutes.ts index 29a653c2b..e2b209beb 100644 --- a/apps/console/src/routes/vendorRoutes.ts +++ b/apps/console/src/routes/vendorRoutes.ts @@ -30,10 +30,9 @@ export const vendorRoutes = [ { path: "vendors/:vendorId", fallback: PageSkeleton, - queryLoader: ({ vendorId, organizationId }) => + queryLoader: ({ vendorId }) => loadQuery(relayEnvironment, vendorNodeQuery, { vendorId, - organizationId, }), Component: lazy( () => import("../pages/organizations/vendors/VendorDetailPage") @@ -95,10 +94,9 @@ export const vendorRoutes = [ { path: "snapshots/:snapshotId/vendors/:vendorId", fallback: PageSkeleton, - queryLoader: ({ vendorId, organizationId }) => + queryLoader: ({ vendorId }) => loadQuery(relayEnvironment, vendorNodeQuery, { vendorId, - organizationId, }), Component: lazy( () => import("../pages/organizations/vendors/VendorDetailPage") diff --git a/packages/ui/src/Layouts/AuthLayout.tsx b/packages/ui/src/Layouts/AuthLayout.tsx index 2f95f431d..f0fb18382 100644 --- a/packages/ui/src/Layouts/AuthLayout.tsx +++ b/packages/ui/src/Layouts/AuthLayout.tsx @@ -1,9 +1,7 @@ -import logo from "../assets/android-chrome-512x512.png"; -import { useTranslate } from "@probo/i18n"; import { Outlet } from "react-router"; +import { Logo } from "../Atoms/Logo/Logo"; export function AuthLayout() { - const { __ } = useTranslate(); return (
@@ -11,13 +9,9 @@ export function AuthLayout() {
-
-
- Probo logo - - {__("Navigate compliance with confidence thanks to")} - probo - +
+
+
diff --git a/pkg/auth/service.go b/pkg/auth/service.go index 8f76c1a46..adc585b27 100644 --- a/pkg/auth/service.go +++ b/pkg/auth/service.go @@ -600,3 +600,107 @@ func (s Service) ResetPassword(ctx context.Context, tokenString string, newPassw }, ) } + +func (s Service) SignupFromInvitation( + ctx context.Context, + token string, + password string, + fullName string, +) (*coredata.User, *coredata.Session, error) { + payload, err := statelesstoken.ValidateToken[coredata.InvitationData]( + s.tokenSecret, + "organization_invitation", + token, + ) + if err != nil { + return nil, nil, &ErrInvalidTokenType{"invalid invitation token"} + } + invitationData := payload.Data + + if len(password) < 8 || len(password) > 128 { + return nil, nil, &ErrInvalidPassword{minLength: 8, maxLength: 128} + } + + if _, err := mail.ParseAddress(invitationData.Email); err != nil { + return nil, nil, &ErrInvalidEmail{invitationData.Email} + } + + if fullName == "" { + fullName = invitationData.FullName + } + + if fullName == "" { + return nil, nil, &ErrInvalidFullName{fullName} + } + + hashedPassword, err := s.hp.HashPassword([]byte(password)) + if err != nil { + return nil, nil, fmt.Errorf("cannot hash password: %w", err) + } + + var user *coredata.User + var session *coredata.Session + + scope := coredata.NewScope(invitationData.InvitationID.TenantID()) + err = s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + invitation := &coredata.Invitation{} + if err := invitation.LoadByID(ctx, tx, scope, invitationData.InvitationID); err != nil { + var errInvitationNotFound *coredata.ErrInvitationNotFound + if errors.As(err, &errInvitationNotFound) { + return fmt.Errorf("invitation was deleted or no longer exists") + } + return fmt.Errorf("cannot load invitation: %w", err) + } + + if invitation.AcceptedAt != nil { + return fmt.Errorf("invitation already accepted") + } + + if time.Now().After(invitation.ExpiresAt) { + return fmt.Errorf("invitation expired") + } + + now := time.Now() + user = &coredata.User{ + ID: gid.New(gid.NilTenant, coredata.UserEntityType), + EmailAddress: invitationData.Email, + HashedPassword: hashedPassword, + EmailAddressVerified: true, + FullName: fullName, + CreatedAt: now, + UpdatedAt: now, + } + + if err := user.Insert(ctx, tx); err != nil { + var errUserAlreadyExists *coredata.ErrUserAlreadyExists + if errors.As(err, &errUserAlreadyExists) { + return &ErrUserAlreadyExists{errUserAlreadyExists.Error()} + } + return fmt.Errorf("cannot insert user: %w", err) + } + + session = &coredata.Session{ + ID: gid.New(gid.NilTenant, coredata.SessionEntityType), + UserID: user.ID, + Data: coredata.SessionData{}, + ExpiredAt: now.Add(24 * time.Hour * 7), + CreatedAt: now, + UpdatedAt: now, + } + + if err := session.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot insert session: %w", err) + } + + return nil + }, + ) + + if err != nil { + return nil, nil, err + } + + return user, session, nil +} diff --git a/pkg/authz/service.go b/pkg/authz/service.go index 153d23a43..d057f9269 100644 --- a/pkg/authz/service.go +++ b/pkg/authz/service.go @@ -20,7 +20,8 @@ import ( _ "embed" "errors" "fmt" - "html/template" + "net/url" + "text/template" "time" "github.com/getprobo/probo/pkg/coredata" @@ -41,6 +42,14 @@ type ( invitationTokenValidity time.Duration } + TenantAuthzService struct { + pg *pg.Client + hostname string + tokenSecret string + invitationTokenValidity time.Duration + scope coredata.Scoper + } + Role string ) @@ -78,6 +87,18 @@ func NewService( }, nil } +func (s *Service) WithTenant(tenantID gid.TenantID) *TenantAuthzService { + return &TenantAuthzService{ + pg: s.pg, + hostname: s.hostname, + tokenSecret: s.tokenSecret, + invitationTokenValidity: s.invitationTokenValidity, + scope: coredata.NewScope(tenantID), + } +} + +// This method is on Service (not TenantAuthzService) because it operates across tenants +// and doesn't require tenant-scoped access. func (s *Service) GetAllUserOrganizations( ctx context.Context, userID gid.GID, @@ -98,6 +119,8 @@ func (s *Service) GetAllUserOrganizations( return organizations, err } +// This method is on Service (not TenantAuthzService) because it operates across tenants +// and doesn't require tenant-scoped access. func (s *Service) GetUserOrganizations( ctx context.Context, userID gid.GID, @@ -115,7 +138,253 @@ func (s *Service) GetUserOrganizations( return organizations, err } -func (s *Service) GetAllOrganizationInvitations( +// This method is on Service (not TenantAuthzService) because the user accepting +// the invitation doesn't have tenant access yet. +func (s *Service) AcceptInvitation( + ctx context.Context, + token string, + userID gid.GID, +) error { + payload, err := statelesstoken.ValidateToken[coredata.InvitationData]( + s.tokenSecret, + TokenTypeOrganizationInvitation, + token, + ) + if err != nil { + return fmt.Errorf("invalid invitation token: %w", err) + } + invitationData := payload.Data + scope := coredata.NewScope(invitationData.InvitationID.TenantID()) + + return s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + invitation := &coredata.Invitation{} + if err := invitation.LoadByID(ctx, tx, scope, invitationData.InvitationID); err != nil { + var errInvitationNotFound *coredata.ErrInvitationNotFound + if errors.As(err, &errInvitationNotFound) { + return fmt.Errorf("invitation was deleted or no longer exists") + } + return fmt.Errorf("cannot load invitation: %w", err) + } + + if invitation.AcceptedAt != nil { + return fmt.Errorf("invitation already accepted") + } + + if time.Now().After(invitation.ExpiresAt) { + return fmt.Errorf("invitation expired") + } + + now := time.Now() + membershipID := gid.New(scope.GetTenantID(), coredata.MembershipEntityType) + + membership := &coredata.Membership{ + ID: membershipID, + UserID: userID, + OrganizationID: invitation.OrganizationID, + Role: invitation.Role, + CreatedAt: now, + UpdatedAt: now, + } + + if err := membership.Create(ctx, tx, scope); err != nil { + return fmt.Errorf("failed to add user to organization: %w", err) + } + + invitation.AcceptedAt = &now + if err := invitation.Update(ctx, tx, scope); err != nil { + return fmt.Errorf("failed to mark invitation as accepted: %w", err) + } + + return nil + }, + ) +} + +// This method is on Service (not TenantAuthzService) because the user accepting +// the invitation doesn't have tenant access yet. +func (s *Service) AcceptInvitationByID( + ctx context.Context, + invitationID gid.GID, + userID gid.GID, +) (*coredata.Invitation, error) { + var acceptedInvitation *coredata.Invitation + scope := coredata.NewScope(invitationID.TenantID()) + err := s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + invitation := &coredata.Invitation{} + if err := invitation.LoadByID(ctx, tx, scope, invitationID); err != nil { + var errInvitationNotFound *coredata.ErrInvitationNotFound + if errors.As(err, &errInvitationNotFound) { + return fmt.Errorf("invitation was deleted or no longer exists") + } + return fmt.Errorf("cannot load invitation: %w", err) + } + + if invitation.AcceptedAt != nil { + return fmt.Errorf("invitation already accepted") + } + + if time.Now().After(invitation.ExpiresAt) { + return fmt.Errorf("invitation expired") + } + + user := &coredata.User{} + if err := user.LoadByID(ctx, tx, userID); err != nil { + return fmt.Errorf("cannot load user: %w", err) + } + + if invitation.Email != user.EmailAddress { + return fmt.Errorf("invitation email does not match user email") + } + + now := time.Now() + membershipID := gid.New(scope.GetTenantID(), coredata.MembershipEntityType) + + membership := &coredata.Membership{ + ID: membershipID, + UserID: userID, + OrganizationID: invitation.OrganizationID, + Role: invitation.Role, + CreatedAt: now, + UpdatedAt: now, + } + + if err := membership.Create(ctx, tx, scope); err != nil { + return fmt.Errorf("failed to add user to organization: %w", err) + } + + invitation.AcceptedAt = &now + if err := invitation.Update(ctx, tx, scope); err != nil { + return fmt.Errorf("failed to mark invitation as accepted: %w", err) + } + + acceptedInvitation = invitation + return nil + }, + ) + if err != nil { + return nil, err + } + return acceptedInvitation, nil +} + +// This method is on Service (not TenantAuthzService) because the user viewing +// their invitations doesn't have tenant access yet, and it operates across multiple tenants. +func (s *Service) GetUserInvitations( + ctx context.Context, + email string, + cursor *page.Cursor[coredata.InvitationOrderField], + filter *coredata.InvitationFilter, +) (*page.Page[*coredata.Invitation, coredata.InvitationOrderField], error) { + var invitations coredata.Invitations + + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := invitations.LoadByEmail(ctx, conn, email, cursor, filter); err != nil { + return fmt.Errorf("failed to load invitations: %w", err) + } + + return nil + }, + ) + if err != nil { + return nil, err + } + + return page.NewPage(invitations, cursor), nil +} + +// This method is on Service (not TenantAuthzService) because the user viewing +// their invitations doesn't have tenant access yet, and it operates across multiple tenants. +func (s *Service) CountUserInvitations( + ctx context.Context, + email string, + filter *coredata.InvitationFilter, +) (int, error) { + var count int + + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + var invitations coredata.Invitations + var err error + count, err = invitations.CountByEmail(ctx, conn, email, filter) + return err + }, + ) + + return count, err +} + +// This method is on Service (not TenantAuthzService) because the user viewing +// the invitation organization doesn't have tenant access yet. +func (s *Service) GetOrganizationByInvitationID( + ctx context.Context, + invitationID gid.GID, +) (*coredata.Organization, error) { + scope := coredata.NewScope(invitationID.TenantID()) + + var organization coredata.Organization + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + var invitation coredata.Invitation + if err := invitation.LoadByID(ctx, conn, scope, invitationID); err != nil { + return fmt.Errorf("failed to load invitation: %w", err) + } + + if err := organization.LoadByID(ctx, conn, scope, invitation.OrganizationID); err != nil { + return fmt.Errorf("failed to load organization: %w", err) + } + + return nil + }, + ) + if err != nil { + return nil, err + } + + return &organization, nil +} + +// This method is on Service (not TenantAuthzService) because the user added to the organization +// doesn't have tenant access yet +func (s *Service) AddUserToOrganization( + ctx context.Context, + userID gid.GID, + orgID gid.GID, + role string, +) error { + now := time.Now() + tenantID := orgID.TenantID() + membershipID := gid.New(tenantID, coredata.MembershipEntityType) + scope := coredata.NewScope(tenantID) + + membership := &coredata.Membership{ + ID: membershipID, + UserID: userID, + OrganizationID: orgID, + Role: role, + CreatedAt: now, + UpdatedAt: now, + } + + return s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := membership.Create(ctx, conn, scope); err != nil { + return fmt.Errorf("failed to add user to organization: %w", err) + } + return nil + }, + ) +} + +func (s *TenantAuthzService) GetInvitationsByOrganizationID( ctx context.Context, orgID gid.GID, cursor *page.Cursor[coredata.InvitationOrderField], @@ -125,7 +394,7 @@ func (s *Service) GetAllOrganizationInvitations( err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := invitations.LoadByOrganizationID(ctx, conn, orgID, cursor); err != nil { + if err := invitations.LoadByOrganizationID(ctx, conn, s.scope, orgID, cursor); err != nil { return fmt.Errorf("failed to load organization invitations: %w", err) } @@ -139,7 +408,7 @@ func (s *Service) GetAllOrganizationInvitations( return page.NewPage(invitations, cursor), nil } -func (s *Service) CountOrganizationInvitations( +func (s *TenantAuthzService) CountOrganizationInvitations( ctx context.Context, orgID gid.GID, ) (int, error) { @@ -149,7 +418,7 @@ func (s *Service) CountOrganizationInvitations( func(conn pg.Conn) error { var invitations coredata.Invitations var err error - count, err = invitations.CountByOrganizationID(ctx, conn, orgID) + count, err = invitations.CountByOrganizationID(ctx, conn, s.scope, orgID) return err }, ) @@ -160,7 +429,27 @@ func (s *Service) CountOrganizationInvitations( return count, nil } -func (s *Service) DeleteInvitation( +func (s *TenantAuthzService) GetInvitationByID( + ctx context.Context, + invitationID gid.GID, +) (*coredata.Invitation, error) { + invitation := &coredata.Invitation{} + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil { + return fmt.Errorf("failed to load invitation: %w", err) + } + return nil + }, + ) + if err != nil { + return nil, err + } + return invitation, nil +} + +func (s *TenantAuthzService) DeleteInvitation( ctx context.Context, invitationID gid.GID, ) error { @@ -168,11 +457,11 @@ func (s *Service) DeleteInvitation( ctx, func(conn pg.Conn) error { invitation := &coredata.Invitation{} - if err := invitation.LoadByID(ctx, conn, invitationID); err != nil { + if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil { return fmt.Errorf("failed to load invitation: %w", err) } - if err := invitation.Delete(ctx, conn); err != nil { + if err := invitation.Delete(ctx, conn, s.scope); err != nil { return fmt.Errorf("failed to delete invitation: %w", err) } @@ -181,7 +470,7 @@ func (s *Service) DeleteInvitation( ) } -func (s *Service) GetAllOrganizationMemberships( +func (s *TenantAuthzService) GetMembershipsByOrganizationID( ctx context.Context, orgID gid.GID, cursor *page.Cursor[coredata.MembershipOrderField], @@ -191,7 +480,7 @@ func (s *Service) GetAllOrganizationMemberships( err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := memberships.LoadByOrganizationID(ctx, conn, orgID, cursor); err != nil { + if err := memberships.LoadByOrganizationID(ctx, conn, s.scope, orgID, cursor); err != nil { return fmt.Errorf("failed to load organization memberships: %w", err) } @@ -205,7 +494,7 @@ func (s *Service) GetAllOrganizationMemberships( return page.NewPage(memberships, cursor), nil } -func (s *Service) CountOrganizationMemberships( +func (s *TenantAuthzService) CountOrganizationMemberships( ctx context.Context, orgID gid.GID, ) (int, error) { @@ -215,7 +504,7 @@ func (s *Service) CountOrganizationMemberships( func(conn pg.Conn) error { var memberships coredata.Memberships var err error - count, err = memberships.CountByOrganizationID(ctx, conn, orgID) + count, err = memberships.CountByOrganizationID(ctx, conn, s.scope, orgID) return err }, ) @@ -226,7 +515,28 @@ func (s *Service) CountOrganizationMemberships( return count, nil } -func (s *Service) CanUserAccessOrganization( +func (s *TenantAuthzService) CountOrganizationUsers( + ctx context.Context, + orgID gid.GID, +) (int, error) { + var count int + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + var users coredata.Users + var err error + count, err = users.CountByOrganizationID(ctx, conn, s.scope, orgID) + return err + }, + ) + if err != nil { + return 0, fmt.Errorf("failed to count users: %w", err) + } + + return count, nil +} + +func (s *TenantAuthzService) CanUserAccessOrganization( ctx context.Context, userID gid.GID, orgID gid.GID, @@ -238,7 +548,7 @@ func (s *Service) CanUserAccessOrganization( err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := membership.LoadByUserAndOrg(ctx, conn, userID, orgID); err != nil { + if err := membership.LoadByUserAndOrg(ctx, conn, s.scope, userID, orgID); err != nil { if _, ok := err.(coredata.ErrMembershipNotFound); ok { return nil // Not an error, just no access } @@ -256,7 +566,7 @@ func (s *Service) CanUserAccessOrganization( return haveAccess, nil } -func (s *Service) GetUserRoleInOrganization( +func (s *TenantAuthzService) GetUserRoleInOrganization( ctx context.Context, userID gid.GID, orgID gid.GID, @@ -266,7 +576,7 @@ func (s *Service) GetUserRoleInOrganization( err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := membership.LoadByUserAndOrg(ctx, conn, userID, orgID); err != nil { + if err := membership.LoadByUserAndOrg(ctx, conn, s.scope, userID, orgID); err != nil { return fmt.Errorf("failed to get user role: %w", err) } return nil @@ -280,7 +590,7 @@ func (s *Service) GetUserRoleInOrganization( return membership.Role, nil } -func (s *Service) RemoveMemberFromOrganization( +func (s *TenantAuthzService) RemoveMemberFromOrganization( ctx context.Context, orgID gid.GID, memberID gid.GID, @@ -290,7 +600,7 @@ func (s *Service) RemoveMemberFromOrganization( return s.pg.WithTx( ctx, func(tx pg.Conn) error { - if err := membership.LoadByID(ctx, tx, memberID); err != nil { + if err := membership.LoadByID(ctx, tx, s.scope, memberID); err != nil { return fmt.Errorf("failed to load membership: %w", err) } @@ -298,7 +608,7 @@ func (s *Service) RemoveMemberFromOrganization( return fmt.Errorf("membership does not belong to organization") } - if err := membership.Delete(ctx, tx); err != nil { + if err := membership.Delete(ctx, tx, s.scope); err != nil { return fmt.Errorf("failed to delete membership: %w", err) } @@ -307,32 +617,7 @@ func (s *Service) RemoveMemberFromOrganization( ) } -func (s *Service) AddUserToOrganization( - ctx context.Context, - userID gid.GID, - orgID gid.GID, - role string, -) error { - membership := &coredata.Membership{ - UserID: userID, - OrganizationID: orgID, - Role: role, - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - } - - return s.pg.WithConn( - ctx, - func(conn pg.Conn) error { - if err := membership.Create(ctx, conn); err != nil { - return fmt.Errorf("failed to add user to organization: %w", err) - } - return nil - }, - ) -} - -func (s *Service) UpdateUserRole( +func (s *TenantAuthzService) UpdateUserRole( ctx context.Context, userID gid.GID, orgID gid.GID, @@ -342,14 +627,14 @@ func (s *Service) UpdateUserRole( ctx, func(tx pg.Conn) error { membership := &coredata.Membership{} - if err := membership.LoadByUserAndOrg(ctx, tx, userID, orgID); err != nil { + if err := membership.LoadByUserAndOrg(ctx, tx, s.scope, userID, orgID); err != nil { return fmt.Errorf("failed to find membership: %w", err) } membership.Role = newRole membership.UpdatedAt = time.Now() - if err := membership.Update(ctx, tx); err != nil { + if err := membership.Update(ctx, tx, s.scope); err != nil { return fmt.Errorf("failed to update user role: %w", err) } @@ -358,7 +643,7 @@ func (s *Service) UpdateUserRole( ) } -func (s *Service) InviteUserToOrganization( +func (s *TenantAuthzService) InviteUserToOrganization( ctx context.Context, organizationID gid.GID, emailAddress string, @@ -380,12 +665,11 @@ func (s *Service) InviteUserToOrganization( } organization := &coredata.Organization{} - scope := coredata.NewScope(organizationID.TenantID()) - if err := organization.LoadByID(ctx, tx, scope, organizationID); err != nil { + if err := organization.LoadByID(ctx, tx, s.scope, organizationID); err != nil { return fmt.Errorf("failed to load organization: %w", err) } - invitationID := gid.New(organizationID.TenantID(), coredata.InvitationEntityType) + invitationID := gid.New(s.scope.GetTenantID(), coredata.InvitationEntityType) now := time.Now() invitation = &coredata.Invitation{ ID: invitationID, @@ -397,19 +681,20 @@ func (s *Service) InviteUserToOrganization( CreatedAt: now, } + body := bytes.NewBuffer(nil) + var err error if userExists { - membership := &coredata.Membership{ - UserID: user.ID, - OrganizationID: organizationID, - Role: role, - CreatedAt: now, - UpdatedAt: now, + err = invitationEmailBodyTemplate.Execute( + body, + map[string]string{ + "FullName": user.FullName, + "OrganizationName": organization.Name, + "InvitationURL": fmt.Sprintf("https://%s/", s.hostname), + }, + ) + if err != nil { + return fmt.Errorf("failed to execute template: %w", err) } - if err := membership.Create(ctx, tx); err != nil { - return fmt.Errorf("failed to add user to organization: %w", err) - } - - invitation.AcceptedAt = &now } else { invitationData := coredata.InvitationData{ InvitationID: invitationID, @@ -429,32 +714,31 @@ func (s *Service) InviteUserToOrganization( return fmt.Errorf("failed to generate invitation token: %w", err) } - body := bytes.NewBuffer(nil) err = invitationEmailBodyTemplate.Execute( body, map[string]string{ "FullName": fullName, "OrganizationName": organization.Name, - "InvitationURL": fmt.Sprintf("https://%s/auth/confirm-invitation?token=%s", s.hostname, invitationToken), + "InvitationURL": fmt.Sprintf("https://%s/auth/signup-from-invitation?token=%s&fullName=%s", s.hostname, invitationToken, url.QueryEscape(fullName)), }, ) if err != nil { return fmt.Errorf("failed to execute template: %w", err) } - - email := coredata.NewEmail( - fullName, - emailAddress, - invitationEmailSubject, - body.String(), - ) - - if err := email.Insert(ctx, tx); err != nil { - return fmt.Errorf("cannot insert email: %w", err) - } } - if err := invitation.Create(ctx, tx); err != nil { + email := coredata.NewEmail( + fullName, + emailAddress, + invitationEmailSubject, + body.String(), + ) + + if err := email.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot insert email: %w", err) + } + + if err := invitation.Create(ctx, tx, s.scope); err != nil { return fmt.Errorf("cannot create invitation: %w", err) } @@ -468,66 +752,8 @@ func (s *Service) InviteUserToOrganization( return invitation, nil } -func (s *Service) AcceptInvitation( - ctx context.Context, - token string, - userID gid.GID, -) error { - payload, err := statelesstoken.ValidateToken[coredata.InvitationData]( - s.tokenSecret, - TokenTypeOrganizationInvitation, - token, - ) - if err != nil { - return fmt.Errorf("invalid invitation token: %w", err) - } - invitationData := payload.Data - - return s.pg.WithTx( - ctx, - func(tx pg.Conn) error { - invitation := &coredata.Invitation{} - if err := invitation.LoadByID(ctx, tx, invitationData.InvitationID); err != nil { - var errInvitationNotFound *coredata.ErrInvitationNotFound - if errors.As(err, &errInvitationNotFound) { - return fmt.Errorf("invitation was deleted or no longer exists") - } - return fmt.Errorf("cannot load invitation: %w", err) - } - - if invitation.AcceptedAt != nil { - return fmt.Errorf("invitation already accepted") - } - - if time.Now().After(invitation.ExpiresAt) { - return fmt.Errorf("invitation expired") - } - - membership := &coredata.Membership{ - UserID: userID, - OrganizationID: invitation.OrganizationID, - Role: invitation.Role, - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - } - - if err := membership.Create(ctx, tx); err != nil { - return fmt.Errorf("failed to add user to organization: %w", err) - } - - now := time.Now() - invitation.AcceptedAt = &now - if err := invitation.Update(ctx, tx); err != nil { - return fmt.Errorf("failed to mark invitation as accepted: %w", err) - } - - return nil - }, - ) -} - // This is a placeholder for future permission system -func (s *Service) HasPermission( +func (s *TenantAuthzService) HasPermission( ctx context.Context, userID gid.GID, orgID gid.GID, @@ -538,22 +764,3 @@ func (s *Service) HasPermission( // In the future, this will check specific permissions based on role return s.CanUserAccessOrganization(ctx, userID, orgID) } - -func (s *Service) ListUserInvitations( - ctx context.Context, - email string, -) ([]*coredata.Invitation, error) { - var invitations coredata.Invitations - - err := s.pg.WithConn( - ctx, - func(conn pg.Conn) error { - if err := invitations.LoadByEmail(ctx, conn, email); err != nil { - return fmt.Errorf("failed to load invitations: %w", err) - } - return nil - }, - ) - - return invitations, err -} diff --git a/pkg/coredata/asset_vendor.go b/pkg/coredata/asset_vendor.go index e0e0884cf..4d4c1e826 100644 --- a/pkg/coredata/asset_vendor.go +++ b/pkg/coredata/asset_vendor.go @@ -163,7 +163,6 @@ JOIN snapshot_vendors sv ON sv.source_id = av.vendor_id query = fmt.Sprintf(query, scope.SQLFragment()) args := pgx.StrictNamedArgs{ - "tenant_id": scope.GetTenantID(), "snapshot_id": snapshotID, "organization_id": organizationID, } diff --git a/pkg/coredata/datum_vendor.go b/pkg/coredata/datum_vendor.go index a14693299..b3158a8fa 100644 --- a/pkg/coredata/datum_vendor.go +++ b/pkg/coredata/datum_vendor.go @@ -158,7 +158,6 @@ JOIN snapshot_vendors sv ON sv.source_id = dv.vendor_id query = fmt.Sprintf(query, scope.SQLFragment()) args := pgx.StrictNamedArgs{ - "tenant_id": scope.GetTenantID(), "snapshot_id": snapshotID, "organization_id": organizationID, } diff --git a/pkg/coredata/invitation.go b/pkg/coredata/invitation.go index 5a5be5010..8f50d6c53 100644 --- a/pkg/coredata/invitation.go +++ b/pkg/coredata/invitation.go @@ -81,17 +81,17 @@ func (i Invitation) CursorKey(orderBy InvitationOrderField) page.CursorKey { panic(fmt.Sprintf("unsupported order by: %s", orderBy)) } -// Tenant id scope is not applied because invitations are managed at the organization level and don't require tenant isolation. -func (i *Invitation) Create(ctx context.Context, conn pg.Conn) error { +func (i *Invitation) Create(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` INSERT INTO authz_invitations ( - id, organization_id, email, full_name, role, expires_at, created_at + tenant_id, id, organization_id, email, full_name, role, expires_at, created_at ) VALUES ( - @id, @organization_id, @email, @full_name, @role, @expires_at, @created_at + @tenant_id, @id, @organization_id, @email, @full_name, @role, @expires_at, @created_at ) ` args := pgx.StrictNamedArgs{ + "tenant_id": scope.GetTenantID(), "id": i.ID, "organization_id": i.OrganizationID, "email": i.Email, @@ -109,21 +109,24 @@ func (i *Invitation) Create(ctx context.Context, conn pg.Conn) error { return nil } -// Tenant id scope is not applied because we want to access invitations across all tenants for authentication purposes. func (i *Invitation) LoadByID( ctx context.Context, conn pg.Conn, + scope Scoper, id gid.GID, ) error { query := ` SELECT id, organization_id, email, full_name, role, expires_at, accepted_at, created_at FROM authz_invitations - WHERE id = @id + WHERE id = @id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "id": id, } + maps.Copy(args, scope.SQLArguments()) rows, err := conn.Query(ctx, query, args) if err != nil { @@ -142,18 +145,20 @@ func (i *Invitation) LoadByID( return nil } -// Tenant id scope is not applied because invitations are managed at the organization level and don't require tenant isolation. -func (i *Invitation) Update(ctx context.Context, conn pg.Conn) error { +func (i *Invitation) Update(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` UPDATE authz_invitations SET accepted_at = @accepted_at - WHERE id = @id + WHERE id = @id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "id": i.ID, "accepted_at": i.AcceptedAt, } + maps.Copy(args, scope.SQLArguments()) result, err := conn.Exec(ctx, query, args) if err != nil { @@ -167,16 +172,18 @@ func (i *Invitation) Update(ctx context.Context, conn pg.Conn) error { return nil } -// Tenant id scope is not applied because invitations are managed at the organization level and don't require tenant isolation. -func (i *Invitation) Delete(ctx context.Context, conn pg.Conn) error { +func (i *Invitation) Delete(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` DELETE FROM authz_invitations - WHERE id = @id + WHERE id = @id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "id": i.ID, } + maps.Copy(args, scope.SQLArguments()) result, err := conn.Exec(ctx, query, args) if err != nil { @@ -190,21 +197,30 @@ func (i *Invitation) Delete(ctx context.Context, conn pg.Conn) error { return nil } +// Tenant scope is not applied because this is used to query invitations across all tenants +// for a user who doesn't have tenant access yet (before accepting an invitation). func (i *Invitations) LoadByEmail( ctx context.Context, conn pg.Conn, email string, + cursor *page.Cursor[InvitationOrderField], + filter *InvitationFilter, ) error { query := ` SELECT id, organization_id, email, full_name, role, expires_at, accepted_at, created_at FROM authz_invitations - WHERE email = @email AND accepted_at IS NULL - ORDER BY created_at DESC + WHERE email = @email + AND %s + AND %s ` + query = fmt.Sprintf(query, filter.SQLFragment(), cursor.SQLFragment()) + args := pgx.StrictNamedArgs{ "email": email, } + maps.Copy(args, filter.SQLArguments()) + maps.Copy(args, cursor.SQLArguments()) rows, err := conn.Query(ctx, query, args) if err != nil { @@ -223,19 +239,23 @@ func (i *Invitations) LoadByEmail( func (i *Invitations) LoadByOrganizationID( ctx context.Context, conn pg.Conn, + scope Scoper, orgID gid.GID, cursor *page.Cursor[InvitationOrderField], ) error { query := ` SELECT id, organization_id, email, full_name, role, expires_at, accepted_at, created_at FROM authz_invitations - WHERE organization_id = @organization_id + WHERE organization_id = @organization_id AND %s AND %s ` - query = fmt.Sprintf(query, cursor.SQLFragment()) + query = fmt.Sprintf(query, scope.SQLFragment(), cursor.SQLFragment()) - args := pgx.StrictNamedArgs{"organization_id": orgID} + args := pgx.StrictNamedArgs{ + "organization_id": orgID, + } + maps.Copy(args, scope.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) rows, err := conn.Query(ctx, query, args) @@ -255,6 +275,7 @@ func (i *Invitations) LoadByOrganizationID( func (i *Invitations) CountByOrganizationID( ctx context.Context, conn pg.Conn, + scope Scoper, orgID gid.GID, ) (int, error) { q := ` @@ -263,10 +284,51 @@ SELECT FROM authz_invitations WHERE - organization_id = @organization_id + organization_id = @organization_id AND %s ` - args := pgx.StrictNamedArgs{"organization_id": orgID} + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{ + "organization_id": orgID, + } + maps.Copy(args, scope.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + err := row.Scan(&count) + if err != nil { + return 0, fmt.Errorf("cannot count invitations: %w", err) + } + + return count, nil +} + +// Tenant scope is not applied because this is used to count invitations across all tenants +// for a user who doesn't have tenant access yet (before accepting an invitation). +func (i *Invitations) CountByEmail( + ctx context.Context, + conn pg.Conn, + email string, + filter *InvitationFilter, +) (int, error) { + q := ` +SELECT + COUNT(*) +FROM + authz_invitations +WHERE + email = @email + AND %s +` + + q = fmt.Sprintf(q, filter.SQLFragment()) + + args := pgx.StrictNamedArgs{ + "email": email, + } + maps.Copy(args, filter.SQLArguments()) row := conn.QueryRow(ctx, q, args) diff --git a/pkg/coredata/invitation_filter.go b/pkg/coredata/invitation_filter.go new file mode 100644 index 000000000..7ae1935f7 --- /dev/null +++ b/pkg/coredata/invitation_filter.go @@ -0,0 +1,49 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package coredata + +import ( + "github.com/jackc/pgx/v5" +) + +type ( + InvitationFilter struct { + onlyPending *bool + } +) + +func NewInvitationFilter(onlyPending *bool) *InvitationFilter { + return &InvitationFilter{ + onlyPending: onlyPending, + } +} + +func (f *InvitationFilter) SQLArguments() pgx.NamedArgs { + return pgx.NamedArgs{ + "only_pending": f.onlyPending, + } +} + +func (f *InvitationFilter) SQLFragment() string { + return ` +( + CASE + WHEN @only_pending::boolean IS NOT NULL AND @only_pending::boolean = true THEN + (accepted_at IS NULL AND expires_at > NOW()) + ELSE TRUE + END +)` +} + diff --git a/pkg/coredata/membership.go b/pkg/coredata/membership.go index 5b8140336..295e1210f 100644 --- a/pkg/coredata/membership.go +++ b/pkg/coredata/membership.go @@ -76,28 +76,20 @@ func (m Membership) CursorKey(orderBy MembershipOrderField) page.CursorKey { panic(fmt.Sprintf("unsupported order by: %s", orderBy)) } -// Tenant id scope is not applied because memberships are managed at the organization level and don't require tenant isolation. -func (m *Membership) Create(ctx context.Context, conn pg.Conn) error { +func (m *Membership) Create(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` - INSERT INTO authz_memberships (id, user_id, organization_id, role, created_at, updated_at) - SELECT - generate_gid(decode_base64_unpadded(o.tenant_id), @entity_type), - @user_id, - @organization_id, - @role, - @created_at, - @updated_at - FROM organizations o - WHERE o.id = @organization_id + INSERT INTO authz_memberships (tenant_id, id, user_id, organization_id, role, created_at, updated_at) + VALUES (@tenant_id, @id, @user_id, @organization_id, @role, @created_at, @updated_at) ` args := pgx.StrictNamedArgs{ + "tenant_id": scope.GetTenantID(), + "id": m.ID, "user_id": m.UserID, "organization_id": m.OrganizationID, "role": m.Role, "created_at": m.CreatedAt, "updated_at": m.UpdatedAt, - "entity_type": MembershipEntityType, } result, err := conn.Exec(ctx, query, args) @@ -116,10 +108,10 @@ func (m *Membership) Create(ctx context.Context, conn pg.Conn) error { return nil } -// Tenant id scope is not applied because we want to access memberships across all tenants for authentication purposes. func (m *Membership) LoadByID( ctx context.Context, conn pg.Conn, + scope Scoper, membershipID gid.GID, ) error { query := ` @@ -134,12 +126,15 @@ func (m *Membership) LoadByID( m.updated_at FROM authz_memberships m JOIN users u ON m.user_id = u.id - WHERE m.id = @membership_id + WHERE m.id = @membership_id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "membership_id": membershipID, } + maps.Copy(args, scope.SQLArguments()) rows, err := conn.Query(ctx, query, args) if err != nil { @@ -158,10 +153,10 @@ func (m *Membership) LoadByID( return nil } -// Tenant id scope is not applied because we want to access memberships across all tenants for authentication purposes. func (m *Membership) LoadByUserAndOrg( ctx context.Context, conn pg.Conn, + scope Scoper, userID gid.GID, orgID gid.GID, ) error { @@ -177,13 +172,16 @@ func (m *Membership) LoadByUserAndOrg( m.updated_at FROM authz_memberships m JOIN users u ON m.user_id = u.id - WHERE m.user_id = @user_id AND m.organization_id = @organization_id + WHERE m.user_id = @user_id AND m.organization_id = @organization_id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "user_id": userID, "organization_id": orgID, } + maps.Copy(args, scope.SQLArguments()) rows, err := conn.Query(ctx, query, args) if err != nil { @@ -202,19 +200,21 @@ func (m *Membership) LoadByUserAndOrg( return nil } -// Tenant id scope is not applied because memberships are managed at the organization level and don't require tenant isolation. -func (m *Membership) Update(ctx context.Context, conn pg.Conn) error { +func (m *Membership) Update(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` UPDATE authz_memberships SET role = @role, updated_at = @updated_at - WHERE id = @id + WHERE id = @id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "id": m.ID, "role": m.Role, "updated_at": m.UpdatedAt, } + maps.Copy(args, scope.SQLArguments()) result, err := conn.Exec(ctx, query, args) if err != nil { @@ -228,16 +228,18 @@ func (m *Membership) Update(ctx context.Context, conn pg.Conn) error { return nil } -// Tenant id scope is not applied because memberships are managed at the organization level and don't require tenant isolation. -func (m *Membership) Delete(ctx context.Context, conn pg.Conn) error { +func (m *Membership) Delete(ctx context.Context, conn pg.Conn, scope Scoper) error { query := ` DELETE FROM authz_memberships - WHERE id = @id + WHERE id = @id AND %s ` + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ "id": m.ID, } + maps.Copy(args, scope.SQLArguments()) result, err := conn.Exec(ctx, query, args) if err != nil { @@ -251,10 +253,10 @@ func (m *Membership) Delete(ctx context.Context, conn pg.Conn) error { return nil } -// Tenant id scope is not applied because we want to access all user's memberships across tenants for authentication purposes. func (m *Memberships) LoadByUserID( ctx context.Context, conn pg.Conn, + scope Scoper, userID gid.GID, ) error { query := ` @@ -272,11 +274,17 @@ FROM JOIN users u ON m.user_id = u.id WHERE m.user_id = @user_id + AND %s ORDER BY m.created_at DESC ` - args := pgx.StrictNamedArgs{"user_id": userID} + query = fmt.Sprintf(query, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{ + "user_id": userID, + } + maps.Copy(args, scope.SQLArguments()) rows, err := conn.Query(ctx, query, args) if err != nil { @@ -292,10 +300,10 @@ ORDER BY return nil } -// Tenant id scope is not applied because we want to access memberships across all tenants for authentication purposes. func (m *Memberships) LoadByOrganizationID( ctx context.Context, conn pg.Conn, + scope Scoper, organizationID gid.GID, cursor *page.Cursor[MembershipOrderField], ) error { @@ -315,11 +323,15 @@ JOIN users u ON m.user_id = u.id WHERE m.organization_id = @organization_id AND %s + AND %s ` - query = fmt.Sprintf(query, cursor.SQLFragment()) + query = fmt.Sprintf(query, scope.SQLFragment(), cursor.SQLFragment()) - args := pgx.StrictNamedArgs{"organization_id": organizationID} + args := pgx.StrictNamedArgs{ + "organization_id": organizationID, + } + maps.Copy(args, scope.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) rows, err := conn.Query(ctx, query, args) @@ -339,14 +351,19 @@ WHERE func (m *Memberships) CountByOrganizationID( ctx context.Context, conn pg.Conn, + scope Scoper, organizationID gid.GID, ) (int, error) { query := ` SELECT COUNT(*) FROM authz_memberships - WHERE organization_id = @organization_id + WHERE organization_id = @organization_id AND %s ` - args := pgx.StrictNamedArgs{"organization_id": organizationID} + query = fmt.Sprintf(query, scope.SQLFragment()) + args := pgx.StrictNamedArgs{ + "organization_id": organizationID, + } + maps.Copy(args, scope.SQLArguments()) row := conn.QueryRow(ctx, query, args) var count int if err := row.Scan(&count); err != nil { diff --git a/pkg/coredata/migrations/20251006T220024Z.sql b/pkg/coredata/migrations/20251006T220024Z.sql index b520f1959..5722f31f2 100644 --- a/pkg/coredata/migrations/20251006T220024Z.sql +++ b/pkg/coredata/migrations/20251006T220024Z.sql @@ -8,6 +8,7 @@ CREATE TYPE authz_role AS ENUM ('OWNER', 'ADMIN', 'MEMBER', 'VIEWER'); -- Create authz_memberships table with id as primary key CREATE TABLE authz_memberships ( id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, user_id TEXT NOT NULL, organization_id TEXT NOT NULL, role authz_role NOT NULL, @@ -19,6 +20,7 @@ CREATE TABLE authz_memberships ( -- Create authz_invitations table CREATE TABLE authz_invitations ( id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, organization_id TEXT NOT NULL, email TEXT NOT NULL, full_name TEXT NOT NULL, @@ -29,8 +31,9 @@ CREATE TABLE authz_invitations ( ); -- Copy data from users_organizations to authz_memberships -INSERT INTO authz_memberships (id, user_id, organization_id, role, created_at, updated_at) +INSERT INTO authz_memberships (tenant_id, id, user_id, organization_id, role, created_at, updated_at) SELECT + organizations.tenant_id, generate_gid(decode_base64_unpadded(organizations.tenant_id), 38) as id, users_organizations.user_id, users_organizations.organization_id, diff --git a/pkg/coredata/people.go b/pkg/coredata/people.go index d7466b4cc..6ee477f1f 100644 --- a/pkg/coredata/people.go +++ b/pkg/coredata/people.go @@ -32,7 +32,6 @@ type ( ID gid.GID `db:"id"` OrganizationID gid.GID `db:"organization_id"` Kind PeopleKind `db:"kind"` - UserID *gid.GID `db:"user_id"` FullName string `db:"full_name"` PrimaryEmailAddress string `db:"primary_email_address"` AdditionalEmailAddresses []string `db:"additional_email_addresses"` @@ -78,7 +77,6 @@ SELECT id, organization_id, kind, - user_id, full_name, primary_email_address, additional_email_addresses, @@ -126,7 +124,6 @@ func (p *People) LoadByEmail( id, organization_id, kind, - user_id, full_name, primary_email_address, additional_email_addresses, @@ -167,58 +164,6 @@ func (p *People) LoadByEmail( return nil } -func (p *People) LoadByUserID( - ctx context.Context, - conn pg.Conn, - scope Scoper, - userID gid.GID, -) error { - q := ` -SELECT - id, - organization_id, - kind, - user_id, - full_name, - primary_email_address, - additional_email_addresses, - position, - contract_start_date, - contract_end_date, - created_at, - updated_at -FROM - peoples -WHERE - %s - AND user_id = @user_id -LIMIT 1; -` - - q = fmt.Sprintf(q, scope.SQLFragment()) - - args := pgx.StrictNamedArgs{"user_id": userID} - maps.Copy(args, scope.SQLArguments()) - - rows, err := conn.Query(ctx, q, args) - if err != nil { - return fmt.Errorf("cannot query people: %w", err) - } - - people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) - if err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return &ErrPeopleNotFound{Identifier: userID.String()} - } - - return fmt.Errorf("cannot collect people: %w", err) - } - - *p = people - - return nil -} - func (p People) Insert( ctx context.Context, conn pg.Conn, @@ -230,7 +175,6 @@ INSERT INTO tenant_id, id, organization_id, - user_id, kind, full_name, primary_email_address, @@ -245,7 +189,6 @@ VALUES ( @tenant_id, @people_id, @organization_id, - @user_id, @kind, @full_name, @primary_email_address, @@ -262,7 +205,6 @@ VALUES ( "tenant_id": scope.GetTenantID(), "people_id": p.ID, "organization_id": p.OrganizationID, - "user_id": p.UserID, "kind": p.Kind, "full_name": p.FullName, "primary_email_address": p.PrimaryEmailAddress, @@ -343,7 +285,6 @@ SELECT id, organization_id, kind, - user_id, full_name, primary_email_address, additional_email_addresses, @@ -390,7 +331,6 @@ func (p *People) Update( ) error { q := ` UPDATE peoples SET - user_id = @user_id, full_name = @full_name, primary_email_address = @primary_email_address, additional_email_addresses = @additional_email_addresses, @@ -406,7 +346,6 @@ WHERE %s args := pgx.StrictNamedArgs{ "people_id": p.ID, - "user_id": p.UserID, "full_name": p.FullName, "primary_email_address": p.PrimaryEmailAddress, "additional_email_addresses": p.AdditionalEmailAddresses, @@ -447,7 +386,6 @@ SELECT id, organization_id, kind, - user_id, full_name, primary_email_address, additional_email_addresses, diff --git a/pkg/coredata/user.go b/pkg/coredata/user.go index 0229f12ca..787723df0 100644 --- a/pkg/coredata/user.go +++ b/pkg/coredata/user.go @@ -115,6 +115,7 @@ WHERE func (u *Users) CountByOrganizationID( ctx context.Context, conn pg.Conn, + scope Scoper, organizationID gid.GID, ) (int, error) { q := ` @@ -124,11 +125,14 @@ FROM users WHERE id IN ( - SELECT user_id FROM authz_memberships WHERE organization_id = @organization_id + SELECT user_id FROM authz_memberships WHERE organization_id = @organization_id AND %s ) ` + q = fmt.Sprintf(q, scope.SQLFragment()) + args := pgx.StrictNamedArgs{"organization_id": organizationID} + maps.Copy(args, scope.SQLArguments()) row := conn.QueryRow(ctx, q, args) diff --git a/pkg/coredata/user_organization.go b/pkg/coredata/user_organization.go deleted file mode 100644 index 0009d7be0..000000000 --- a/pkg/coredata/user_organization.go +++ /dev/null @@ -1,81 +0,0 @@ -// Copyright (c) 2025 Probo Inc . -// -// Permission to use, copy, modify, and/or distribute this software for any -// purpose with or without fee is hereby granted, provided that the above -// copyright notice and this permission notice appear in all copies. -// -// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH -// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY -// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, -// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM -// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR -// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR -// PERFORMANCE OF THIS SOFTWARE. - -package coredata - -import ( - "context" - "time" - - "github.com/getprobo/probo/pkg/gid" - "github.com/jackc/pgx/v5" - "go.gearno.de/kit/pg" -) - -type ( - UserOrganization struct { - UserID gid.GID `db:"user_id"` - OrganizationID gid.GID `db:"organization_id"` - CreatedAt time.Time `db:"created_at"` - } - - UserOrganizations []*UserOrganization -) - -func (uo UserOrganization) Insert( - ctx context.Context, - conn pg.Conn, -) error { - q := ` -INSERT INTO users_organizations (user_id, organization_id, created_at) -VALUES (@user_id, @organization_id, @created_at) -` - - _, err := conn.Exec(ctx, q, pgx.StrictNamedArgs{"user_id": uo.UserID, "organization_id": uo.OrganizationID, "created_at": uo.CreatedAt}) - return err -} - -// Tenant id scope is not applied because user organizations are managed at the organization level and don't require tenant isolation. -func (uo UserOrganization) Delete(ctx context.Context, conn pg.Conn) error { - q := ` -DELETE FROM users_organizations WHERE user_id = @user_id AND organization_id = @organization_id -` - - _, err := conn.Exec(ctx, q, pgx.StrictNamedArgs{"user_id": uo.UserID, "organization_id": uo.OrganizationID}) - return err -} - -func (uo *UserOrganizations) ForUserID( - ctx context.Context, - conn pg.Conn, - userID gid.GID, -) error { - q := ` -SELECT user_id, organization_id, created_at FROM users_organizations WHERE user_id = @user_id -` - - rows, err := conn.Query(ctx, q, pgx.StrictNamedArgs{"user_id": userID}) - if err != nil { - return err - } - - userOrganizations, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[UserOrganization]) - if err != nil { - return err - } - - *uo = userOrganizations - - return nil -} diff --git a/pkg/probo/people_service.go b/pkg/probo/people_service.go index 1a9d00755..1b468d954 100644 --- a/pkg/probo/people_service.go +++ b/pkg/probo/people_service.go @@ -32,7 +32,6 @@ type ( UpdatePeopleRequest struct { ID gid.GID - UserID *gid.GID Kind *coredata.PeopleKind FullName *string PrimaryEmailAddress *string @@ -44,7 +43,6 @@ type ( CreatePeopleRequest struct { OrganizationID gid.GID - UserID *gid.GID FullName string PrimaryEmailAddress string AdditionalEmailAddresses []string @@ -75,26 +73,6 @@ func (s PeopleService) Get( return people, nil } -func (s PeopleService) GetByUserID( - ctx context.Context, - userID gid.GID, -) (*coredata.People, error) { - people := &coredata.People{} - - err := s.svc.pg.WithConn( - ctx, - func(conn pg.Conn) error { - return people.LoadByUserID(ctx, conn, s.svc.scope, userID) - }, - ) - - if err != nil { - return nil, err - } - - return people, nil -} - func (s PeopleService) CountForOrganizationID( ctx context.Context, organizationID gid.GID, @@ -164,10 +142,6 @@ func (s PeopleService) Update( return fmt.Errorf("cannot load people: %w", err) } - if req.UserID != nil { - people.UserID = req.UserID - } - if req.Kind != nil { people.Kind = *req.Kind } @@ -234,7 +208,6 @@ func (s PeopleService) Create( FullName: req.FullName, PrimaryEmailAddress: req.PrimaryEmailAddress, AdditionalEmailAddresses: req.AdditionalEmailAddresses, - UserID: req.UserID, Position: req.Position, ContractStartDate: req.ContractStartDate, ContractEndDate: req.ContractEndDate, diff --git a/pkg/server/api/console/v1/invitation_confirmation_handler.go b/pkg/server/api/console/v1/invitation_confirmation_handler.go deleted file mode 100644 index 1654abf43..000000000 --- a/pkg/server/api/console/v1/invitation_confirmation_handler.go +++ /dev/null @@ -1,81 +0,0 @@ -// Copyright (c) 2025 Probo Inc . -// -// Permission to use, copy, modify, and/or distribute this software for any -// purpose with or without fee is hereby granted, provided that the above -// copyright notice and this permission notice appear in all copies. -// -// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH -// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY -// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, -// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM -// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR -// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR -// PERFORMANCE OF THIS SOFTWARE. - -package console_v1 - -import ( - "encoding/json" - "errors" - "fmt" - "net/http" - - "github.com/getprobo/probo/pkg/auth" - "github.com/getprobo/probo/pkg/authz" - "github.com/getprobo/probo/pkg/coredata" - "github.com/getprobo/probo/pkg/statelesstoken" - "go.gearno.de/kit/httpserver" -) - -type ( - InvitationConfirmationRequest struct { - Token string `json:"token"` - Password string `json:"password"` - } - - InvitationConfirmationResponse struct { - } -) - -func InvitationConfirmationHandler(authSvc *auth.Service, authzSvc *authz.Service, authCfg AuthConfig) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - var req InvitationConfirmationRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("cannot decode body: %w", err)) - return - } - - payload, err := statelesstoken.ValidateToken[coredata.InvitationData]( - authCfg.CookieSecret, - authz.TokenTypeOrganizationInvitation, - req.Token, - ) - if err != nil { - httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("invalid invitation token: %w", err)) - return - } - - user, _, err := authSvc.SignUp(r.Context(), payload.Data.Email, req.Password, payload.Data.FullName) - if err != nil { - var errUserAlreadyExists *auth.ErrUserAlreadyExists - if errors.As(err, &errUserAlreadyExists) { - user, err = authSvc.GetUserByEmail(r.Context(), payload.Data.Email) - if err != nil { - httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("failed to load existing user: %w", err)) - return - } - } else { - httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("failed to create user: %w", err)) - return - } - } - - err = authzSvc.AcceptInvitation(r.Context(), req.Token, user.ID) - if err != nil { - httpserver.RenderError(w, http.StatusInternalServerError, err) - return - } - - httpserver.RenderJSON(w, http.StatusOK, InvitationConfirmationResponse{}) - } -} diff --git a/pkg/server/api/console/v1/resolver.go b/pkg/server/api/console/v1/resolver.go index 4337d4172..7deb3905c 100644 --- a/pkg/server/api/console/v1/resolver.go +++ b/pkg/server/api/console/v1/resolver.go @@ -21,6 +21,7 @@ import ( "encoding/json" "fmt" "net/http" + "slices" "strings" "time" @@ -157,7 +158,7 @@ func NewMux( r.Post("/auth/register", SignUpHandler(authSvc, authCfg)) r.Post("/auth/login", SignInHandler(authSvc, authCfg)) r.Delete("/auth/logout", SignOutHandler(authSvc, authCfg)) - r.Post("/auth/invitation", InvitationConfirmationHandler(authSvc, authzSvc, authCfg)) + r.Post("/auth/signup-from-invitation", SignupFromInvitationHandler(authSvc, authCfg)) r.Post("/auth/forget-password", ForgetPasswordHandler(authSvc, authCfg)) r.Post("/auth/reset-password", ResetPasswordHandler(authSvc, authCfg)) @@ -316,18 +317,28 @@ func (r *Resolver) ProboService(ctx context.Context, tenantID gid.TenantID) *pro return GetTenantService(ctx, r.proboSvc, tenantID) } +func (r *Resolver) AuthzService(ctx context.Context, tenantID gid.TenantID) *authz.TenantAuthzService { + return GetTenantAuthzService(ctx, r.authzSvc, tenantID) +} + func GetTenantService(ctx context.Context, proboSvc *probo.Service, tenantID gid.TenantID) *probo.TenantService { + validateTenantAccess(ctx, tenantID) + return proboSvc.WithTenant(tenantID) +} + +func GetTenantAuthzService(ctx context.Context, authzSvc *authz.Service, tenantID gid.TenantID) *authz.TenantAuthzService { + validateTenantAccess(ctx, tenantID) + return authzSvc.WithTenant(tenantID) +} + +func validateTenantAccess(ctx context.Context, tenantID gid.TenantID) { tenantIDs, _ := ctx.Value(userTenantContextKey).(*[]gid.TenantID) if tenantIDs == nil { panic(fmt.Errorf("tenant not found")) } - for _, id := range *tenantIDs { - if id == tenantID { - return proboSvc.WithTenant(tenantID) - } + if !slices.Contains(*tenantIDs, tenantID) { + panic(fmt.Errorf("tenant not found")) } - - panic(fmt.Errorf("tenant not found")) } diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index d4afddacb..cf717d889 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1465,6 +1465,10 @@ input InvitationOrder { field: InvitationOrderField! } +input InvitationFilter { + onlyPending: Boolean +} + input DocumentVersionFilter { status: DocumentStatus } @@ -1572,6 +1576,7 @@ type Organization implements Node { last: Int before: CursorKey orderBy: InvitationOrder + filter: InvitationFilter ): InvitationConnection! @goField(forceResolver: true) connectors( @@ -1736,8 +1741,6 @@ type User implements Node { email: String! createdAt: Datetime! updatedAt: Datetime! - - people(organizationId: ID!): People @goField(forceResolver: true) } type Membership implements Node { @@ -1759,6 +1762,7 @@ type Invitation implements Node { expiresAt: Datetime! acceptedAt: Datetime createdAt: Datetime! + organization: Organization! @goField(forceResolver: true) } type Connector implements Node { @@ -2295,6 +2299,15 @@ type Viewer { before: CursorKey orderBy: OrganizationOrder ): OrganizationConnection! @goField(forceResolver: true) + + invitations( + first: Int + after: CursorKey + last: Int + before: CursorKey + orderBy: InvitationOrder + filter: InvitationFilter + ): InvitationConnection! @goField(forceResolver: true) } # Connection Types @@ -2394,13 +2407,19 @@ type TrustCenterReferenceEdge { node: TrustCenterReference! } -type UserConnection { +type UserConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.UserConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [UserEdge!]! pageInfo: PageInfo! } -type MembershipConnection { +type MembershipConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.MembershipConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [MembershipEdge!]! pageInfo: PageInfo! @@ -2709,7 +2728,10 @@ type File { updatedAt: Datetime! } -type InvitationConnection { +type InvitationConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.InvitationConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [InvitationEdge!]! pageInfo: PageInfo! @@ -2780,6 +2802,7 @@ type Mutation { # User mutations confirmEmail(input: ConfirmEmailInput!): ConfirmEmailPayload! inviteUser(input: InviteUserInput!): InviteUserPayload! + acceptInvitation(input: AcceptInvitationInput!): AcceptInvitationPayload! deleteInvitation(input: DeleteInvitationInput!): DeleteInvitationPayload! removeMember(input: RemoveMemberInput!): RemoveMemberPayload! @@ -3540,6 +3563,10 @@ input InviteUserInput { createPeople: Boolean! } +input AcceptInvitationInput { + invitationId: ID! +} + input DeleteInvitationInput { invitationId: ID! } @@ -4066,12 +4093,16 @@ type InviteUserPayload { invitationEdge: InvitationEdge! } +type AcceptInvitationPayload { + invitation: Invitation! +} + type DeleteInvitationPayload { deletedInvitationId: ID! } type RemoveMemberPayload { - success: Boolean! + deletedMemberId: ID! } input VendorRiskAssessmentOrder { diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index d55c227f9..4ac937b81 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -64,6 +64,7 @@ type ResolverRoot interface { File() FileResolver Framework() FrameworkResolver FrameworkConnection() FrameworkConnectionResolver + Invitation() InvitationResolver InvitationConnection() InvitationConnectionResolver Measure() MeasureResolver MeasureConnection() MeasureConnectionResolver @@ -91,7 +92,6 @@ type ResolverRoot interface { TrustCenterDocumentAccessConnection() TrustCenterDocumentAccessConnectionResolver TrustCenterReference() TrustCenterReferenceResolver TrustCenterReferenceConnection() TrustCenterReferenceConnectionResolver - User() UserResolver UserConnection() UserConnectionResolver Vendor() VendorResolver VendorBusinessAssociateAgreement() VendorBusinessAssociateAgreementResolver @@ -108,6 +108,10 @@ type DirectiveRoot struct { } type ComplexityRoot struct { + AcceptInvitationPayload struct { + Invitation func(childComplexity int) int + } + AssessVendorPayload struct { Vendor func(childComplexity int) int } @@ -759,13 +763,14 @@ type ComplexityRoot struct { } Invitation struct { - AcceptedAt func(childComplexity int) int - CreatedAt func(childComplexity int) int - Email func(childComplexity int) int - ExpiresAt func(childComplexity int) int - FullName func(childComplexity int) int - ID func(childComplexity int) int - Role func(childComplexity int) int + AcceptedAt func(childComplexity int) int + CreatedAt func(childComplexity int) int + Email func(childComplexity int) int + ExpiresAt func(childComplexity int) int + FullName func(childComplexity int) int + ID func(childComplexity int) int + Organization func(childComplexity int) int + Role func(childComplexity int) int } InvitationConnection struct { @@ -831,6 +836,7 @@ type ComplexityRoot struct { } Mutation struct { + AcceptInvitation func(childComplexity int, input types.AcceptInvitationInput) int AssessVendor func(childComplexity int, input types.AssessVendorInput) int AssignTask func(childComplexity int, input types.AssignTaskInput) int BulkDeleteDocuments func(childComplexity int, input types.BulkDeleteDocumentsInput) int @@ -1026,7 +1032,7 @@ type ComplexityRoot struct { HeadquarterAddress func(childComplexity int) int HorizontalLogoURL func(childComplexity int) int ID func(childComplexity int) int - Invitations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder) int + Invitations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) int LogoURL func(childComplexity int) int Measures func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.MeasureOrderBy, filter *types.MeasureFilter) int Memberships func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.MembershipOrderBy) int @@ -1131,7 +1137,7 @@ type ComplexityRoot struct { } RemoveMemberPayload struct { - Success func(childComplexity int) int + DeletedMemberID func(childComplexity int) int } Report struct { @@ -1459,7 +1465,6 @@ type ComplexityRoot struct { Email func(childComplexity int) int FullName func(childComplexity int) int ID func(childComplexity int) int - People func(childComplexity int, organizationID gid.GID) int UpdatedAt func(childComplexity int) int } @@ -1627,6 +1632,7 @@ type ComplexityRoot struct { Viewer struct { ID func(childComplexity int) int + Invitations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) int Organizations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.OrganizationOrder) int User func(childComplexity int) int } @@ -1719,6 +1725,9 @@ type FrameworkResolver interface { type FrameworkConnectionResolver interface { TotalCount(ctx context.Context, obj *types.FrameworkConnection) (int, error) } +type InvitationResolver interface { + Organization(ctx context.Context, obj *types.Invitation) (*types.Organization, error) +} type InvitationConnectionResolver interface { TotalCount(ctx context.Context, obj *types.InvitationConnection) (int, error) } @@ -1750,6 +1759,7 @@ type MutationResolver interface { DeleteTrustCenterReference(ctx context.Context, input types.DeleteTrustCenterReferenceInput) (*types.DeleteTrustCenterReferencePayload, error) ConfirmEmail(ctx context.Context, input types.ConfirmEmailInput) (*types.ConfirmEmailPayload, error) InviteUser(ctx context.Context, input types.InviteUserInput) (*types.InviteUserPayload, error) + AcceptInvitation(ctx context.Context, input types.AcceptInvitationInput) (*types.AcceptInvitationPayload, error) DeleteInvitation(ctx context.Context, input types.DeleteInvitationInput) (*types.DeleteInvitationPayload, error) RemoveMember(ctx context.Context, input types.RemoveMemberInput) (*types.RemoveMemberPayload, error) CreatePeople(ctx context.Context, input types.CreatePeopleInput) (*types.CreatePeoplePayload, error) @@ -1878,7 +1888,7 @@ type OrganizationResolver interface { HorizontalLogoURL(ctx context.Context, obj *types.Organization) (*string, error) Memberships(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.MembershipOrderBy) (*types.MembershipConnection, error) - Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder) (*types.InvitationConnection, error) + Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) Connectors(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.ConnectorOrder) (*types.ConnectorConnection, error) Frameworks(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.FrameworkOrderBy) (*types.FrameworkConnection, error) Controls(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.ControlOrderBy, filter *types.ControlFilter) (*types.ControlConnection, error) @@ -1969,9 +1979,6 @@ type TrustCenterReferenceResolver interface { type TrustCenterReferenceConnectionResolver interface { TotalCount(ctx context.Context, obj *types.TrustCenterReferenceConnection) (int, error) } -type UserResolver interface { - People(ctx context.Context, obj *types.User, organizationID gid.GID) (*types.People, error) -} type UserConnectionResolver interface { TotalCount(ctx context.Context, obj *types.UserConnection) (int, error) } @@ -2015,6 +2022,7 @@ type VendorServiceResolver interface { } type ViewerResolver interface { Organizations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.OrganizationOrder) (*types.OrganizationConnection, error) + Invitations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) } type executableSchema struct { @@ -2036,6 +2044,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin _ = ec switch typeName + "." + field { + case "AcceptInvitationPayload.invitation": + if e.complexity.AcceptInvitationPayload.Invitation == nil { + break + } + + return e.complexity.AcceptInvitationPayload.Invitation(childComplexity), true + case "AssessVendorPayload.vendor": if e.complexity.AssessVendorPayload.Vendor == nil { break @@ -4142,6 +4157,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.Invitation.ID(childComplexity), true + case "Invitation.organization": + if e.complexity.Invitation.Organization == nil { + break + } + + return e.complexity.Invitation.Organization(childComplexity), true + case "Invitation.role": if e.complexity.Invitation.Role == nil { break @@ -4414,6 +4436,18 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.MembershipEdge.Node(childComplexity), true + case "Mutation.acceptInvitation": + if e.complexity.Mutation.AcceptInvitation == nil { + break + } + + args, err := ec.field_Mutation_acceptInvitation_args(ctx, rawArgs) + if err != nil { + return 0, false + } + + return e.complexity.Mutation.AcceptInvitation(childComplexity, args["input"].(types.AcceptInvitationInput)), true + case "Mutation.assessVendor": if e.complexity.Mutation.AssessVendor == nil { break @@ -6277,7 +6311,7 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return 0, false } - return e.complexity.Organization.Invitations(childComplexity, args["first"].(*int), args["after"].(*page.CursorKey), args["last"].(*int), args["before"].(*page.CursorKey), args["orderBy"].(*types.InvitationOrder)), true + return e.complexity.Organization.Invitations(childComplexity, args["first"].(*int), args["after"].(*page.CursorKey), args["last"].(*int), args["before"].(*page.CursorKey), args["orderBy"].(*types.InvitationOrder), args["filter"].(*types.InvitationFilter)), true case "Organization.logoUrl": if e.complexity.Organization.LogoURL == nil { @@ -6810,12 +6844,12 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.Query.Viewer(childComplexity), true - case "RemoveMemberPayload.success": - if e.complexity.RemoveMemberPayload.Success == nil { + case "RemoveMemberPayload.deletedMemberId": + if e.complexity.RemoveMemberPayload.DeletedMemberID == nil { break } - return e.complexity.RemoveMemberPayload.Success(childComplexity), true + return e.complexity.RemoveMemberPayload.DeletedMemberID(childComplexity), true case "Report.audit": if e.complexity.Report.Audit == nil { @@ -7933,18 +7967,6 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.User.ID(childComplexity), true - case "User.people": - if e.complexity.User.People == nil { - break - } - - args, err := ec.field_User_people_args(ctx, rawArgs) - if err != nil { - return 0, false - } - - return e.complexity.User.People(childComplexity, args["organizationId"].(gid.GID)), true - case "User.updatedAt": if e.complexity.User.UpdatedAt == nil { break @@ -8714,6 +8736,18 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.Viewer.ID(childComplexity), true + case "Viewer.invitations": + if e.complexity.Viewer.Invitations == nil { + break + } + + args, err := ec.field_Viewer_invitations_args(ctx, rawArgs) + if err != nil { + return 0, false + } + + return e.complexity.Viewer.Invitations(childComplexity, args["first"].(*int), args["after"].(*page.CursorKey), args["last"].(*int), args["before"].(*page.CursorKey), args["orderBy"].(*types.InvitationOrder), args["filter"].(*types.InvitationFilter)), true + case "Viewer.organizations": if e.complexity.Viewer.Organizations == nil { break @@ -8741,6 +8775,7 @@ func (e *executableSchema) Exec(ctx context.Context) graphql.ResponseHandler { opCtx := graphql.GetOperationContext(ctx) ec := executionContext{opCtx, e, 0, 0, make(chan graphql.DeferredResult)} inputUnmarshalMap := graphql.BuildUnmarshalerMap( + ec.unmarshalInputAcceptInvitationInput, ec.unmarshalInputAssessVendorInput, ec.unmarshalInputAssetFilter, ec.unmarshalInputAssetOrder, @@ -8843,6 +8878,7 @@ func (e *executableSchema) Exec(ctx context.Context) graphql.ResponseHandler { ec.unmarshalInputGenerateFrameworkStateOfApplicabilityInput, ec.unmarshalInputImportFrameworkInput, ec.unmarshalInputImportMeasureInput, + ec.unmarshalInputInvitationFilter, ec.unmarshalInputInvitationOrder, ec.unmarshalInputInviteUserInput, ec.unmarshalInputMeasureFilter, @@ -10471,6 +10507,10 @@ input InvitationOrder { field: InvitationOrderField! } +input InvitationFilter { + onlyPending: Boolean +} + input DocumentVersionFilter { status: DocumentStatus } @@ -10578,6 +10618,7 @@ type Organization implements Node { last: Int before: CursorKey orderBy: InvitationOrder + filter: InvitationFilter ): InvitationConnection! @goField(forceResolver: true) connectors( @@ -10742,8 +10783,6 @@ type User implements Node { email: String! createdAt: Datetime! updatedAt: Datetime! - - people(organizationId: ID!): People @goField(forceResolver: true) } type Membership implements Node { @@ -10765,6 +10804,7 @@ type Invitation implements Node { expiresAt: Datetime! acceptedAt: Datetime createdAt: Datetime! + organization: Organization! @goField(forceResolver: true) } type Connector implements Node { @@ -11301,6 +11341,15 @@ type Viewer { before: CursorKey orderBy: OrganizationOrder ): OrganizationConnection! @goField(forceResolver: true) + + invitations( + first: Int + after: CursorKey + last: Int + before: CursorKey + orderBy: InvitationOrder + filter: InvitationFilter + ): InvitationConnection! @goField(forceResolver: true) } # Connection Types @@ -11400,13 +11449,19 @@ type TrustCenterReferenceEdge { node: TrustCenterReference! } -type UserConnection { +type UserConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.UserConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [UserEdge!]! pageInfo: PageInfo! } -type MembershipConnection { +type MembershipConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.MembershipConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [MembershipEdge!]! pageInfo: PageInfo! @@ -11715,7 +11770,10 @@ type File { updatedAt: Datetime! } -type InvitationConnection { +type InvitationConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.InvitationConnection" + ) { totalCount: Int! @goField(forceResolver: true) edges: [InvitationEdge!]! pageInfo: PageInfo! @@ -11786,6 +11844,7 @@ type Mutation { # User mutations confirmEmail(input: ConfirmEmailInput!): ConfirmEmailPayload! inviteUser(input: InviteUserInput!): InviteUserPayload! + acceptInvitation(input: AcceptInvitationInput!): AcceptInvitationPayload! deleteInvitation(input: DeleteInvitationInput!): DeleteInvitationPayload! removeMember(input: RemoveMemberInput!): RemoveMemberPayload! @@ -12546,6 +12605,10 @@ input InviteUserInput { createPeople: Boolean! } +input AcceptInvitationInput { + invitationId: ID! +} + input DeleteInvitationInput { invitationId: ID! } @@ -13072,12 +13135,16 @@ type InviteUserPayload { invitationEdge: InvitationEdge! } +type AcceptInvitationPayload { + invitation: Invitation! +} + type DeleteInvitationPayload { deletedInvitationId: ID! } type RemoveMemberPayload { - success: Boolean! + deletedMemberId: ID! } input VendorRiskAssessmentOrder { @@ -15186,6 +15253,29 @@ func (ec *executionContext) field_Measure_tasks_argsOrderBy( return zeroVal, nil } +func (ec *executionContext) field_Mutation_acceptInvitation_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { + var err error + args := map[string]any{} + arg0, err := ec.field_Mutation_acceptInvitation_argsInput(ctx, rawArgs) + if err != nil { + return nil, err + } + args["input"] = arg0 + return args, nil +} +func (ec *executionContext) field_Mutation_acceptInvitation_argsInput( + ctx context.Context, + rawArgs map[string]any, +) (types.AcceptInvitationInput, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("input")) + if tmp, ok := rawArgs["input"]; ok { + return ec.unmarshalNAcceptInvitationInput2githubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAcceptInvitationInput(ctx, tmp) + } + + var zeroVal types.AcceptInvitationInput + return zeroVal, nil +} + func (ec *executionContext) field_Mutation_assessVendor_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -18801,6 +18891,11 @@ func (ec *executionContext) field_Organization_invitations_args(ctx context.Cont return nil, err } args["orderBy"] = arg4 + arg5, err := ec.field_Organization_invitations_argsFilter(ctx, rawArgs) + if err != nil { + return nil, err + } + args["filter"] = arg5 return args, nil } func (ec *executionContext) field_Organization_invitations_argsFirst( @@ -18868,6 +18963,19 @@ func (ec *executionContext) field_Organization_invitations_argsOrderBy( return zeroVal, nil } +func (ec *executionContext) field_Organization_invitations_argsFilter( + ctx context.Context, + rawArgs map[string]any, +) (*types.InvitationFilter, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("filter")) + if tmp, ok := rawArgs["filter"]; ok { + return ec.unmarshalOInvitationFilter2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationFilter(ctx, tmp) + } + + var zeroVal *types.InvitationFilter + return zeroVal, nil +} + func (ec *executionContext) field_Organization_measures_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -20935,29 +21043,6 @@ func (ec *executionContext) field_TrustCenter_references_argsOrderBy( return zeroVal, nil } -func (ec *executionContext) field_User_people_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { - var err error - args := map[string]any{} - arg0, err := ec.field_User_people_argsOrganizationID(ctx, rawArgs) - if err != nil { - return nil, err - } - args["organizationId"] = arg0 - return args, nil -} -func (ec *executionContext) field_User_people_argsOrganizationID( - ctx context.Context, - rawArgs map[string]any, -) (gid.GID, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("organizationId")) - if tmp, ok := rawArgs["organizationId"]; ok { - return ec.unmarshalNID2githubᚗcomᚋgetproboᚋproboᚋpkgᚋgidᚐGID(ctx, tmp) - } - - var zeroVal gid.GID - return zeroVal, nil -} - func (ec *executionContext) field_Vendor_complianceReports_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -21338,6 +21423,119 @@ func (ec *executionContext) field_Vendor_services_argsOrderBy( return zeroVal, nil } +func (ec *executionContext) field_Viewer_invitations_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { + var err error + args := map[string]any{} + arg0, err := ec.field_Viewer_invitations_argsFirst(ctx, rawArgs) + if err != nil { + return nil, err + } + args["first"] = arg0 + arg1, err := ec.field_Viewer_invitations_argsAfter(ctx, rawArgs) + if err != nil { + return nil, err + } + args["after"] = arg1 + arg2, err := ec.field_Viewer_invitations_argsLast(ctx, rawArgs) + if err != nil { + return nil, err + } + args["last"] = arg2 + arg3, err := ec.field_Viewer_invitations_argsBefore(ctx, rawArgs) + if err != nil { + return nil, err + } + args["before"] = arg3 + arg4, err := ec.field_Viewer_invitations_argsOrderBy(ctx, rawArgs) + if err != nil { + return nil, err + } + args["orderBy"] = arg4 + arg5, err := ec.field_Viewer_invitations_argsFilter(ctx, rawArgs) + if err != nil { + return nil, err + } + args["filter"] = arg5 + return args, nil +} +func (ec *executionContext) field_Viewer_invitations_argsFirst( + ctx context.Context, + rawArgs map[string]any, +) (*int, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("first")) + if tmp, ok := rawArgs["first"]; ok { + return ec.unmarshalOInt2ᚖint(ctx, tmp) + } + + var zeroVal *int + return zeroVal, nil +} + +func (ec *executionContext) field_Viewer_invitations_argsAfter( + ctx context.Context, + rawArgs map[string]any, +) (*page.CursorKey, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("after")) + if tmp, ok := rawArgs["after"]; ok { + return ec.unmarshalOCursorKey2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋpageᚐCursorKey(ctx, tmp) + } + + var zeroVal *page.CursorKey + return zeroVal, nil +} + +func (ec *executionContext) field_Viewer_invitations_argsLast( + ctx context.Context, + rawArgs map[string]any, +) (*int, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("last")) + if tmp, ok := rawArgs["last"]; ok { + return ec.unmarshalOInt2ᚖint(ctx, tmp) + } + + var zeroVal *int + return zeroVal, nil +} + +func (ec *executionContext) field_Viewer_invitations_argsBefore( + ctx context.Context, + rawArgs map[string]any, +) (*page.CursorKey, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("before")) + if tmp, ok := rawArgs["before"]; ok { + return ec.unmarshalOCursorKey2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋpageᚐCursorKey(ctx, tmp) + } + + var zeroVal *page.CursorKey + return zeroVal, nil +} + +func (ec *executionContext) field_Viewer_invitations_argsOrderBy( + ctx context.Context, + rawArgs map[string]any, +) (*types.InvitationOrder, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("orderBy")) + if tmp, ok := rawArgs["orderBy"]; ok { + return ec.unmarshalOInvitationOrder2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationOrder(ctx, tmp) + } + + var zeroVal *types.InvitationOrder + return zeroVal, nil +} + +func (ec *executionContext) field_Viewer_invitations_argsFilter( + ctx context.Context, + rawArgs map[string]any, +) (*types.InvitationFilter, error) { + ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("filter")) + if tmp, ok := rawArgs["filter"]; ok { + return ec.unmarshalOInvitationFilter2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationFilter(ctx, tmp) + } + + var zeroVal *types.InvitationFilter + return zeroVal, nil +} + func (ec *executionContext) field_Viewer_organizations_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -21533,6 +21731,68 @@ func (ec *executionContext) field___Type_fields_argsIncludeDeprecated( // region **************************** field.gotpl ***************************** +func (ec *executionContext) _AcceptInvitationPayload_invitation(ctx context.Context, field graphql.CollectedField, obj *types.AcceptInvitationPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_AcceptInvitationPayload_invitation(ctx, field) + if err != nil { + return graphql.Null + } + ctx = graphql.WithFieldContext(ctx, fc) + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { + ctx = rctx // use context from middleware stack in children + return obj.Invitation, nil + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(*types.Invitation) + fc.Result = res + return ec.marshalNInvitation2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitation(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_AcceptInvitationPayload_invitation(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "AcceptInvitationPayload", + Field: field, + IsMethod: false, + IsResolver: false, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + switch field.Name { + case "id": + return ec.fieldContext_Invitation_id(ctx, field) + case "email": + return ec.fieldContext_Invitation_email(ctx, field) + case "fullName": + return ec.fieldContext_Invitation_fullName(ctx, field) + case "role": + return ec.fieldContext_Invitation_role(ctx, field) + case "expiresAt": + return ec.fieldContext_Invitation_expiresAt(ctx, field) + case "acceptedAt": + return ec.fieldContext_Invitation_acceptedAt(ctx, field) + case "createdAt": + return ec.fieldContext_Invitation_createdAt(ctx, field) + case "organization": + return ec.fieldContext_Invitation_organization(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type Invitation", field.Name) + }, + } + return fc, nil +} + func (ec *executionContext) _AssessVendorPayload_vendor(ctx context.Context, field graphql.CollectedField, obj *types.AssessVendorPayload) (ret graphql.Marshaler) { fc, err := ec.fieldContext_AssessVendorPayload_vendor(ctx, field) if err != nil { @@ -36289,6 +36549,114 @@ func (ec *executionContext) fieldContext_Invitation_createdAt(_ context.Context, return fc, nil } +func (ec *executionContext) _Invitation_organization(ctx context.Context, field graphql.CollectedField, obj *types.Invitation) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_Invitation_organization(ctx, field) + if err != nil { + return graphql.Null + } + ctx = graphql.WithFieldContext(ctx, fc) + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { + ctx = rctx // use context from middleware stack in children + return ec.resolvers.Invitation().Organization(rctx, obj) + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(*types.Organization) + fc.Result = res + return ec.marshalNOrganization2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐOrganization(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_Invitation_organization(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "Invitation", + Field: field, + IsMethod: true, + IsResolver: true, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + switch field.Name { + case "id": + return ec.fieldContext_Organization_id(ctx, field) + case "name": + return ec.fieldContext_Organization_name(ctx, field) + case "logoUrl": + return ec.fieldContext_Organization_logoUrl(ctx, field) + case "horizontalLogoUrl": + return ec.fieldContext_Organization_horizontalLogoUrl(ctx, field) + case "description": + return ec.fieldContext_Organization_description(ctx, field) + case "websiteUrl": + return ec.fieldContext_Organization_websiteUrl(ctx, field) + case "email": + return ec.fieldContext_Organization_email(ctx, field) + case "headquarterAddress": + return ec.fieldContext_Organization_headquarterAddress(ctx, field) + case "memberships": + return ec.fieldContext_Organization_memberships(ctx, field) + case "invitations": + return ec.fieldContext_Organization_invitations(ctx, field) + case "connectors": + return ec.fieldContext_Organization_connectors(ctx, field) + case "frameworks": + return ec.fieldContext_Organization_frameworks(ctx, field) + case "controls": + return ec.fieldContext_Organization_controls(ctx, field) + case "vendors": + return ec.fieldContext_Organization_vendors(ctx, field) + case "peoples": + return ec.fieldContext_Organization_peoples(ctx, field) + case "documents": + return ec.fieldContext_Organization_documents(ctx, field) + case "measures": + return ec.fieldContext_Organization_measures(ctx, field) + case "risks": + return ec.fieldContext_Organization_risks(ctx, field) + case "tasks": + return ec.fieldContext_Organization_tasks(ctx, field) + case "assets": + return ec.fieldContext_Organization_assets(ctx, field) + case "data": + return ec.fieldContext_Organization_data(ctx, field) + case "audits": + return ec.fieldContext_Organization_audits(ctx, field) + case "nonconformities": + return ec.fieldContext_Organization_nonconformities(ctx, field) + case "obligations": + return ec.fieldContext_Organization_obligations(ctx, field) + case "continualImprovements": + return ec.fieldContext_Organization_continualImprovements(ctx, field) + case "processingActivities": + return ec.fieldContext_Organization_processingActivities(ctx, field) + case "snapshots": + return ec.fieldContext_Organization_snapshots(ctx, field) + case "trustCenter": + return ec.fieldContext_Organization_trustCenter(ctx, field) + case "customDomain": + return ec.fieldContext_Organization_customDomain(ctx, field) + case "createdAt": + return ec.fieldContext_Organization_createdAt(ctx, field) + case "updatedAt": + return ec.fieldContext_Organization_updatedAt(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type Organization", field.Name) + }, + } + return fc, nil +} + func (ec *executionContext) _InvitationConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.InvitationConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_InvitationConnection_totalCount(ctx, field) if err != nil { @@ -36534,6 +36902,8 @@ func (ec *executionContext) fieldContext_InvitationEdge_node(_ context.Context, return ec.fieldContext_Invitation_acceptedAt(ctx, field) case "createdAt": return ec.fieldContext_Invitation_createdAt(ctx, field) + case "organization": + return ec.fieldContext_Invitation_organization(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type Invitation", field.Name) }, @@ -38902,6 +39272,65 @@ func (ec *executionContext) fieldContext_Mutation_inviteUser(ctx context.Context return fc, nil } +func (ec *executionContext) _Mutation_acceptInvitation(ctx context.Context, field graphql.CollectedField) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_Mutation_acceptInvitation(ctx, field) + if err != nil { + return graphql.Null + } + ctx = graphql.WithFieldContext(ctx, fc) + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { + ctx = rctx // use context from middleware stack in children + return ec.resolvers.Mutation().AcceptInvitation(rctx, fc.Args["input"].(types.AcceptInvitationInput)) + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(*types.AcceptInvitationPayload) + fc.Result = res + return ec.marshalNAcceptInvitationPayload2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAcceptInvitationPayload(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_Mutation_acceptInvitation(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "Mutation", + Field: field, + IsMethod: true, + IsResolver: true, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + switch field.Name { + case "invitation": + return ec.fieldContext_AcceptInvitationPayload_invitation(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type AcceptInvitationPayload", field.Name) + }, + } + defer func() { + if r := recover(); r != nil { + err = ec.Recover(ctx, r) + ec.Error(ctx, err) + } + }() + ctx = graphql.WithFieldContext(ctx, fc) + if fc.Args, err = ec.field_Mutation_acceptInvitation_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { + ec.Error(ctx, err) + return fc, err + } + return fc, nil +} + func (ec *executionContext) _Mutation_deleteInvitation(ctx context.Context, field graphql.CollectedField) (ret graphql.Marshaler) { fc, err := ec.fieldContext_Mutation_deleteInvitation(ctx, field) if err != nil { @@ -39000,8 +39429,8 @@ func (ec *executionContext) fieldContext_Mutation_removeMember(ctx context.Conte IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { - case "success": - return ec.fieldContext_RemoveMemberPayload_success(ctx, field) + case "deletedMemberId": + return ec.fieldContext_RemoveMemberPayload_deletedMemberId(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type RemoveMemberPayload", field.Name) }, @@ -47494,7 +47923,7 @@ func (ec *executionContext) _Organization_invitations(ctx context.Context, field }() resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { ctx = rctx // use context from middleware stack in children - return ec.resolvers.Organization().Invitations(rctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.InvitationOrder)) + return ec.resolvers.Organization().Invitations(rctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.InvitationOrder), fc.Args["filter"].(*types.InvitationFilter)) }) if err != nil { ec.Error(ctx, err) @@ -51410,6 +51839,8 @@ func (ec *executionContext) fieldContext_Query_viewer(_ context.Context, field g return ec.fieldContext_Viewer_user(ctx, field) case "organizations": return ec.fieldContext_Viewer_organizations(ctx, field) + case "invitations": + return ec.fieldContext_Viewer_invitations(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type Viewer", field.Name) }, @@ -51548,8 +51979,8 @@ func (ec *executionContext) fieldContext_Query___schema(_ context.Context, field return fc, nil } -func (ec *executionContext) _RemoveMemberPayload_success(ctx context.Context, field graphql.CollectedField, obj *types.RemoveMemberPayload) (ret graphql.Marshaler) { - fc, err := ec.fieldContext_RemoveMemberPayload_success(ctx, field) +func (ec *executionContext) _RemoveMemberPayload_deletedMemberId(ctx context.Context, field graphql.CollectedField, obj *types.RemoveMemberPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_RemoveMemberPayload_deletedMemberId(ctx, field) if err != nil { return graphql.Null } @@ -51562,7 +51993,7 @@ func (ec *executionContext) _RemoveMemberPayload_success(ctx context.Context, fi }() resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { ctx = rctx // use context from middleware stack in children - return obj.Success, nil + return obj.DeletedMemberID, nil }) if err != nil { ec.Error(ctx, err) @@ -51574,19 +52005,19 @@ func (ec *executionContext) _RemoveMemberPayload_success(ctx context.Context, fi } return graphql.Null } - res := resTmp.(bool) + res := resTmp.(gid.GID) fc.Result = res - return ec.marshalNBoolean2bool(ctx, field.Selections, res) + return ec.marshalNID2githubᚗcomᚋgetproboᚋproboᚋpkgᚋgidᚐGID(ctx, field.Selections, res) } -func (ec *executionContext) fieldContext_RemoveMemberPayload_success(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { +func (ec *executionContext) fieldContext_RemoveMemberPayload_deletedMemberId(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { fc = &graphql.FieldContext{ Object: "RemoveMemberPayload", 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 ID does not have child fields") }, } return fc, nil @@ -59994,80 +60425,6 @@ func (ec *executionContext) fieldContext_User_updatedAt(_ context.Context, field return fc, nil } -func (ec *executionContext) _User_people(ctx context.Context, field graphql.CollectedField, obj *types.User) (ret graphql.Marshaler) { - fc, err := ec.fieldContext_User_people(ctx, field) - if err != nil { - return graphql.Null - } - ctx = graphql.WithFieldContext(ctx, fc) - defer func() { - if r := recover(); r != nil { - ec.Error(ctx, ec.Recover(ctx, r)) - ret = graphql.Null - } - }() - resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { - ctx = rctx // use context from middleware stack in children - return ec.resolvers.User().People(rctx, obj, fc.Args["organizationId"].(gid.GID)) - }) - if err != nil { - ec.Error(ctx, err) - return graphql.Null - } - if resTmp == nil { - return graphql.Null - } - res := resTmp.(*types.People) - fc.Result = res - return ec.marshalOPeople2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐPeople(ctx, field.Selections, res) -} - -func (ec *executionContext) fieldContext_User_people(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { - fc = &graphql.FieldContext{ - Object: "User", - Field: field, - IsMethod: true, - IsResolver: true, - Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - switch field.Name { - case "id": - return ec.fieldContext_People_id(ctx, field) - case "fullName": - return ec.fieldContext_People_fullName(ctx, field) - case "primaryEmailAddress": - return ec.fieldContext_People_primaryEmailAddress(ctx, field) - case "additionalEmailAddresses": - return ec.fieldContext_People_additionalEmailAddresses(ctx, field) - case "kind": - return ec.fieldContext_People_kind(ctx, field) - case "position": - return ec.fieldContext_People_position(ctx, field) - case "contractStartDate": - return ec.fieldContext_People_contractStartDate(ctx, field) - case "contractEndDate": - return ec.fieldContext_People_contractEndDate(ctx, field) - case "createdAt": - return ec.fieldContext_People_createdAt(ctx, field) - case "updatedAt": - return ec.fieldContext_People_updatedAt(ctx, field) - } - return nil, fmt.Errorf("no field named %q was found under type People", field.Name) - }, - } - defer func() { - if r := recover(); r != nil { - err = ec.Recover(ctx, r) - ec.Error(ctx, err) - } - }() - ctx = graphql.WithFieldContext(ctx, fc) - if fc.Args, err = ec.field_User_people_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { - ec.Error(ctx, err) - return fc, err - } - return fc, nil -} - func (ec *executionContext) _UserConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.UserConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_UserConnection_totalCount(ctx, field) if err != nil { @@ -60309,8 +60666,6 @@ func (ec *executionContext) fieldContext_UserEdge_node(_ context.Context, field return ec.fieldContext_User_createdAt(ctx, field) case "updatedAt": return ec.fieldContext_User_updatedAt(ctx, field) - case "people": - return ec.fieldContext_User_people(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type User", field.Name) }, @@ -65549,8 +65904,6 @@ func (ec *executionContext) fieldContext_Viewer_user(_ context.Context, field gr return ec.fieldContext_User_createdAt(ctx, field) case "updatedAt": return ec.fieldContext_User_updatedAt(ctx, field) - case "people": - return ec.fieldContext_User_people(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type User", field.Name) }, @@ -65619,6 +65972,69 @@ func (ec *executionContext) fieldContext_Viewer_organizations(ctx context.Contex return fc, nil } +func (ec *executionContext) _Viewer_invitations(ctx context.Context, field graphql.CollectedField, obj *types.Viewer) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_Viewer_invitations(ctx, field) + if err != nil { + return graphql.Null + } + ctx = graphql.WithFieldContext(ctx, fc) + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { + ctx = rctx // use context from middleware stack in children + return ec.resolvers.Viewer().Invitations(rctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.InvitationOrder), fc.Args["filter"].(*types.InvitationFilter)) + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(*types.InvitationConnection) + fc.Result = res + return ec.marshalNInvitationConnection2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationConnection(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_Viewer_invitations(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "Viewer", + Field: field, + IsMethod: true, + IsResolver: true, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + switch field.Name { + case "totalCount": + return ec.fieldContext_InvitationConnection_totalCount(ctx, field) + case "edges": + return ec.fieldContext_InvitationConnection_edges(ctx, field) + case "pageInfo": + return ec.fieldContext_InvitationConnection_pageInfo(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type InvitationConnection", field.Name) + }, + } + defer func() { + if r := recover(); r != nil { + err = ec.Recover(ctx, r) + ec.Error(ctx, err) + } + }() + ctx = graphql.WithFieldContext(ctx, fc) + if fc.Args, err = ec.field_Viewer_invitations_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { + ec.Error(ctx, err) + return fc, err + } + return fc, nil +} + func (ec *executionContext) ___Directive_name(ctx context.Context, field graphql.CollectedField, obj *introspection.Directive) (ret graphql.Marshaler) { fc, err := ec.fieldContext___Directive_name(ctx, field) if err != nil { @@ -67570,6 +67986,33 @@ func (ec *executionContext) fieldContext___Type_isOneOf(_ context.Context, field // region **************************** input.gotpl ***************************** +func (ec *executionContext) unmarshalInputAcceptInvitationInput(ctx context.Context, obj any) (types.AcceptInvitationInput, error) { + var it types.AcceptInvitationInput + asMap := map[string]any{} + for k, v := range obj.(map[string]any) { + asMap[k] = v + } + + fieldsInOrder := [...]string{"invitationId"} + for _, k := range fieldsInOrder { + v, ok := asMap[k] + if !ok { + continue + } + switch k { + case "invitationId": + ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("invitationId")) + data, err := ec.unmarshalNID2githubᚗcomᚋgetproboᚋproboᚋpkgᚋgidᚐGID(ctx, v) + if err != nil { + return it, err + } + it.InvitationID = data + } + } + + return it, nil +} + func (ec *executionContext) unmarshalInputAssessVendorInput(ctx context.Context, obj any) (types.AssessVendorInput, error) { var it types.AssessVendorInput asMap := map[string]any{} @@ -71598,6 +72041,33 @@ func (ec *executionContext) unmarshalInputImportMeasureInput(ctx context.Context return it, nil } +func (ec *executionContext) unmarshalInputInvitationFilter(ctx context.Context, obj any) (types.InvitationFilter, error) { + var it types.InvitationFilter + asMap := map[string]any{} + for k, v := range obj.(map[string]any) { + asMap[k] = v + } + + fieldsInOrder := [...]string{"onlyPending"} + for _, k := range fieldsInOrder { + v, ok := asMap[k] + if !ok { + continue + } + switch k { + case "onlyPending": + ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("onlyPending")) + data, err := ec.unmarshalOBoolean2ᚖbool(ctx, v) + if err != nil { + return it, err + } + it.OnlyPending = data + } + } + + return it, nil +} + func (ec *executionContext) unmarshalInputInvitationOrder(ctx context.Context, obj any) (types.InvitationOrder, error) { var it types.InvitationOrder asMap := map[string]any{} @@ -74925,6 +75395,45 @@ func (ec *executionContext) _Node(ctx context.Context, sel ast.SelectionSet, obj // region **************************** object.gotpl **************************** +var acceptInvitationPayloadImplementors = []string{"AcceptInvitationPayload"} + +func (ec *executionContext) _AcceptInvitationPayload(ctx context.Context, sel ast.SelectionSet, obj *types.AcceptInvitationPayload) graphql.Marshaler { + fields := graphql.CollectFields(ec.OperationContext, sel, acceptInvitationPayloadImplementors) + + 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("AcceptInvitationPayload") + case "invitation": + out.Values[i] = ec._AcceptInvitationPayload_invitation(ctx, field, obj) + if out.Values[i] == graphql.Null { + out.Invalids++ + } + 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 assessVendorPayloadImplementors = []string{"AssessVendorPayload"} func (ec *executionContext) _AssessVendorPayload(ctx context.Context, sel ast.SelectionSet, obj *types.AssessVendorPayload) graphql.Marshaler { @@ -81730,35 +82239,71 @@ func (ec *executionContext) _Invitation(ctx context.Context, sel ast.SelectionSe case "id": out.Values[i] = ec._Invitation_id(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "email": out.Values[i] = ec._Invitation_email(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "fullName": out.Values[i] = ec._Invitation_fullName(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "role": out.Values[i] = ec._Invitation_role(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "expiresAt": out.Values[i] = ec._Invitation_expiresAt(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "acceptedAt": out.Values[i] = ec._Invitation_acceptedAt(ctx, field, obj) case "createdAt": out.Values[i] = ec._Invitation_createdAt(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } + case "organization": + field := field + + innerFunc := func(ctx context.Context, fs *graphql.FieldSet) (res graphql.Marshaler) { + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + } + }() + res = ec._Invitation_organization(ctx, field, obj) + if res == graphql.Null { + atomic.AddUint32(&fs.Invalids, 1) + } + return res + } + + if field.Deferrable != nil { + dfs, ok := deferred[field.Deferrable.Label] + di := 0 + if ok { + dfs.AddField(field) + di = len(dfs.Values) - 1 + } else { + dfs = graphql.NewFieldSet([]graphql.CollectedField{field}) + deferred[field.Deferrable.Label] = dfs + } + dfs.Concurrently(di, func(ctx context.Context) graphql.Marshaler { + return innerFunc(ctx, dfs) + }) + + // don't run the out.Concurrently() call below + out.Values[i] = graphql.Null + continue + } + + out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) default: panic("unknown field " + strconv.Quote(field.Name)) } @@ -82604,6 +83149,13 @@ func (ec *executionContext) _Mutation(ctx context.Context, sel ast.SelectionSet) if out.Values[i] == graphql.Null { out.Invalids++ } + case "acceptInvitation": + out.Values[i] = ec.OperationContext.RootResolverMiddleware(innerCtx, func(ctx context.Context) (res graphql.Marshaler) { + return ec._Mutation_acceptInvitation(ctx, field) + }) + if out.Values[i] == graphql.Null { + out.Invalids++ + } case "deleteInvitation": out.Values[i] = ec.OperationContext.RootResolverMiddleware(innerCtx, func(ctx context.Context) (res graphql.Marshaler) { return ec._Mutation_deleteInvitation(ctx, field) @@ -85554,8 +86106,8 @@ func (ec *executionContext) _RemoveMemberPayload(ctx context.Context, sel ast.Se switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("RemoveMemberPayload") - case "success": - out.Values[i] = ec._RemoveMemberPayload_success(ctx, field, obj) + case "deletedMemberId": + out.Values[i] = ec._RemoveMemberPayload_deletedMemberId(ctx, field, obj) if out.Values[i] == graphql.Null { out.Invalids++ } @@ -89108,61 +89660,28 @@ func (ec *executionContext) _User(ctx context.Context, sel ast.SelectionSet, obj case "id": out.Values[i] = ec._User_id(ctx, field, obj) if out.Values[i] == graphql.Null { - atomic.AddUint32(&out.Invalids, 1) + out.Invalids++ } case "fullName": out.Values[i] = ec._User_fullName(ctx, field, obj) if out.Values[i] == graphql.Null { - atomic.AddUint32(&out.Invalids, 1) + out.Invalids++ } case "email": out.Values[i] = ec._User_email(ctx, field, obj) if out.Values[i] == graphql.Null { - atomic.AddUint32(&out.Invalids, 1) + out.Invalids++ } case "createdAt": out.Values[i] = ec._User_createdAt(ctx, field, obj) if out.Values[i] == graphql.Null { - atomic.AddUint32(&out.Invalids, 1) + out.Invalids++ } case "updatedAt": out.Values[i] = ec._User_updatedAt(ctx, field, obj) if out.Values[i] == graphql.Null { - atomic.AddUint32(&out.Invalids, 1) + out.Invalids++ } - case "people": - field := field - - 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._User_people(ctx, field, obj) - return res - } - - if field.Deferrable != nil { - dfs, ok := deferred[field.Deferrable.Label] - di := 0 - if ok { - dfs.AddField(field) - di = len(dfs.Values) - 1 - } else { - dfs = graphql.NewFieldSet([]graphql.CollectedField{field}) - deferred[field.Deferrable.Label] = dfs - } - dfs.Concurrently(di, func(ctx context.Context) graphql.Marshaler { - return innerFunc(ctx, dfs) - }) - - // don't run the out.Concurrently() call below - out.Values[i] = graphql.Null - continue - } - - out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) default: panic("unknown field " + strconv.Quote(field.Name)) } @@ -90943,6 +91462,42 @@ func (ec *executionContext) _Viewer(ctx context.Context, sel ast.SelectionSet, o continue } + out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) + case "invitations": + field := field + + innerFunc := func(ctx context.Context, fs *graphql.FieldSet) (res graphql.Marshaler) { + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + } + }() + res = ec._Viewer_invitations(ctx, field, obj) + if res == graphql.Null { + atomic.AddUint32(&fs.Invalids, 1) + } + return res + } + + if field.Deferrable != nil { + dfs, ok := deferred[field.Deferrable.Label] + di := 0 + if ok { + dfs.AddField(field) + di = len(dfs.Values) - 1 + } else { + dfs = graphql.NewFieldSet([]graphql.CollectedField{field}) + deferred[field.Deferrable.Label] = dfs + } + dfs.Concurrently(di, func(ctx context.Context) graphql.Marshaler { + return innerFunc(ctx, dfs) + }) + + // don't run the out.Concurrently() call below + out.Values[i] = graphql.Null + continue + } + out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) default: panic("unknown field " + strconv.Quote(field.Name)) @@ -91302,6 +91857,25 @@ func (ec *executionContext) ___Type(ctx context.Context, sel ast.SelectionSet, o // region ***************************** type.gotpl ***************************** +func (ec *executionContext) unmarshalNAcceptInvitationInput2githubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAcceptInvitationInput(ctx context.Context, v any) (types.AcceptInvitationInput, error) { + res, err := ec.unmarshalInputAcceptInvitationInput(ctx, v) + return res, graphql.ErrorOnPath(ctx, err) +} + +func (ec *executionContext) marshalNAcceptInvitationPayload2githubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAcceptInvitationPayload(ctx context.Context, sel ast.SelectionSet, v types.AcceptInvitationPayload) graphql.Marshaler { + return ec._AcceptInvitationPayload(ctx, sel, &v) +} + +func (ec *executionContext) marshalNAcceptInvitationPayload2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAcceptInvitationPayload(ctx context.Context, sel ast.SelectionSet, v *types.AcceptInvitationPayload) graphql.Marshaler { + if v == nil { + if !graphql.HasFieldError(ctx, graphql.GetFieldContext(ctx)) { + ec.Errorf(ctx, "the requested element is null which the schema does not allow") + } + return graphql.Null + } + return ec._AcceptInvitationPayload(ctx, sel, v) +} + func (ec *executionContext) unmarshalNAssessVendorInput2githubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐAssessVendorInput(ctx context.Context, v any) (types.AssessVendorInput, error) { res, err := ec.unmarshalInputAssessVendorInput(ctx, v) return res, graphql.ErrorOnPath(ctx, err) @@ -100660,6 +101234,14 @@ func (ec *executionContext) marshalOInt2ᚖint(ctx context.Context, sel ast.Sele return res } +func (ec *executionContext) unmarshalOInvitationFilter2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationFilter(ctx context.Context, v any) (*types.InvitationFilter, error) { + if v == nil { + return nil, nil + } + res, err := ec.unmarshalInputInvitationFilter(ctx, v) + return &res, graphql.ErrorOnPath(ctx, err) +} + func (ec *executionContext) unmarshalOInvitationOrder2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationOrder(ctx context.Context, v any) (*types.InvitationOrder, error) { if v == nil { return nil, nil diff --git a/pkg/server/api/console/v1/signup_from_invitation_handler.go b/pkg/server/api/console/v1/signup_from_invitation_handler.go new file mode 100644 index 000000000..9d1c11301 --- /dev/null +++ b/pkg/server/api/console/v1/signup_from_invitation_handler.go @@ -0,0 +1,75 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package console_v1 + +import ( + "encoding/json" + "fmt" + "net/http" + + "github.com/getprobo/probo/pkg/auth" + "github.com/getprobo/probo/pkg/securecookie" + "go.gearno.de/kit/httpserver" +) + +type ( + SignupFromInvitationRequest struct { + Token string `json:"token"` + Password string `json:"password"` + FullName string `json:"fullName"` + } + + SignupFromInvitationResponse struct { + } +) + +func SignupFromInvitationHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var req SignupFromInvitationRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("cannot decode body: %w", err)) + return + } + + user, session, err := authSvc.SignupFromInvitation(r.Context(), req.Token, req.Password, req.FullName) + if err != nil { + httpserver.RenderError(w, http.StatusBadRequest, err) + return + } + + securecookie.Set( + w, + securecookie.DefaultConfig( + authCfg.CookieName, + authCfg.CookieSecret, + ), + session.ID.String(), + ) + + httpserver.RenderJSON( + w, + http.StatusOK, + SignUpResponse{ + User: UserResponse{ + ID: user.ID, + Email: user.EmailAddress, + FullName: user.FullName, + CreatedAt: user.CreatedAt, + UpdatedAt: user.UpdatedAt, + }, + }, + ) + } +} diff --git a/pkg/server/api/console/v1/types/invitation.go b/pkg/server/api/console/v1/types/invitation.go index d2954c24f..3f355619e 100644 --- a/pkg/server/api/console/v1/types/invitation.go +++ b/pkg/server/api/console/v1/types/invitation.go @@ -16,10 +16,28 @@ package types import ( "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" "github.com/getprobo/probo/pkg/page" ) -func NewInvitationConnection(p *page.Page[*coredata.Invitation, coredata.InvitationOrderField]) *InvitationConnection { +type ( + InvitationConnection struct { + TotalCount int `json:"totalCount"` + Edges []*InvitationEdge `json:"edges"` + PageInfo *PageInfo `json:"pageInfo"` + + Resolver any + ParentID gid.GID + Filter *InvitationFilter + } +) + +func NewInvitationConnection( + p *page.Page[*coredata.Invitation, coredata.InvitationOrderField], + resolver any, + parentID gid.GID, + filter *InvitationFilter, +) *InvitationConnection { var edges = make([]*InvitationEdge, len(p.Data)) for i := range edges { @@ -29,6 +47,9 @@ func NewInvitationConnection(p *page.Page[*coredata.Invitation, coredata.Invitat return &InvitationConnection{ Edges: edges, PageInfo: NewPageInfo(p), + Resolver: resolver, + ParentID: parentID, + Filter: filter, } } diff --git a/pkg/server/api/console/v1/types/membership.go b/pkg/server/api/console/v1/types/membership.go index 9d512890e..09afc01f6 100644 --- a/pkg/server/api/console/v1/types/membership.go +++ b/pkg/server/api/console/v1/types/membership.go @@ -16,14 +16,28 @@ package types import ( "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" "github.com/getprobo/probo/pkg/page" ) type ( + MembershipConnection struct { + TotalCount int `json:"totalCount"` + Edges []*MembershipEdge `json:"edges"` + PageInfo *PageInfo `json:"pageInfo"` + + Resolver any + ParentID gid.GID + } + MembershipOrderBy OrderBy[coredata.MembershipOrderField] ) -func NewMembershipConnection(p *page.Page[*coredata.Membership, coredata.MembershipOrderField]) *MembershipConnection { +func NewMembershipConnection( + p *page.Page[*coredata.Membership, coredata.MembershipOrderField], + resolver any, + parentID gid.GID, +) *MembershipConnection { var edges = make([]*MembershipEdge, len(p.Data)) for i := range edges { @@ -33,6 +47,8 @@ func NewMembershipConnection(p *page.Page[*coredata.Membership, coredata.Members return &MembershipConnection{ Edges: edges, PageInfo: NewPageInfo(p), + Resolver: resolver, + ParentID: parentID, } } diff --git a/pkg/server/api/console/v1/types/types.go b/pkg/server/api/console/v1/types/types.go index c4072dd2a..a5bf219fe 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -16,6 +16,14 @@ type Node interface { GetID() gid.GID } +type AcceptInvitationInput struct { + InvitationID gid.GID `json:"invitationId"` +} + +type AcceptInvitationPayload struct { + Invitation *Invitation `json:"invitation"` +} + type AssessVendorInput struct { ID gid.GID `json:"id"` WebsiteURL string `json:"websiteUrl"` @@ -1200,29 +1208,28 @@ type ImportMeasurePayload struct { } type Invitation struct { - ID gid.GID `json:"id"` - Email string `json:"email"` - FullName string `json:"fullName"` - Role string `json:"role"` - ExpiresAt time.Time `json:"expiresAt"` - AcceptedAt *time.Time `json:"acceptedAt,omitempty"` - CreatedAt time.Time `json:"createdAt"` + ID gid.GID `json:"id"` + Email string `json:"email"` + FullName string `json:"fullName"` + Role string `json:"role"` + ExpiresAt time.Time `json:"expiresAt"` + AcceptedAt *time.Time `json:"acceptedAt,omitempty"` + CreatedAt time.Time `json:"createdAt"` + Organization *Organization `json:"organization"` } func (Invitation) IsNode() {} func (this Invitation) GetID() gid.GID { return this.ID } -type InvitationConnection struct { - TotalCount int `json:"totalCount"` - Edges []*InvitationEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type InvitationEdge struct { Cursor page.CursorKey `json:"cursor"` Node *Invitation `json:"node"` } +type InvitationFilter struct { + OnlyPending *bool `json:"onlyPending,omitempty"` +} + type InvitationOrder struct { Direction page.OrderDirection `json:"direction"` Field coredata.InvitationOrderField `json:"field"` @@ -1280,12 +1287,6 @@ type Membership struct { func (Membership) IsNode() {} func (this Membership) GetID() gid.GID { return this.ID } -type MembershipConnection struct { - TotalCount int `json:"totalCount"` - Edges []*MembershipEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type MembershipEdge struct { Cursor page.CursorKey `json:"cursor"` Node *Membership `json:"node"` @@ -1493,7 +1494,7 @@ type RemoveMemberInput struct { } type RemoveMemberPayload struct { - Success bool `json:"success"` + DeletedMemberID gid.GID `json:"deletedMemberId"` } type Report struct { @@ -2120,18 +2121,11 @@ type User struct { Email string `json:"email"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt"` - People *People `json:"people,omitempty"` } func (User) IsNode() {} func (this User) GetID() gid.GID { return this.ID } -type UserConnection struct { - TotalCount int `json:"totalCount"` - Edges []*UserEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type UserEdge struct { Cursor page.CursorKey `json:"cursor"` Node *User `json:"node"` @@ -2316,4 +2310,5 @@ type Viewer struct { ID gid.GID `json:"id"` User *User `json:"user"` Organizations *OrganizationConnection `json:"organizations"` + Invitations *InvitationConnection `json:"invitations"` } diff --git a/pkg/server/api/console/v1/types/user.go b/pkg/server/api/console/v1/types/user.go index 36130a221..04647ef2f 100644 --- a/pkg/server/api/console/v1/types/user.go +++ b/pkg/server/api/console/v1/types/user.go @@ -16,14 +16,28 @@ package types import ( "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" "github.com/getprobo/probo/pkg/page" ) type ( + UserConnection struct { + TotalCount int `json:"totalCount"` + Edges []*UserEdge `json:"edges"` + PageInfo *PageInfo `json:"pageInfo"` + + Resolver any + ParentID gid.GID + } + UserOrderBy OrderBy[coredata.UserOrderField] ) -func NewUserConnection(p *page.Page[*coredata.User, coredata.UserOrderField]) *UserConnection { +func NewUserConnection( + p *page.Page[*coredata.User, coredata.UserOrderField], + resolver any, + parentID gid.GID, +) *UserConnection { var edges = make([]*UserEdge, len(p.Data)) for i := range edges { @@ -33,6 +47,8 @@ func NewUserConnection(p *page.Page[*coredata.User, coredata.UserOrderField]) *U return &UserConnection{ Edges: edges, PageInfo: NewPageInfo(p), + Resolver: resolver, + ParentID: parentID, } } diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 5be863675..5c52f9014 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -891,25 +891,45 @@ func (r *frameworkConnectionResolver) TotalCount(ctx context.Context, obj *types panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) } +// Organization is the resolver for the organization field. +func (r *invitationResolver) Organization(ctx context.Context, obj *types.Invitation) (*types.Organization, error) { + organization, err := r.authzSvc.GetOrganizationByInvitationID(ctx, obj.ID) + if err != nil { + panic(fmt.Errorf("cannot load organization: %w", err)) + } + + return types.NewOrganization(organization), nil +} + // TotalCount is the resolver for the totalCount field. func (r *invitationConnectionResolver) TotalCount(ctx context.Context, obj *types.InvitationConnection) (int, error) { - currentUser := UserFromContext(ctx) - if currentUser == nil { - return 0, fmt.Errorf("no authenticated user") + switch obj.Resolver.(type) { + case *organizationResolver: + authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID()) + count, err := authzSvc.CountOrganizationInvitations(ctx, obj.ParentID) + if err != nil { + panic(fmt.Errorf("failed to count organization invitations: %w", err)) + } + return count, nil + case *viewerResolver: + user := UserFromContext(ctx) + if user == nil { + panic(fmt.Errorf("no authenticated user")) + } + + invitationFilter := coredata.NewInvitationFilter(nil) + if obj.Filter != nil { + invitationFilter = coredata.NewInvitationFilter(obj.Filter.OnlyPending) + } + + count, err := r.authzSvc.CountUserInvitations(ctx, user.EmailAddress, invitationFilter) + if err != nil { + panic(fmt.Errorf("failed to count user invitations: %w", err)) + } + return count, nil } - memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID) - if err != nil || len(memberships) == 0 { - return 0, fmt.Errorf("user has no organization memberships") - } - - orgID := memberships[0].ID - count, err := r.authzSvc.CountOrganizationInvitations(ctx, orgID) - if err != nil { - return 0, fmt.Errorf("failed to count invitations: %w", err) - } - - return count, nil + panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) } // Evidences is the resolver for the evidences field. @@ -1052,23 +1072,17 @@ func (r *measureConnectionResolver) TotalCount(ctx context.Context, obj *types.M // TotalCount is the resolver for the totalCount field. func (r *membershipConnectionResolver) TotalCount(ctx context.Context, obj *types.MembershipConnection) (int, error) { - currentUser := UserFromContext(ctx) - if currentUser == nil { - return 0, fmt.Errorf("no authenticated user") + switch obj.Resolver.(type) { + case *organizationResolver: + authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID()) + count, err := authzSvc.CountOrganizationMemberships(ctx, obj.ParentID) + if err != nil { + panic(fmt.Errorf("failed to count organization memberships: %w", err)) + } + return count, nil + default: + panic(fmt.Errorf("unknown resolver type for membership connection")) } - - memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID) - if err != nil || len(memberships) == 0 { - return 0, fmt.Errorf("user has no organization memberships") - } - - orgID := memberships[0].ID - count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID) - if err != nil { - return 0, fmt.Errorf("failed to count memberships: %w", err) - } - - return count, nil } // CreateOrganization is the resolver for the createOrganization field. @@ -1104,7 +1118,6 @@ func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.C ctx, probo.CreatePeopleRequest{ OrganizationID: organization.ID, - UserID: &UserFromContext(ctx).ID, FullName: UserFromContext(ctx).FullName, PrimaryEmailAddress: UserFromContext(ctx).EmailAddress, AdditionalEmailAddresses: []string{}, @@ -1379,48 +1392,49 @@ func (r *mutationResolver) ConfirmEmail(ctx context.Context, input types.Confirm // InviteUser is the resolver for the inviteUser field. func (r *mutationResolver) InviteUser(ctx context.Context, input types.InviteUserInput) (*types.InviteUserPayload, error) { - user := UserFromContext(ctx) - - organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID) + authzSvc := r.AuthzService(ctx, input.OrganizationID.TenantID()) + invitation, err := authzSvc.InviteUserToOrganization(ctx, input.OrganizationID, input.Email, input.FullName, string(authz.RoleMember)) if err != nil { - panic(fmt.Errorf("failed to list organizations for user: %w", err)) + panic(fmt.Errorf("failed to invite user to organization: %w", err)) } - for _, organization := range organizations { - if organization.ID == input.OrganizationID { - invitation, err := r.authzSvc.InviteUserToOrganization(ctx, input.OrganizationID, input.Email, input.FullName, string(authz.RoleMember)) - if err != nil { - return nil, err - } - - if input.CreatePeople { - prb := r.ProboService(ctx, input.OrganizationID.TenantID()) - _, err := prb.Peoples.Create(ctx, probo.CreatePeopleRequest{ - OrganizationID: input.OrganizationID, - FullName: input.FullName, - PrimaryEmailAddress: input.Email, - AdditionalEmailAddresses: []string{}, - Kind: coredata.PeopleKindEmployee, - }) - if err != nil { - return nil, fmt.Errorf("failed to create people record: %w", err) - } - } - - return &types.InviteUserPayload{ - InvitationEdge: types.NewInvitationEdge(invitation, coredata.InvitationOrderFieldCreatedAt), - }, nil + if input.CreatePeople { + prb := r.ProboService(ctx, input.OrganizationID.TenantID()) + _, err := prb.Peoples.Create(ctx, probo.CreatePeopleRequest{ + OrganizationID: input.OrganizationID, + FullName: input.FullName, + PrimaryEmailAddress: input.Email, + AdditionalEmailAddresses: []string{}, + Kind: coredata.PeopleKindEmployee, + }) + if err != nil { + return nil, fmt.Errorf("failed to create people record: %w", err) } } - return nil, fmt.Errorf("organization not found") + return &types.InviteUserPayload{ + InvitationEdge: types.NewInvitationEdge(invitation, coredata.InvitationOrderFieldCreatedAt), + }, nil +} + +// AcceptInvitation is the resolver for the acceptInvitation field. +func (r *mutationResolver) AcceptInvitation(ctx context.Context, input types.AcceptInvitationInput) (*types.AcceptInvitationPayload, error) { + user := UserFromContext(ctx) + + invitation, err := r.authzSvc.AcceptInvitationByID(ctx, input.InvitationID, user.ID) + if err != nil { + panic(fmt.Errorf("failed to accept invitation: %w", err)) + } + + return &types.AcceptInvitationPayload{Invitation: types.NewInvitation(invitation)}, nil } // DeleteInvitation is the resolver for the deleteInvitation field. func (r *mutationResolver) DeleteInvitation(ctx context.Context, input types.DeleteInvitationInput) (*types.DeleteInvitationPayload, error) { - err := r.authzSvc.DeleteInvitation(ctx, input.InvitationID) + authzSvc := r.AuthzService(ctx, input.InvitationID.TenantID()) + err := authzSvc.DeleteInvitation(ctx, input.InvitationID) if err != nil { - return nil, err + panic(fmt.Errorf("failed to delete invitation: %w", err)) } return &types.DeleteInvitationPayload{ @@ -1430,25 +1444,13 @@ func (r *mutationResolver) DeleteInvitation(ctx context.Context, input types.Del // RemoveMember is the resolver for the removeMember field. func (r *mutationResolver) RemoveMember(ctx context.Context, input types.RemoveMemberInput) (*types.RemoveMemberPayload, error) { - user := UserFromContext(ctx) - - organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID) + authzSvc := r.AuthzService(ctx, input.OrganizationID.TenantID()) + err := authzSvc.RemoveMemberFromOrganization(ctx, input.OrganizationID, input.MemberID) if err != nil { - panic(fmt.Errorf("failed to list organizations for user: %w", err)) + return nil, err } - for _, organization := range organizations { - if organization.ID == input.OrganizationID { - err := r.authzSvc.RemoveMemberFromOrganization(ctx, input.OrganizationID, input.MemberID) - if err != nil { - return nil, err - } - - return &types.RemoveMemberPayload{Success: true}, nil - } - } - - return nil, fmt.Errorf("organization not found") + return &types.RemoveMemberPayload{DeletedMemberID: input.MemberID}, nil } // CreatePeople is the resolver for the createPeople field. @@ -3611,16 +3613,17 @@ func (r *organizationResolver) Memberships(ctx context.Context, obj *types.Organ cursor := types.NewCursor(first, after, last, before, pageOrderBy) - page, err := r.authzSvc.GetAllOrganizationMemberships(ctx, obj.ID, cursor) + authzSvc := r.AuthzService(ctx, obj.ID.TenantID()) + page, err := authzSvc.GetMembershipsByOrganizationID(ctx, obj.ID, cursor) if err != nil { panic(fmt.Errorf("cannot list memberships: %w", err)) } - return types.NewMembershipConnection(page), nil + return types.NewMembershipConnection(page, r, obj.ID), nil } // Invitations is the resolver for the invitations field. -func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder) (*types.InvitationConnection, error) { +func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) { pageOrderBy := page.OrderBy[coredata.InvitationOrderField]{ Field: coredata.InvitationOrderFieldCreatedAt, Direction: page.OrderDirectionDesc, @@ -3634,12 +3637,13 @@ func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organ cursor := types.NewCursor(first, after, last, before, pageOrderBy) - page, err := r.authzSvc.GetAllOrganizationInvitations(ctx, obj.ID, cursor) + authzSvc := r.AuthzService(ctx, obj.ID.TenantID()) + page, err := authzSvc.GetInvitationsByOrganizationID(ctx, obj.ID, cursor) if err != nil { panic(fmt.Errorf("cannot list invitations: %w", err)) } - return types.NewInvitationConnection(page), nil + return types.NewInvitationConnection(page, r, obj.ID, filter), nil } // Connectors is the resolver for the connectors field. @@ -4943,41 +4947,19 @@ func (r *trustCenterReferenceConnectionResolver) TotalCount(ctx context.Context, return count, nil } -// People is the resolver for the people field. -func (r *userResolver) People(ctx context.Context, obj *types.User, organizationID gid.GID) (*types.People, error) { - prb := r.ProboService(ctx, organizationID.TenantID()) - - people, err := prb.Peoples.GetByUserID(ctx, obj.ID) - if err != nil { - var errPeopleNotFound *coredata.ErrPeopleNotFound - if errors.As(err, &errPeopleNotFound) { - return nil, nil - } - panic(fmt.Errorf("failed to get people: %w", err)) - } - - return types.NewPeople(people), nil -} - // TotalCount is the resolver for the totalCount field. func (r *userConnectionResolver) TotalCount(ctx context.Context, obj *types.UserConnection) (int, error) { - currentUser := UserFromContext(ctx) - if currentUser == nil { - return 0, fmt.Errorf("no authenticated user") + switch obj.Resolver.(type) { + case *organizationResolver: + authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID()) + count, err := authzSvc.CountOrganizationUsers(ctx, obj.ParentID) + if err != nil { + panic(fmt.Errorf("failed to count organization users: %w", err)) + } + return count, nil + default: + panic(fmt.Errorf("unknown resolver type for user connection")) } - - memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID) - if err != nil || len(memberships) == 0 { - return 0, fmt.Errorf("user has no organization memberships") - } - - orgID := memberships[0].ID - count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID) - if err != nil { - return 0, fmt.Errorf("failed to count memberships: %w", err) - } - - return count, nil } // Organization is the resolver for the organization field. @@ -5353,6 +5335,35 @@ func (r *viewerResolver) Organizations(ctx context.Context, obj *types.Viewer, f return types.NewOrganizationConnection(page), nil } +// Invitations is the resolver for the invitations field. +func (r *viewerResolver) Invitations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) { + user := UserFromContext(ctx) + + pageOrderBy := page.OrderBy[coredata.InvitationOrderField]{ + Field: coredata.InvitationOrderFieldCreatedAt, + Direction: page.OrderDirectionDesc, + } + if orderBy != nil { + pageOrderBy = page.OrderBy[coredata.InvitationOrderField]{ + Field: orderBy.Field, + Direction: orderBy.Direction, + } + } + cursor := types.NewCursor(first, after, last, before, pageOrderBy) + + invitationFilter := coredata.NewInvitationFilter(nil) + if filter != nil { + invitationFilter = coredata.NewInvitationFilter(filter.OnlyPending) + } + + invitations, err := r.authzSvc.GetUserInvitations(ctx, user.EmailAddress, cursor, invitationFilter) + if err != nil { + panic(fmt.Errorf("failed to list invitations for user: %w", err)) + } + + return types.NewInvitationConnection(invitations, r, gid.GID{}, filter), nil +} + // Asset returns schema.AssetResolver implementation. func (r *Resolver) Asset() schema.AssetResolver { return &assetResolver{r} } @@ -5432,6 +5443,9 @@ func (r *Resolver) FrameworkConnection() schema.FrameworkConnectionResolver { return &frameworkConnectionResolver{r} } +// Invitation returns schema.InvitationResolver implementation. +func (r *Resolver) Invitation() schema.InvitationResolver { return &invitationResolver{r} } + // InvitationConnection returns schema.InvitationConnectionResolver implementation. func (r *Resolver) InvitationConnection() schema.InvitationConnectionResolver { return &invitationConnectionResolver{r} @@ -5541,9 +5555,6 @@ func (r *Resolver) TrustCenterReferenceConnection() schema.TrustCenterReferenceC return &trustCenterReferenceConnectionResolver{r} } -// User returns schema.UserResolver implementation. -func (r *Resolver) User() schema.UserResolver { return &userResolver{r} } - // UserConnection returns schema.UserConnectionResolver implementation. func (r *Resolver) UserConnection() schema.UserConnectionResolver { return &userConnectionResolver{r} } @@ -5603,6 +5614,7 @@ type evidenceConnectionResolver struct{ *Resolver } type fileResolver struct{ *Resolver } type frameworkResolver struct{ *Resolver } type frameworkConnectionResolver struct{ *Resolver } +type invitationResolver struct{ *Resolver } type invitationConnectionResolver struct{ *Resolver } type measureResolver struct{ *Resolver } type measureConnectionResolver struct{ *Resolver } @@ -5630,7 +5642,6 @@ type trustCenterDocumentAccessResolver struct{ *Resolver } type trustCenterDocumentAccessConnectionResolver struct{ *Resolver } type trustCenterReferenceResolver struct{ *Resolver } type trustCenterReferenceConnectionResolver struct{ *Resolver } -type userResolver struct{ *Resolver } type userConnectionResolver struct{ *Resolver } type vendorResolver struct{ *Resolver } type vendorBusinessAssociateAgreementResolver struct{ *Resolver }