Require explicit portal access request IDs

Drop the request-all shortcut so callers always name the
documents, reports, and files to request. TopBar Get Access
now only signs in; bulk selection is the multi-resource path.

Signed-off-by: Émile Ré <emile@probo.com>
This commit is contained in:
Émile Ré
2026-07-27 16:18:30 +02:00
parent baa5a3588a
commit d0c9327e99
10 changed files with 129 additions and 211 deletions

View File

@@ -27,7 +27,7 @@ import { useTranslation } from "react-i18next";
import { graphql, useFragment } from "react-relay"; import { graphql, useFragment } from "react-relay";
import { Link as RouterLink, useLocation } from "react-router"; import { Link as RouterLink, useLocation } from "react-router";
import { buildRequestAllContinueUrl, redirectToInitiate } from "#/lib/auth/continueUrl"; import { getSafeContinueUrl, redirectToInitiate } from "#/lib/auth/continueUrl";
import { useLocalizedPath } from "#/lib/i18n/useLocale"; import { useLocalizedPath } from "#/lib/i18n/useLocale";
import type { TopBar_query$key } from "./__generated__/TopBar_query.graphql"; import type { TopBar_query$key } from "./__generated__/TopBar_query.graphql";
@@ -121,7 +121,7 @@ export function TopBar({ queryKey }: TopBarProps) {
highContrast highContrast
iconStart={<LockSimpleIcon />} iconStart={<LockSimpleIcon />}
onClick={() => { onClick={() => {
redirectToInitiate(buildRequestAllContinueUrl()); redirectToInitiate(getSafeContinueUrl(window.location.href));
}} }}
> >
{t("topBar.getAccess")} {t("topBar.getAccess")}

View File

@@ -43,7 +43,7 @@ import { useTranslation } from "react-i18next";
import { graphql, useFragment } from "react-relay"; import { graphql, useFragment } from "react-relay";
import { useLocation } from "react-router"; import { useLocation } from "react-router";
import { buildRequestAllContinueUrl, redirectToInitiate } from "#/lib/auth/continueUrl"; import { getSafeContinueUrl, redirectToInitiate } from "#/lib/auth/continueUrl";
import { useSignOut } from "#/lib/auth/useSignOut"; import { useSignOut } from "#/lib/auth/useSignOut";
import { useLocalizedPath } from "#/lib/i18n/useLocale"; import { useLocalizedPath } from "#/lib/i18n/useLocale";
import { useSubscribeDialog } from "#/lib/mailingList/subscribeDialogContext"; import { useSubscribeDialog } from "#/lib/mailingList/subscribeDialogContext";
@@ -164,7 +164,7 @@ export function TopBarMobileNav({ identityKey }: TopBarMobileNavProps) {
iconStart={<LockSimpleIcon />} iconStart={<LockSimpleIcon />}
onClick={() => { onClick={() => {
close(); close();
redirectToInitiate(buildRequestAllContinueUrl()); redirectToInitiate(getSafeContinueUrl(window.location.href));
}} }}
> >
{t("topBar.getAccess")} {t("topBar.getAccess")}

View File

@@ -23,10 +23,9 @@ import { FullNameRequiredError, NDASignatureRequiredError } from "@probo/relay";
import { localizedPath, resolveUrlLocale, type UrlLocale } from "#/lib/i18n/locale"; import { localizedPath, resolveUrlLocale, type UrlLocale } from "#/lib/i18n/locale";
// Markers appended to a post-auth `continue` URL so the portal fires the pending // Markers appended to a post-auth `continue` URL so the portal fires the pending
// "request access" mutation once the user lands back authenticated. `request-all` // "request access" mutation once the user lands back authenticated. The
// covers the top-bar "Get Access"; the per-resource markers carry the id of a // per-resource markers carry the id of a single document / report / file whose
// single document / report / file whose access was requested from a locked row. // access was requested from a locked row.
export const REQUEST_ALL_PARAM = "request-all";
export const REQUEST_DOCUMENT_PARAM = "request-document-id"; export const REQUEST_DOCUMENT_PARAM = "request-document-id";
export const REQUEST_REPORT_PARAM = "request-report-id"; export const REQUEST_REPORT_PARAM = "request-report-id";
export const REQUEST_FILE_PARAM = "request-file-id"; export const REQUEST_FILE_PARAM = "request-file-id";
@@ -59,14 +58,6 @@ export function getSafeContinueUrl(param: string | null | undefined): string {
return fallback; return fallback;
} }
// Absolute URL of the current page with the request-all marker set, used as the
// `continue` target so the access request resumes after sign-in.
export function buildRequestAllContinueUrl(): string {
const url = new URL(window.location.href);
url.searchParams.set(REQUEST_ALL_PARAM, "true");
return url.toString();
}
// Absolute URL of the current page with a per-resource marker set, so a single // Absolute URL of the current page with a per-resource marker set, so a single
// document / report / file access request resumes after sign-in. // document / report / file access request resumes after sign-in.
export function buildRequestAccessContinueUrl(param: string, id: string): string { export function buildRequestAccessContinueUrl(param: string, id: string): string {

View File

@@ -27,9 +27,7 @@ import { graphql } from "relay-runtime";
import { import {
buildRequestAccessContinueUrl, buildRequestAccessContinueUrl,
buildRequestAllContinueUrl,
gateRedirectPath, gateRedirectPath,
REQUEST_ALL_PARAM,
REQUEST_DOCUMENT_PARAM, REQUEST_DOCUMENT_PARAM,
REQUEST_FILE_PARAM, REQUEST_FILE_PARAM,
REQUEST_REPORT_PARAM, REQUEST_REPORT_PARAM,
@@ -40,17 +38,6 @@ import { useMutation } from "#/lib/relay/useMutation";
import type { useResumeAccessRequest_documentMutation } from "./__generated__/useResumeAccessRequest_documentMutation.graphql"; import type { useResumeAccessRequest_documentMutation } from "./__generated__/useResumeAccessRequest_documentMutation.graphql";
import type { useResumeAccessRequest_fileMutation } from "./__generated__/useResumeAccessRequest_fileMutation.graphql"; import type { useResumeAccessRequest_fileMutation } from "./__generated__/useResumeAccessRequest_fileMutation.graphql";
import type { useResumeAccessRequest_reportMutation } from "./__generated__/useResumeAccessRequest_reportMutation.graphql"; import type { useResumeAccessRequest_reportMutation } from "./__generated__/useResumeAccessRequest_reportMutation.graphql";
import type { useResumeAccessRequestMutation } from "./__generated__/useResumeAccessRequestMutation.graphql";
const requestAllAccessesMutation = graphql`
mutation useResumeAccessRequestMutation {
requestAllAccesses {
compliancePortalAccess {
id
}
}
}
`;
const requestDocumentMutation = graphql` const requestDocumentMutation = graphql`
mutation useResumeAccessRequest_documentMutation($input: RequestDocumentAccessInput!) { mutation useResumeAccessRequest_documentMutation($input: RequestDocumentAccessInput!) {
@@ -99,10 +86,9 @@ const requestFileMutation = graphql`
// After a user signs in through OAuth /initiate, they land back on the page that // After a user signs in through OAuth /initiate, they land back on the page that
// carried a deferred access marker. This hook fires the matching mutation once // carried a deferred access marker. This hook fires the matching mutation once
// (when authenticated) — request-all from the top bar, or a single // (when authenticated) — a single document / report / file requested from a
// document / report / file requested from a locked row — routes to the // locked row — routes to the full-name gate when the backend asks for it, and
// full-name gate when the backend asks for it, and clears the marker so a // clears the marker so a refresh never re-triggers it.
// refresh never re-triggers it.
export function useResumeAccessRequest(isAuthenticated: boolean) { export function useResumeAccessRequest(isAuthenticated: boolean) {
const [searchParams, setSearchParams] = useSearchParams(); const [searchParams, setSearchParams] = useSearchParams();
const navigate = useNavigate(); const navigate = useNavigate();
@@ -111,10 +97,6 @@ export function useResumeAccessRequest(isAuthenticated: boolean) {
const { t } = useTranslation(); const { t } = useTranslation();
const firedRef = useRef(false); const firedRef = useRef(false);
const [requestAllAccesses] = useMutation<useResumeAccessRequestMutation>(
requestAllAccessesMutation,
{ errorToast: false },
);
const [requestDocumentAccess] = useMutation<useResumeAccessRequest_documentMutation>( const [requestDocumentAccess] = useMutation<useResumeAccessRequest_documentMutation>(
requestDocumentMutation, requestDocumentMutation,
{ errorToast: false }, { errorToast: false },
@@ -136,9 +118,8 @@ export function useResumeAccessRequest(isAuthenticated: boolean) {
const documentId = searchParams.get(REQUEST_DOCUMENT_PARAM); const documentId = searchParams.get(REQUEST_DOCUMENT_PARAM);
const reportId = searchParams.get(REQUEST_REPORT_PARAM); const reportId = searchParams.get(REQUEST_REPORT_PARAM);
const fileId = searchParams.get(REQUEST_FILE_PARAM); const fileId = searchParams.get(REQUEST_FILE_PARAM);
const all = searchParams.get(REQUEST_ALL_PARAM) === "true";
if (!documentId && !reportId && !fileId && !all) { if (!documentId && !reportId && !fileId) {
return; return;
} }
@@ -199,20 +180,11 @@ export function useResumeAccessRequest(isAuthenticated: boolean) {
variables: { input: { compliancePortalFileId: fileId } }, variables: { input: { compliancePortalFileId: fileId } },
...makeHandlers(continueUrl), ...makeHandlers(continueUrl),
}).catch(() => {}); }).catch(() => {});
return;
} }
const allContinueUrl = buildRequestAllContinueUrl();
clear(REQUEST_ALL_PARAM);
void requestAllAccesses({
variables: {},
...makeHandlers(allContinueUrl),
}).catch(() => {});
}, [ }, [
isAuthenticated, isAuthenticated,
locale, locale,
navigate, navigate,
requestAllAccesses,
requestDocumentAccess, requestDocumentAccess,
requestReportAccess, requestReportAccess,
requestFileAccess, requestFileAccess,

View File

@@ -81,9 +81,9 @@ func TestSecurity_WriteGap_PublishRiskListApproverIDs(t *testing.T) {
// not per-tenant. // not per-tenant.
// //
// CompliancePortalAccess rows are normally created through the trust/v1 public // CompliancePortalAccess rows are normally created through the trust/v1 public
// portal's visitor request flow (requestAllAccesses), which needs a // portal's visitor request flow (requestAccesses / requestDocumentAccess), which
// separate authenticated visitor identity and NDA acceptance. To keep this // needs a separate authenticated visitor identity and NDA acceptance. To keep
// test focused on the fix under test (the Update mutation's FK validation) // this test focused on the fix under test (the Update mutation's FK validation)
// rather than that unrelated flow, the access row's prerequisite state is // rather than that unrelated flow, the access row's prerequisite state is
// seeded directly via SQL against the same Postgres database the e2e probod // seeded directly via SQL against the same Postgres database the e2e probod
// instance runs against, then the real updateCompliancePortalAccess mutation is // instance runs against, then the real updateCompliancePortalAccess mutation is

View File

@@ -132,6 +132,35 @@ func TestCompliancePortal_RequestAccesses_TenantIsolation(t *testing.T) {
) )
} }
// TestCompliancePortal_RequestAccesses_EmptyRejects verifies that requestAccesses
// with no document, report, or file ids is rejected — there is no "request all"
// shortcut; callers must always name the targets explicitly.
func TestCompliancePortal_RequestAccesses_EmptyRejects(t *testing.T) {
t.Parallel()
owner := testutil.NewClient(t, testutil.RoleOwner)
compliancePortalID := lookupCompliancePortalID(t, owner)
trustHost := lookupTrustHost(t, owner, compliancePortalID)
visitor := testutil.SelfProvisionCompliancePortalVisitor(t, trustHost)
err := visitor.ExecuteTrust(trustHost, requestAccessesMutation, map[string]any{
"input": map[string]any{
"documentIds": []string{},
"reportIds": []string{},
"compliancePortalFileIds": []string{},
},
}, nil)
require.Error(t, err, "requestAccesses with empty id lists must be rejected")
assert.Contains(
t,
err.Error(),
"at least one document, report, or file id is required",
"empty request must surface a client validation error",
)
}
// setupPrivatePortalDocument creates a document and marks it privately visible on // setupPrivatePortalDocument creates a document and marks it privately visible on
// the owner's compliance portal, returning the document ID. // the owner's compliance portal, returning the document ID.
func setupPrivatePortalDocument(t *testing.T, owner *testutil.Client) string { func setupPrivatePortalDocument(t *testing.T, owner *testutil.Client) string {

View File

@@ -35,4 +35,5 @@ var (
ErrReportNotFound = errors.New("report not found") ErrReportNotFound = errors.New("report not found")
ErrPortalFileNotFound = errors.New("portal file not found") ErrPortalFileNotFound = errors.New("portal file not found")
ErrPortalFileNotVisible = errors.New("portal file not visible") ErrPortalFileNotVisible = errors.New("portal file not visible")
ErrNoAccessTargets = errors.New("at least one document, report, or file id is required")
) )

View File

@@ -35,6 +35,9 @@ import (
"go.probo.inc/probo/pkg/page" "go.probo.inc/probo/pkg/page"
) )
// PortalAccessRequest carries the explicit resource IDs to request access for.
// Callers must supply at least one ID across the three slices; nil and empty
// both mean "none of that type" (there is no "request all" expansion).
type PortalAccessRequest struct { type PortalAccessRequest struct {
CompliancePortalID gid.GID CompliancePortalID gid.GID
IdentityID gid.GID IdentityID gid.GID
@@ -52,6 +55,10 @@ func (s *Service) RequestPortalAccess(
scope coredata.Scoper, scope coredata.Scoper,
req *PortalAccessRequest, req *PortalAccessRequest,
) (*coredata.CompliancePortalAccess, error) { ) (*coredata.CompliancePortalAccess, error) {
if len(req.DocumentIDs) == 0 && len(req.ReportIDs) == 0 && len(req.CompliancePortalFileIDs) == 0 {
return nil, ErrNoAccessTargets
}
var ( var (
now = time.Now() now = time.Now()
access *coredata.CompliancePortalAccess access *coredata.CompliancePortalAccess
@@ -70,96 +77,6 @@ func (s *Service) RequestPortalAccess(
return fmt.Errorf("cannot load compliance page membership: %w", err) return fmt.Errorf("cannot load compliance page membership: %w", err)
} }
organizationID := compliancePage.OrganizationID
documentIDs := req.DocumentIDs
if req.DocumentIDs == nil {
filter := coredata.NewDocumentCompliancePortalFilter()
allDocuments, err := page.LoadAll(
ctx,
page.OrderBy[coredata.DocumentOrderField]{
Field: coredata.DocumentOrderFieldTitle,
Direction: page.OrderDirectionAsc,
},
func(ctx context.Context, cursor *page.Cursor[coredata.DocumentOrderField]) ([]*coredata.Document, error) {
var batch coredata.Documents
if err := batch.LoadByOrganizationID(ctx, tx, scope, organizationID, cursor, filter); err != nil {
return nil, fmt.Errorf("cannot list documents: %w", err)
}
return batch, nil
},
)
if err != nil {
return err
}
for _, doc := range allDocuments {
documentIDs = append(documentIDs, doc.ID)
}
}
reportIDs := req.ReportIDs
if req.ReportIDs == nil {
auditFilter := coredata.NewAuditCompliancePortalFilter()
allAudits, err := page.LoadAll(
ctx,
page.OrderBy[coredata.AuditOrderField]{
Field: coredata.AuditOrderFieldCreatedAt,
Direction: page.OrderDirectionAsc,
},
func(ctx context.Context, cursor *page.Cursor[coredata.AuditOrderField]) ([]*coredata.Audit, error) {
var batch coredata.Audits
if err := batch.LoadByOrganizationID(ctx, tx, scope, organizationID, cursor, auditFilter); err != nil {
return nil, fmt.Errorf("cannot list audits: %w", err)
}
return batch, nil
},
)
if err != nil {
return err
}
for _, audit := range allAudits {
if audit.ReportFileID != nil {
reportIDs = append(reportIDs, *audit.ReportFileID)
}
}
}
compliancePortalFileIDs := req.CompliancePortalFileIDs
if req.CompliancePortalFileIDs == nil {
filter := coredata.NewCompliancePortalFileFilter(
coredata.WithCompliancePortalFileVisibilities(coredata.CompliancePortalVisibilityPrivate, coredata.CompliancePortalVisibilityNone),
)
allCompliancePortalFiles, err := page.LoadAll(
ctx,
page.OrderBy[coredata.CompliancePortalFileOrderField]{
Field: coredata.CompliancePortalFileOrderFieldCreatedAt,
Direction: page.OrderDirectionDesc,
},
func(ctx context.Context, cursor *page.Cursor[coredata.CompliancePortalFileOrderField]) ([]*coredata.CompliancePortalFile, error) {
var batch coredata.CompliancePortalFiles
if err := batch.LoadByOrganizationID(ctx, tx, scope, organizationID, cursor, filter); err != nil {
return nil, fmt.Errorf("cannot list compliance page files: %w", err)
}
return batch, nil
},
)
if err != nil {
return err
}
for _, file := range allCompliancePortalFiles {
compliancePortalFileIDs = append(compliancePortalFileIDs, file.ID)
}
}
existingAccesses, err := page.LoadAll( existingAccesses, err := page.LoadAll(
ctx, ctx,
page.OrderBy[coredata.CompliancePortalDocumentAccessOrderField]{ page.OrderBy[coredata.CompliancePortalDocumentAccessOrderField]{
@@ -180,9 +97,9 @@ func (s *Service) RequestPortalAccess(
} }
existingDocumentIDs, existingReportIDs, existingCompliancePortalFileIDs := extractExistingIDs(existingAccesses) existingDocumentIDs, existingReportIDs, existingCompliancePortalFileIDs := extractExistingIDs(existingAccesses)
newDocumentIDs := filterExistingIDs(documentIDs, existingDocumentIDs) newDocumentIDs := filterExistingIDs(req.DocumentIDs, existingDocumentIDs)
newReportIDs := filterExistingIDs(reportIDs, existingReportIDs) newReportIDs := filterExistingIDs(req.ReportIDs, existingReportIDs)
newCompliancePortalFileIDs := filterExistingIDs(compliancePortalFileIDs, existingCompliancePortalFileIDs) newCompliancePortalFileIDs := filterExistingIDs(req.CompliancePortalFileIDs, existingCompliancePortalFileIDs)
var accesses coredata.CompliancePortalDocumentAccesses var accesses coredata.CompliancePortalDocumentAccesses

View File

@@ -756,40 +756,6 @@ func (r *frameworkResolver) DarkLogo(ctx context.Context, obj *types.Framework)
return r.loadPublicFile(ctx, *framework.DarkLogoFileID) return r.loadPublicFile(ctx, *framework.DarkLogoFileID)
} }
// RequestAllAccesses is the resolver for the requestAllAccesses field.
func (r *mutationResolver) RequestAllAccesses(ctx context.Context) (*types.RequestAccessesPayload, error) {
compliancePortal := complianceportal.CompliancePortalFromContext(ctx)
scope := coredata.NewScopeFromObjectID(compliancePortal.ID)
visitorService := r.visitor
identity := authn.IdentityFromContext(ctx)
if identity == nil {
return nil, gqlutils.Unauthenticatedf(ctx, "authentication is required to request access")
}
access, err := visitorService.RequestPortalAccess(
ctx, scope,
&visitor.PortalAccessRequest{
CompliancePortalID: compliancePortal.ID,
IdentityID: identity.ID,
DocumentIDs: nil,
ReportIDs: nil,
},
)
if err != nil {
r.logger.ErrorCtx(ctx, "cannot create compliance portal access", log.Error(err))
return nil, gqlutils.Internal(ctx)
}
return &types.RequestAccessesPayload{
CompliancePortalAccess: &types.CompliancePortalAccess{
ID: access.ID,
CreatedAt: access.CreatedAt,
UpdatedAt: access.UpdatedAt,
},
}, nil
}
// ExportDocumentPDF is the resolver for the exportDocumentPDF field. // ExportDocumentPDF is the resolver for the exportDocumentPDF field.
func (r *mutationResolver) ExportDocumentPDF(ctx context.Context, input types.ExportDocumentPDFInput) (*types.ExportDocumentPDFPayload, error) { func (r *mutationResolver) ExportDocumentPDF(ctx context.Context, input types.ExportDocumentPDFInput) (*types.ExportDocumentPDFPayload, error) {
scope := coredata.NewScopeFromObjectID(input.DocumentID) scope := coredata.NewScopeFromObjectID(input.DocumentID)
@@ -1008,6 +974,10 @@ func (r *mutationResolver) RequestDocumentAccess(ctx context.Context, input type
CompliancePortalFileIDs: []gid.GID{}, CompliancePortalFileIDs: []gid.GID{},
}, },
); err != nil { ); err != nil {
if errors.Is(err, visitor.ErrNoAccessTargets) {
return nil, gqlutils.Invalidf(ctx, "at least one document, report, or file id is required")
}
r.logger.ErrorCtx(ctx, "cannot request document access", log.Error(err)) r.logger.ErrorCtx(ctx, "cannot request document access", log.Error(err))
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
@@ -1025,10 +995,21 @@ func (r *mutationResolver) RequestReportAccess(ctx context.Context, input types.
audit, err := visitorService.GetAuditByReportFileID(ctx, scope, input.ReportID) audit, err := visitorService.GetAuditByReportFileID(ctx, scope, input.ReportID)
if err != nil { if err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return nil, gqlutils.NotFoundf(ctx, "report %q not found", input.ReportID)
}
r.logger.ErrorCtx(ctx, "cannot load audit", log.Error(err)) r.logger.ErrorCtx(ctx, "cannot load audit", log.Error(err))
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
// GetAuditByReportFileID is only tenant-scoped, so a report belonging to
// another organization in the same tenant would otherwise be reachable.
// Reject it as not found before an access row can be written.
if audit.OrganizationID != compliancePortal.OrganizationID {
return nil, gqlutils.NotFoundf(ctx, "report %q not found", input.ReportID)
}
if audit.CompliancePortalVisibility == coredata.CompliancePortalVisibilityPublic { if audit.CompliancePortalVisibility == coredata.CompliancePortalVisibilityPublic {
return nil, gqlutils.Invalidf( return nil, gqlutils.Invalidf(
ctx, ctx,
@@ -1051,6 +1032,10 @@ func (r *mutationResolver) RequestReportAccess(ctx context.Context, input types.
CompliancePortalFileIDs: []gid.GID{}, CompliancePortalFileIDs: []gid.GID{},
}, },
); err != nil { ); err != nil {
if errors.Is(err, visitor.ErrNoAccessTargets) {
return nil, gqlutils.Invalidf(ctx, "at least one document, report, or file id is required")
}
r.logger.ErrorCtx(ctx, "cannot request report access", log.Error(err)) r.logger.ErrorCtx(ctx, "cannot request report access", log.Error(err))
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
@@ -1099,6 +1084,10 @@ func (r *mutationResolver) RequestCompliancePortalFileAccess(ctx context.Context
CompliancePortalFileIDs: []gid.GID{input.CompliancePortalFileID}, CompliancePortalFileIDs: []gid.GID{input.CompliancePortalFileID},
}, },
); err != nil { ); err != nil {
if errors.Is(err, visitor.ErrNoAccessTargets) {
return nil, gqlutils.Invalidf(ctx, "at least one document, report, or file id is required")
}
r.logger.ErrorCtx(ctx, "cannot request compliance portal file access", log.Error(err)) r.logger.ErrorCtx(ctx, "cannot request compliance portal file access", log.Error(err))
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
@@ -1119,33 +1108,25 @@ func (r *mutationResolver) RequestAccesses(ctx context.Context, input types.Requ
return nil, gqlutils.Unauthenticatedf(ctx, "authentication is required to request access") return nil, gqlutils.Unauthenticatedf(ctx, "authentication is required to request access")
} }
// Coerce to non-nil slices: an empty list means "none of that type", whereas if len(input.DocumentIds) == 0 && len(input.ReportIds) == 0 && len(input.CompliancePortalFileIds) == 0 {
// a nil slice is interpreted by RequestPortalAccess as "all of that type". return nil, gqlutils.Invalidf(ctx, "at least one document, report, or file id is required")
documentIDs := input.DocumentIds
if documentIDs == nil {
documentIDs = []gid.GID{}
} }
reportIDs := input.ReportIds // Load and validate every target before requesting so a foreign, invisible,
if reportIDs == nil { // or public GID is handled before any access row is written (mirrors the
reportIDs = []gid.GID{} // per-resource resolvers, which load and guard ahead of the request). Only
} // the resolved, non-public ids are forwarded to RequestPortalAccess.
compliancePortalFileIDs := input.CompliancePortalFileIds
if compliancePortalFileIDs == nil {
compliancePortalFileIDs = []gid.GID{}
}
// Load and tenant-check every target before requesting so a foreign or
// invisible GID is rejected before any access row is written (mirrors the
// per-resource resolvers, which guard with a load ahead of the request).
payload := &types.RequestAccessesResultPayload{ payload := &types.RequestAccessesResultPayload{
Documents: make([]*types.Document, 0, len(documentIDs)), Documents: make([]*types.Document, 0, len(input.DocumentIds)),
Audits: make([]*types.Audit, 0, len(reportIDs)), Audits: make([]*types.Audit, 0, len(input.ReportIds)),
Files: make([]*types.CompliancePortalFile, 0, len(compliancePortalFileIDs)), Files: make([]*types.CompliancePortalFile, 0, len(input.CompliancePortalFileIds)),
} }
for _, documentID := range documentIDs { requestDocumentIDs := make([]gid.GID, 0, len(input.DocumentIds))
requestReportIDs := make([]gid.GID, 0, len(input.ReportIds))
requestFileIDs := make([]gid.GID, 0, len(input.CompliancePortalFileIds))
for _, documentID := range input.DocumentIds {
document, err := visitorService.GetDocument(ctx, scope, compliancePortal.OrganizationID, documentID) document, err := visitorService.GetDocument(ctx, scope, compliancePortal.OrganizationID, documentID)
if err != nil { if err != nil {
if errors.Is(err, visitor.ErrDocumentNotFound) || errors.Is(err, visitor.ErrDocumentNotVisible) || errors.Is(err, coredata.ErrResourceNotFound) { if errors.Is(err, visitor.ErrDocumentNotFound) || errors.Is(err, visitor.ErrDocumentNotVisible) || errors.Is(err, coredata.ErrResourceNotFound) {
@@ -1161,10 +1142,17 @@ func (r *mutationResolver) RequestAccesses(ctx context.Context, input types.Requ
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
// Public resources are already accessible; skip them so no needless
// REQUESTED row is created (mirrors the per-resource resolver's guard).
if document.CompliancePortalVisibility == coredata.CompliancePortalVisibilityPublic {
continue
}
requestDocumentIDs = append(requestDocumentIDs, documentID)
payload.Documents = append(payload.Documents, types.NewDocument(document)) payload.Documents = append(payload.Documents, types.NewDocument(document))
} }
for _, reportID := range reportIDs { for _, reportID := range input.ReportIds {
audit, err := visitorService.GetAuditByReportFileID(ctx, scope, reportID) audit, err := visitorService.GetAuditByReportFileID(ctx, scope, reportID)
if err != nil { if err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) { if errors.Is(err, coredata.ErrResourceNotFound) {
@@ -1176,10 +1164,22 @@ func (r *mutationResolver) RequestAccesses(ctx context.Context, input types.Requ
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
// GetAuditByReportFileID is only tenant-scoped, so a report belonging to
// another organization in the same tenant would otherwise be reachable.
// Reject it as not found before an access row can be written.
if audit.OrganizationID != compliancePortal.OrganizationID {
return nil, gqlutils.NotFoundf(ctx, "report %q not found", reportID)
}
if audit.CompliancePortalVisibility == coredata.CompliancePortalVisibilityPublic {
continue
}
requestReportIDs = append(requestReportIDs, reportID)
payload.Audits = append(payload.Audits, types.NewAudit(audit)) payload.Audits = append(payload.Audits, types.NewAudit(audit))
} }
for _, fileID := range compliancePortalFileIDs { for _, fileID := range input.CompliancePortalFileIds {
portalFile, err := visitorService.GetPortalFile(ctx, scope, compliancePortal.OrganizationID, fileID) portalFile, err := visitorService.GetPortalFile(ctx, scope, compliancePortal.OrganizationID, fileID)
if err != nil { if err != nil {
if errors.Is(err, visitor.ErrPortalFileNotFound) || errors.Is(err, visitor.ErrPortalFileNotVisible) { if errors.Is(err, visitor.ErrPortalFileNotFound) || errors.Is(err, visitor.ErrPortalFileNotVisible) {
@@ -1191,19 +1191,33 @@ func (r *mutationResolver) RequestAccesses(ctx context.Context, input types.Requ
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)
} }
if portalFile.CompliancePortalVisibility == coredata.CompliancePortalVisibilityPublic {
continue
}
requestFileIDs = append(requestFileIDs, fileID)
payload.Files = append(payload.Files, types.NewCompliancePortalFile(portalFile)) payload.Files = append(payload.Files, types.NewCompliancePortalFile(portalFile))
} }
// All supplied ids were public (already accessible); nothing to request.
if len(requestDocumentIDs) == 0 && len(requestReportIDs) == 0 && len(requestFileIDs) == 0 {
return payload, nil
}
if _, err := visitorService.RequestPortalAccess( if _, err := visitorService.RequestPortalAccess(
ctx, scope, ctx, scope,
&visitor.PortalAccessRequest{ &visitor.PortalAccessRequest{
CompliancePortalID: compliancePortal.ID, CompliancePortalID: compliancePortal.ID,
IdentityID: identity.ID, IdentityID: identity.ID,
DocumentIDs: documentIDs, DocumentIDs: requestDocumentIDs,
ReportIDs: reportIDs, ReportIDs: requestReportIDs,
CompliancePortalFileIDs: compliancePortalFileIDs, CompliancePortalFileIDs: requestFileIDs,
}, },
); err != nil { ); err != nil {
if errors.Is(err, visitor.ErrNoAccessTargets) {
return nil, gqlutils.Invalidf(ctx, "at least one document, report, or file id is required")
}
r.logger.ErrorCtx(ctx, "cannot request accesses", log.Error(err)) r.logger.ErrorCtx(ctx, "cannot request accesses", log.Error(err))
return nil, gqlutils.Internal(ctx) return nil, gqlutils.Internal(ctx)

View File

@@ -466,8 +466,6 @@ type DocumentAccess implements Node {
} }
extend type Mutation { extend type Mutation {
requestAllAccesses: RequestAccessesPayload! @authentication(required: PRESENT) @nda
exportDocumentPDF(input: ExportDocumentPDFInput!): ExportDocumentPDFPayload! exportDocumentPDF(input: ExportDocumentPDFInput!): ExportDocumentPDFPayload!
@authentication(required: OPTIONAL) @nda @authentication(required: OPTIONAL) @nda
@@ -507,10 +505,6 @@ type RequestFileAccessPayload {
file: CompliancePortalFile file: CompliancePortalFile
} }
type RequestAccessesPayload {
compliancePortalAccess: CompliancePortalAccess!
}
# Returns the affected nodes so the client can update each row in place. Mirrors # Returns the affected nodes so the client can update each row in place. Mirrors
# the per-resource payloads but for a selection-scoped batch request. # the per-resource payloads but for a selection-scoped batch request.
type RequestAccessesResultPayload { type RequestAccessesResultPayload {