diff --git a/src/govoplan_core/core/modules.py b/src/govoplan_core/core/modules.py index c64552b..732dbf1 100644 --- a/src/govoplan_core/core/modules.py +++ b/src/govoplan_core/core/modules.py @@ -8,7 +8,10 @@ from govoplan_core.core.views import ViewSurface if TYPE_CHECKING: from fastapi import APIRouter - from govoplan_core.core.search import SearchProviderRegistration + from govoplan_core.core.search import ( + SearchProviderRegistration, + SearchSourceProviderRegistration, + ) SUPPORTED_MANIFEST_CONTRACT_VERSION = "1" @@ -358,6 +361,7 @@ class ModuleManifest: uninstall_guard_providers: tuple[UninstallGuardProvider, ...] = () capability_factories: Mapping[str, CapabilityFactory] = field(default_factory=dict) search_providers: tuple["SearchProviderRegistration", ...] = () + search_sources: tuple["SearchSourceProviderRegistration", ...] = () compatibility: ModuleCompatibility = field(default_factory=ModuleCompatibility) on_activate: LifecycleHook | None = None on_deactivate: LifecycleHook | None = None diff --git a/src/govoplan_core/core/registry.py b/src/govoplan_core/core/registry.py index f3c56e4..3371a6b 100644 --- a/src/govoplan_core/core/registry.py +++ b/src/govoplan_core/core/registry.py @@ -25,7 +25,12 @@ from govoplan_core.core.modules import ( user_workflow_scope_condition_issues, ) from govoplan_core.core.versioning import format_version_range, version_range_is_valid, version_satisfies_range -from govoplan_core.core.search import RegisteredSearchProvider, SearchProvider +from govoplan_core.core.search import ( + RegisteredSearchProvider, + RegisteredSearchSourceProvider, + SearchProvider, + SearchSourceProvider, +) from govoplan_core.core.views import ( ViewSurface, module_view_surface_id, @@ -63,6 +68,10 @@ class PlatformRegistry: self._capability_context: ModuleContext | None = None self._search_provider_registrations: list[RegisteredSearchProvider] = [] self._search_providers: dict[str, SearchProvider] = {} + self._search_source_registrations: list[ + RegisteredSearchSourceProvider + ] = [] + self._search_sources: dict[str, SearchSourceProvider] = {} def register(self, manifest: ModuleManifest) -> ModuleManifest: if manifest.id in self._manifests: @@ -82,6 +91,13 @@ class PlatformRegistry: registration=registration, ) ) + for registration in manifest.search_sources: + self._search_source_registrations.append( + RegisteredSearchSourceProvider( + module_id=manifest.id, + registration=registration, + ) + ) return manifest def replace(self, manifests: Iterable[ModuleManifest]) -> RegistrySnapshot: @@ -102,8 +118,12 @@ class PlatformRegistry: self._search_provider_registrations = list( replacement._search_provider_registrations ) + self._search_source_registrations = list( + replacement._search_source_registrations + ) self._capabilities.clear() self._search_providers.clear() + self._search_sources.clear() return snapshot def get(self, module_id: str) -> ModuleManifest | None: @@ -148,6 +168,7 @@ class PlatformRegistry: self._capability_context = context self._capabilities.clear() self._search_providers.clear() + self._search_sources.clear() def register_capability_factory(self, module_id: str, name: str, factory: CapabilityFactory) -> None: if name in self._capability_factories: @@ -206,6 +227,46 @@ class PlatformRegistry: providers.append((registered, provider)) return tuple(providers) + def search_source_registrations( + self, + ) -> tuple[RegisteredSearchSourceProvider, ...]: + return tuple( + sorted( + self._search_source_registrations, + key=lambda item: ( + item.registration.order, + item.module_id, + item.registration.id, + ), + ) + ) + + def search_sources( + self, + ) -> tuple[ + tuple[RegisteredSearchSourceProvider, SearchSourceProvider], + ..., + ]: + if self._capability_context is None: + if self._search_source_registrations: + raise RegistryError( + "Search source context is not configured." + ) + return () + providers: list[ + tuple[RegisteredSearchSourceProvider, SearchSourceProvider] + ] = [] + for registered in self.search_source_registrations(): + key = f"{registered.module_id}:{registered.registration.id}" + provider = self._search_sources.get(key) + if provider is None: + provider = registered.registration.create( + self._capability_context + ) + self._search_sources[key] = provider + providers.append((registered, provider)) + return tuple(providers) + def register_tenant_summary_provider(self, module_id: str, provider: TenantSummaryProvider) -> None: self._tenant_summary_providers[module_id] = provider @@ -521,6 +582,19 @@ def _validate_manifest_contract_lists(manifest: ModuleManifest) -> None: f"{registration.id!r}" ) provider_ids.add(registration.id) + source_ids: set[str] = set() + for registration in manifest.search_sources: + if not _INTERFACE_NAME_RE.match(registration.id): + raise RegistryError( + f"Module {manifest.id!r} search source id must be " + f"namespaced: {registration.id!r}" + ) + if registration.id in source_ids: + raise RegistryError( + f"Module {manifest.id!r} declares duplicate search source " + f"{registration.id!r}" + ) + source_ids.add(registration.id) def _validate_manifest_overlaps(manifest: ModuleManifest) -> None: diff --git a/src/govoplan_core/core/search.py b/src/govoplan_core/core/search.py index ee60627..7f9e9b1 100644 --- a/src/govoplan_core/core/search.py +++ b/src/govoplan_core/core/search.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass, field +from datetime import datetime from typing import Literal, Protocol, runtime_checkable from govoplan_core.core.external_references import ExternalObjectReference @@ -10,6 +11,8 @@ from govoplan_core.core.modules import ModuleContext SearchContextKind = Literal["global", "module", "resource"] SearchVisibility = Literal["tenant", "restricted"] +SearchIndexChangeKind = Literal["upsert", "delete"] +SEARCH_SOURCE_CONTRACT_VERSION = "1" @dataclass(frozen=True, slots=True) @@ -22,6 +25,8 @@ class SearchQuery: context_id: str | None = None limit: int = 25 offset: int = 0 + cursor: str | None = None + language: str = "simple" def __post_init__(self) -> None: normalized_text = self.text.strip() @@ -33,6 +38,16 @@ class SearchQuery: raise ValueError("Search result limits must be between 1 and 200.") if self.offset < 0: raise ValueError("Search offsets cannot be negative.") + if self.cursor and self.offset: + raise ValueError( + "Search cursor and offset pagination cannot be combined." + ) + if not self.language.replace("_", "").isalnum(): + raise ValueError("Search language configuration is invalid.") + if len(self.module_ids) > 50 or len(self.resource_types) > 50: + raise ValueError( + "Search queries support at most 50 module and type filters." + ) object.__setattr__(self, "text", normalized_text) @@ -50,6 +65,68 @@ class SearchResult: breadcrumbs: tuple[str, ...] = () external_reference: ExternalObjectReference | None = None metadata: Mapping[str, object] = field(default_factory=dict) + source_revision: str | None = None + provenance: Mapping[str, object] = field(default_factory=dict) + + +@dataclass(frozen=True, slots=True) +class SearchResourceReference: + tenant_id: str + module_id: str + resource_type: str + resource_id: str + + def __post_init__(self) -> None: + for field_name in ( + "tenant_id", + "module_id", + "resource_type", + "resource_id", + ): + if not str(getattr(self, field_name) or "").strip(): + raise ValueError( + f"Search resource reference {field_name} is required." + ) + + @property + def key(self) -> str: + values = ( + self.tenant_id, + self.module_id, + self.resource_type, + self.resource_id, + ) + return "|".join( + f"{len(value)}:{value}" + for value in values + ) + + +@dataclass(frozen=True, slots=True) +class SearchResourceType: + provider_id: str + module_id: str + resource_type: str + label: str + index_version: int = 1 + language: str = "simple" + requires_authorization_recheck: bool = True + + def __post_init__(self) -> None: + if not all( + value.strip() + for value in ( + self.provider_id, + self.module_id, + self.resource_type, + self.label, + ) + ): + raise ValueError( + "Search resource type identifiers and label are required." + ) + if self.index_version < 1: + raise ValueError("Search index versions start at one.") @dataclass(frozen=True, slots=True) @@ -67,6 +144,13 @@ class SearchDocument: acl_tokens: tuple[str, ...] = () external_reference: ExternalObjectReference | None = None metadata: Mapping[str, object] = field(default_factory=dict) + provider_id: str | None = None + source_revision: str = "1" + change_cursor: str | None = None + source_updated_at: datetime | None = None + language: str = "simple" + index_version: int = 1 + requires_authorization_recheck: bool = False def __post_init__(self) -> None: required = { @@ -80,10 +164,152 @@ class SearchDocument: for field_name, value in required.items(): if not str(value or "").strip(): raise ValueError(f"Search document {field_name} is required.") + limits = { + "module_id": 100, + "resource_type": 100, + "resource_id": 255, + "title": 500, + "url": 1500, + "summary": 4_000, + "body": 200_000, + "source_revision": 255, + "change_cursor": 500, + "language": 32, + } + for field_name, limit in limits.items(): + value = getattr(self, field_name) + if value is not None and len(str(value)) > limit: + raise ValueError( + f"Search document {field_name} is limited to " + f"{limit} characters." + ) + if len(self.keywords) > 100 or any( + len(keyword) > 200 + for keyword in self.keywords + ): + raise ValueError( + "Search documents support at most 100 bounded keywords." + ) + if len(self.acl_tokens) > 500 or any( + len(token) > 500 + for token in self.acl_tokens + ): + raise ValueError( + "Search documents support at most 500 bounded ACL tokens." + ) if self.visibility == "restricted" and not self.acl_tokens: raise ValueError( "Restricted search documents require at least one ACL token." ) + if not self.source_revision.strip(): + raise ValueError( + "Search documents require a stable source revision." + ) + if self.index_version < 1: + raise ValueError("Search document index versions start at one.") + if ( + self.requires_authorization_recheck + and not str(self.provider_id or "").strip() + ): + raise ValueError( + "Search documents requiring authorization rechecks must " + "name their source provider." + ) + if self.provider_id is not None and len(self.provider_id) > 200: + raise ValueError( + "Search document provider_id is limited to 200 characters." + ) + + @property + def reference(self) -> SearchResourceReference: + return SearchResourceReference( + tenant_id=self.tenant_id, + module_id=self.module_id, + resource_type=self.resource_type, + resource_id=self.resource_id, + ) + + +@dataclass(frozen=True, slots=True) +class SearchAuthorizationRequest: + reference: SearchResourceReference + source_revision: str + + +@dataclass(frozen=True, slots=True) +class SearchBackfillRequest: + tenant_id: str + provider_id: str + resource_type: str + rebuild_id: str + cursor: str | None = None + limit: int = 100 + + def __post_init__(self) -> None: + if not 1 <= self.limit <= 500: + raise ValueError( + "Search backfill page size must be between 1 and 500." + ) + + +@dataclass(frozen=True, slots=True) +class SearchBackfillPage: + documents: tuple[SearchDocument, ...] + next_cursor: str | None + complete: bool + high_watermark: str | None = None + + def __post_init__(self) -> None: + if len(self.documents) > 500: + raise ValueError( + "Search backfill providers returned more than 500 documents." + ) + if not self.complete and not self.next_cursor: + raise ValueError( + "Incomplete search backfill pages require a next cursor." + ) + + +@dataclass(frozen=True, slots=True) +class SearchIndexChange: + change_id: str + provider_id: str + kind: SearchIndexChangeKind + reference: SearchResourceReference + source_revision: str + cursor: str + document: SearchDocument | None = None + occurred_at: datetime | None = None + + def __post_init__(self) -> None: + if not all( + value.strip() + for value in ( + self.change_id, + self.provider_id, + self.source_revision, + self.cursor, + ) + ): + raise ValueError( + "Search index changes require stable identity and revision." + ) + if self.kind == "upsert": + if ( + self.document is None + or self.document.reference != self.reference + or self.document.provider_id != self.provider_id + or self.document.source_revision + != self.source_revision + or self.document.change_cursor != self.cursor + ): + raise ValueError( + "Search upserts require a matching document." + ) + elif self.document is not None: + raise ValueError( + "Search deletes must not contain a document." + ) @runtime_checkable @@ -121,8 +347,43 @@ class SearchIndexWriter(Protocol): ) -> bool: ... + def enqueue_change( + self, + session: object, + *, + change: SearchIndexChange, + ) -> bool: + ... + + +@runtime_checkable +class SearchSourceProvider(Protocol): + def resource_types(self) -> Sequence[SearchResourceType]: + ... + + def backfill( + self, + session: object, + *, + request: SearchBackfillRequest, + ) -> SearchBackfillPage: + ... + + def authorize( + self, + session: object, + principal: object, + *, + requests: Sequence[SearchAuthorizationRequest], + ) -> Mapping[str, bool]: + """Return an explicit decision for every requested reference key.""" + SearchProviderFactory = Callable[[ModuleContext], SearchProvider] +SearchSourceProviderFactory = Callable[ + [ModuleContext], + SearchSourceProvider, +] @dataclass(frozen=True, slots=True) @@ -147,6 +408,28 @@ class RegisteredSearchProvider: registration: SearchProviderRegistration +@dataclass(frozen=True, slots=True) +class SearchSourceProviderRegistration: + id: str + factory: SearchSourceProviderFactory + order: int = 100 + + def create(self, context: ModuleContext) -> SearchSourceProvider: + provider = self.factory(context) + if not isinstance(provider, SearchSourceProvider): + raise TypeError( + f"Search source {self.id!r} does not implement " + "SearchSourceProvider." + ) + return provider + + +@dataclass(frozen=True, slots=True) +class RegisteredSearchSourceProvider: + module_id: str + registration: SearchSourceProviderRegistration + + CAPABILITY_SEARCH_INDEX_WRITER = "search.index_writer" @@ -164,15 +447,27 @@ def search_index_writer(registry: object | None) -> SearchIndexWriter | None: __all__ = [ "CAPABILITY_SEARCH_INDEX_WRITER", + "SEARCH_SOURCE_CONTRACT_VERSION", "RegisteredSearchProvider", + "RegisteredSearchSourceProvider", + "SearchAuthorizationRequest", + "SearchBackfillPage", + "SearchBackfillRequest", "SearchContextKind", "SearchDocument", + "SearchIndexChange", + "SearchIndexChangeKind", "SearchIndexWriter", "SearchProvider", "SearchProviderFactory", "SearchProviderRegistration", "SearchQuery", + "SearchResourceReference", + "SearchResourceType", "SearchResult", + "SearchSourceProvider", + "SearchSourceProviderFactory", + "SearchSourceProviderRegistration", "SearchVisibility", "search_index_writer", ] diff --git a/tests/test_search_contract.py b/tests/test_search_contract.py index 593de92..5c42726 100644 --- a/tests/test_search_contract.py +++ b/tests/test_search_contract.py @@ -5,11 +5,18 @@ import unittest from govoplan_core.core.modules import ModuleContext, ModuleManifest from govoplan_core.core.registry import PlatformRegistry from govoplan_core.core.search import ( + SearchAuthorizationRequest, + SearchBackfillPage, + SearchBackfillRequest, SearchDocument, SearchProvider, SearchProviderRegistration, SearchQuery, + SearchResourceReference, + SearchResourceType, SearchResult, + SearchSourceProvider, + SearchSourceProviderRegistration, ) @@ -28,6 +35,34 @@ class _Provider: ) +class _Source: + def resource_types(self): + return ( + SearchResourceType( + provider_id="test.records", + module_id="test", + resource_type="record", + label="Records", + ), + ) + + def backfill(self, session, *, request): + del session + return SearchBackfillPage( + documents=(), + next_cursor=None, + complete=True, + high_watermark=request.cursor, + ) + + def authorize(self, session, principal, *, requests): + del session, principal + return { + request.reference.key: True + for request in requests + } + + class SearchContractTests(unittest.TestCase): def test_query_and_restricted_document_are_bounded(self) -> None: query = SearchQuery(text=" permit ", tenant_id="tenant-1", limit=20) @@ -52,6 +87,17 @@ class SearchContractTests(unittest.TestCase): title="Permit", url="/cases/case-1", ) + with self.assertRaisesRegex(ValueError, "body"): + SearchDocument( + tenant_id="tenant-1", + module_id="cases", + resource_type="case", + resource_id="case-1", + title="Permit", + url="/cases/case-1", + body="x" * 200_001, + acl_tokens=("account:account-1",), + ) def test_manifest_provider_registration_resolves_lazily(self) -> None: registry = PlatformRegistry() @@ -67,6 +113,12 @@ class SearchContractTests(unittest.TestCase): resource_types=("record",), ), ), + search_sources=( + SearchSourceProviderRegistration( + id="test.records", + factory=lambda context: _Source(), + ), + ), ) ) registry.configure_capability_context( @@ -76,6 +128,35 @@ class SearchContractTests(unittest.TestCase): providers = registry.search_providers() self.assertEqual(1, len(providers)) self.assertIsInstance(providers[0][1], SearchProvider) + sources = registry.search_sources() + self.assertEqual(1, len(sources)) + self.assertIsInstance(sources[0][1], SearchSourceProvider) + request = SearchAuthorizationRequest( + reference=SearchResourceReference( + tenant_id="tenant-1", + module_id="test", + resource_type="record", + resource_id="1", + ), + source_revision="1", + ) + self.assertTrue( + sources[0][1].authorize( + object(), + object(), + requests=(request,), + )[request.reference.key] + ) + page = sources[0][1].backfill( + object(), + request=SearchBackfillRequest( + tenant_id="tenant-1", + provider_id="test.records", + resource_type="record", + rebuild_id="rebuild-1", + ), + ) + self.assertTrue(page.complete) if __name__ == "__main__":