from __future__ import annotations from collections.abc import Mapping, Sequence from urllib.parse import quote from sqlalchemy import func, select from sqlalchemy.orm import Session from govoplan_core.auth import ApiPrincipal from govoplan_core.core.events import PlatformEvent from govoplan_core.core.modules import ModuleContext from govoplan_core.core.search import ( SearchAuthorizationRequest, SearchBackfillPage, SearchBackfillRequest, SearchDocument, SearchIndexChange, SearchResourceReference, SearchResourceType, ) from govoplan_cases.backend.db.models import ( CaseAccessGrant, CaseIdentity, CaseRecordRevision, ) from govoplan_cases.backend.service import can_access_case PROVIDER_ID = "cases.cases" RESOURCE_TYPE = "case" READ_SCOPE = "cases:case:read" ADMIN_SCOPE = "cases:case:admin" SEARCH_ACCESS_PURPOSE = "cases.search" class CasesSearchSource: def resource_types(self) -> Sequence[SearchResourceType]: return ( SearchResourceType( provider_id=PROVIDER_ID, module_id="cases", resource_type=RESOURCE_TYPE, label="Cases", requires_authorization_recheck=True, ), ) def backfill( self, session: object, *, request: SearchBackfillRequest, ) -> SearchBackfillPage: _assert_source(request.provider_id, request.resource_type) db = _session(session) statement = ( select(CaseRecordRevision, CaseIdentity) .join(CaseIdentity, CaseIdentity.id == CaseRecordRevision.identity_id) .where( CaseRecordRevision.tenant_id == request.tenant_id, CaseRecordRevision.superseded_at.is_(None), ) ) if request.cursor: statement = statement.where(CaseRecordRevision.case_id > request.cursor) rows = list( db.execute( statement.order_by(CaseRecordRevision.case_id).limit(request.limit + 1) ).all() ) has_more = len(rows) > request.limit selected = rows[: request.limit] grants = _grants_by_case( db, request.tenant_id, [row.case_id for row, _identity in selected], ) high_watermark = db.scalar( select(func.max(CaseRecordRevision.recorded_at)).where( CaseRecordRevision.tenant_id == request.tenant_id, CaseRecordRevision.superseded_at.is_(None), ) ) return SearchBackfillPage( documents=tuple( _document( row, identity=identity, grants=grants.get(row.case_id, ()), ) for row, identity in selected ), next_cursor=selected[-1][0].case_id if has_more and selected else None, complete=not has_more, high_watermark=high_watermark.isoformat() if high_watermark else None, ) def authorize( self, session: object, principal: object, *, requests: Sequence[SearchAuthorizationRequest], ) -> Mapping[str, bool]: decisions = {item.reference.key: False for item in requests} if not isinstance(principal, ApiPrincipal) or not principal.has(READ_SCOPE): return decisions db = _session(session) for request in requests: reference = request.reference if ( reference.tenant_id != principal.tenant_id or reference.module_id != "cases" or reference.resource_type != RESOURCE_TYPE ): continue decisions[reference.key] = can_access_case( db, principal, case_id=reference.resource_id, permission="read", purpose=SEARCH_ACCESS_PURPOSE, ) return decisions def index_changes_for_event( self, session: object, *, event: PlatformEvent, delivery_key: str, ) -> Sequence[SearchIndexChange]: if ( event.module_id != "cases" or event.tenant is None or event.resource is None or event.resource.type != RESOURCE_TYPE or event.resource.id is None ): return () db = _session(session) row = db.scalar( select(CaseRecordRevision).where( CaseRecordRevision.tenant_id == event.tenant.id, CaseRecordRevision.case_id == event.resource.id, CaseRecordRevision.superseded_at.is_(None), ) ) identity = db.scalar( select(CaseIdentity).where( CaseIdentity.tenant_id == event.tenant.id, CaseIdentity.case_id == event.resource.id, ) ) deleted = row is None or identity is None cursor = event.event_id document = None if not deleted: document = _document( row, identity=identity, grants=_grants_by_case( db, event.tenant.id, [row.case_id], ).get(row.case_id, ()), change_cursor=cursor, ) reference = SearchResourceReference( tenant_id=event.tenant.id, module_id="cases", resource_type=RESOURCE_TYPE, resource_id=event.resource.id, ) return ( SearchIndexChange( change_id=f"{delivery_key}:{PROVIDER_ID}", provider_id=PROVIDER_ID, kind="delete" if deleted else "upsert", reference=reference, source_revision=(document.source_revision if document else cursor), cursor=cursor, document=document, occurred_at=event.occurred_at, ), ) def create_cases_search_source(_context: ModuleContext) -> CasesSearchSource: return CasesSearchSource() def _document( row: CaseRecordRevision, *, identity: CaseIdentity, grants: Sequence[CaseAccessGrant], change_cursor: str | None = None, ) -> SearchDocument: tokens = [f"scope:{READ_SCOPE}", f"scope:{ADMIN_SCOPE}"] if identity.created_by: tokens.extend( (f"account:{identity.created_by}", f"membership:{identity.created_by}") ) prefixes = { "account": "account", "identity": "identity", "group": "group", "role": "role", "function": "function", "function_assignment": "function", } for grant in grants: prefix = prefixes.get(grant.subject_kind) if prefix: tokens.append(f"{prefix}:{grant.subject_id}") return SearchDocument( tenant_id=row.tenant_id, module_id="cases", provider_id=PROVIDER_ID, resource_type=RESOURCE_TYPE, resource_id=row.case_id, title=row.title, url=( f"/cases/{quote(row.case_id, safe='')}" f"?purpose={quote(SEARCH_ACCESS_PURPOSE, safe='')}" ), summary=f"{identity.case_number} - {row.status_key}", body=row.search_text[:200_000], keywords=( identity.case_number[:200], row.case_type_key[:200], row.status_key[:200], ), visibility="restricted", acl_tokens=tuple(dict.fromkeys(tokens)), metadata={ "case_number": identity.case_number, "case_type_key": row.case_type_key, "status_key": row.status_key, "access_mode": row.access_mode, }, source_revision=str(row.revision), change_cursor=change_cursor, source_updated_at=row.recorded_at, requires_authorization_recheck=True, ) def _grants_by_case( session: Session, tenant_id: str, case_ids: Sequence[str], ) -> dict[str, tuple[CaseAccessGrant, ...]]: grouped: dict[str, list[CaseAccessGrant]] = { case_id: [] for case_id in case_ids } if not case_ids: return {} rows = session.scalars( select(CaseAccessGrant).where( CaseAccessGrant.tenant_id == tenant_id, CaseAccessGrant.case_id.in_(tuple(case_ids)), CaseAccessGrant.active.is_(True), ) ) for row in rows: grouped.setdefault(row.case_id, []).append(row) return {key: tuple(value) for key, value in grouped.items()} def _assert_source(provider_id: str, resource_type: str) -> None: if provider_id != PROVIDER_ID or resource_type != RESOURCE_TYPE: raise ValueError("Unsupported Cases search source.") def _session(value: object) -> Session: if not isinstance(value, Session): raise TypeError("Cases search requires a SQLAlchemy session.") return value __all__ = [ "CasesSearchSource", "PROVIDER_ID", "RESOURCE_TYPE", "SEARCH_ACCESS_PURPOSE", "create_cases_search_source", ]