From a156e3d4fcb12716b9b5183e347308680e7ab889 Mon Sep 17 00:00:00 2001 From: Albrecht Degering Date: Wed, 29 Jul 2026 18:08:52 +0200 Subject: [PATCH] feat: implement durable permission-aware search --- README.md | 20 + src/govoplan_search/backend/db/models.py | 256 +++- src/govoplan_search/backend/manifest.py | 7 +- .../b2c3d4e5f607_v0114_search_lifecycle.py | 311 +++++ src/govoplan_search/backend/router.py | 274 +++- src/govoplan_search/backend/schemas.py | 53 + src/govoplan_search/backend/service.py | 1243 ++++++++++++++++- tests/test_migrations.py | 4 +- tests/test_search_service.py | 359 ++++- webui/src/api/search.ts | 10 +- webui/src/features/search/SearchPage.tsx | 73 +- webui/src/styles/search.css | 30 + 12 files changed, 2540 insertions(+), 100 deletions(-) create mode 100644 src/govoplan_search/backend/migrations/versions/b2c3d4e5f607_v0114_search_lifecycle.py diff --git a/README.md b/README.md index fd7544f..871512c 100644 --- a/README.md +++ b/README.md @@ -16,3 +16,23 @@ development fallback. Other modules may: An optional OpenSearch adapter is a later provider, not a hard dependency. Source modules remain responsible for defining visibility and authorization. + +## Index lifecycle + +Source modules register a versioned `search_sources` provider. A provider +declares its resource types and index version, returns bounded resumable +backfill pages, and batch-rechecks current authorization for sensitive +resources. Incremental writes use `SearchIndexChange` and +`search.index_writer.enqueue_change()` so change IDs are durable and +idempotent in the same database transaction as the caller. + +The built-in backend exposes opaque cursor pagination and does not return +pre-authorization totals. PostgreSQL uses full-text search and will add +trigram indexes when `pg_trgm` is already installed; SQLite remains a bounded +development fallback. + +Tenant search administrators can inspect `/api/v1/search/admin/diagnostics`, +reconcile disabled modules, process queued changes, and start or continue +resumable provider rebuilds. Rows requiring a source authorization recheck +are omitted when their provider is unavailable, stale, or fails to return an +explicit allow decision. diff --git a/src/govoplan_search/backend/db/models.py b/src/govoplan_search/backend/db/models.py index 4a784c9..68d3e26 100644 --- a/src/govoplan_search/backend/db/models.py +++ b/src/govoplan_search/backend/db/models.py @@ -5,9 +5,11 @@ from datetime import datetime from typing import Any from sqlalchemy import ( + Boolean, DateTime, ForeignKey, Index, + Integer, JSON, String, Text, @@ -15,7 +17,7 @@ from sqlalchemy import ( ) from sqlalchemy.orm import Mapped, mapped_column, relationship -from govoplan_core.db.base import Base, TimestampMixin +from govoplan_core.db.base import Base, TimestampMixin, utcnow def new_uuid() -> str: @@ -39,11 +41,23 @@ class SearchIndexDocument(Base, TimestampMixin): "resource_type", ), Index("ix_search_document_visibility", "tenant_id", "visibility"), + Index( + "ix_search_document_provider", + "tenant_id", + "provider_id", + "resource_type", + "active", + ), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid) tenant_id: Mapped[str] = mapped_column(String(36), nullable=False, index=True) module_id: Mapped[str] = mapped_column(String(100), nullable=False, index=True) + provider_id: Mapped[str] = mapped_column( + String(200), + nullable=False, + index=True, + ) resource_type: Mapped[str] = mapped_column( String(100), nullable=False, index=True ) @@ -66,6 +80,44 @@ class SearchIndexDocument(Base, TimestampMixin): "metadata", JSON, default=dict, nullable=False ) content_hash: Mapped[str] = mapped_column(String(64), nullable=False) + source_revision: Mapped[str] = mapped_column( + String(255), + nullable=False, + ) + change_cursor: Mapped[str | None] = mapped_column( + String(500), + nullable=True, + ) + source_updated_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + language: Mapped[str] = mapped_column( + String(32), + default="simple", + nullable=False, + ) + index_version: Mapped[int] = mapped_column( + Integer, + default=1, + nullable=False, + ) + requires_authorization_recheck: Mapped[bool] = mapped_column( + Boolean, + default=False, + nullable=False, + ) + rebuild_id: Mapped[str | None] = mapped_column( + String(36), + nullable=True, + index=True, + ) + active: Mapped[bool] = mapped_column( + Boolean, + default=True, + nullable=False, + index=True, + ) indexed_at: Mapped[datetime] = mapped_column( DateTime(timezone=True), nullable=False ) @@ -100,4 +152,204 @@ class SearchIndexAclToken(Base): ) -__all__ = ["SearchIndexAclToken", "SearchIndexDocument", "new_uuid"] +class SearchIndexState(Base, TimestampMixin): + __tablename__ = "search_index_states" + __table_args__ = ( + UniqueConstraint( + "tenant_id", + "provider_id", + "resource_type", + name="uq_search_index_state_source", + ), + Index( + "ix_search_index_state_status", + "tenant_id", + "status", + ), + ) + + id: Mapped[str] = mapped_column( + String(36), + primary_key=True, + default=new_uuid, + ) + tenant_id: Mapped[str] = mapped_column( + String(36), + nullable=False, + index=True, + ) + provider_id: Mapped[str] = mapped_column( + String(200), + nullable=False, + index=True, + ) + module_id: Mapped[str] = mapped_column( + String(100), + nullable=False, + index=True, + ) + resource_type: Mapped[str] = mapped_column( + String(100), + nullable=False, + index=True, + ) + index_version: Mapped[int] = mapped_column( + Integer, + default=1, + nullable=False, + ) + status: Mapped[str] = mapped_column( + String(30), + default="idle", + nullable=False, + index=True, + ) + rebuild_id: Mapped[str | None] = mapped_column( + String(36), + nullable=True, + ) + checkpoint_cursor: Mapped[str | None] = mapped_column( + String(500), + nullable=True, + ) + high_watermark: Mapped[str | None] = mapped_column( + String(500), + nullable=True, + ) + last_change_cursor: Mapped[str | None] = mapped_column( + String(500), + nullable=True, + ) + indexed_documents: Mapped[int] = mapped_column( + Integer, + default=0, + nullable=False, + ) + rejected_documents: Mapped[int] = mapped_column( + Integer, + default=0, + nullable=False, + ) + rebuild_started_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + rebuild_completed_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + last_success_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + last_error: Mapped[str | None] = mapped_column(Text, nullable=True) + + +class SearchIndexChangeQueue(Base, TimestampMixin): + __tablename__ = "search_index_change_queue" + __table_args__ = ( + UniqueConstraint( + "change_id", + name="uq_search_index_change_queue_change", + ), + Index( + "ix_search_index_change_queue_pending", + "status", + "available_at", + "created_at", + ), + Index( + "ix_search_index_change_queue_source", + "tenant_id", + "provider_id", + "resource_type", + ), + ) + + id: Mapped[str] = mapped_column( + String(36), + primary_key=True, + default=new_uuid, + ) + change_id: Mapped[str] = mapped_column( + String(255), + nullable=False, + index=True, + ) + tenant_id: Mapped[str] = mapped_column( + String(36), + nullable=False, + index=True, + ) + provider_id: Mapped[str] = mapped_column( + String(200), + nullable=False, + index=True, + ) + module_id: Mapped[str] = mapped_column( + String(100), + nullable=False, + index=True, + ) + resource_type: Mapped[str] = mapped_column( + String(100), + nullable=False, + index=True, + ) + resource_id: Mapped[str] = mapped_column( + String(255), + nullable=False, + index=True, + ) + kind: Mapped[str] = mapped_column( + String(20), + nullable=False, + ) + source_revision: Mapped[str] = mapped_column( + String(255), + nullable=False, + ) + source_cursor: Mapped[str] = mapped_column( + String(500), + nullable=False, + ) + document_: Mapped[dict[str, Any] | None] = mapped_column( + "document", + JSON, + nullable=True, + ) + occurred_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + status: Mapped[str] = mapped_column( + String(30), + default="queued", + nullable=False, + index=True, + ) + attempts: Mapped[int] = mapped_column( + Integer, + default=0, + nullable=False, + ) + available_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + default=utcnow, + nullable=False, + index=True, + ) + processed_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), + nullable=True, + ) + error: Mapped[str | None] = mapped_column(Text, nullable=True) + + +__all__ = [ + "SearchIndexAclToken", + "SearchIndexChangeQueue", + "SearchIndexDocument", + "SearchIndexState", + "new_uuid", +] diff --git a/src/govoplan_search/backend/manifest.py b/src/govoplan_search/backend/manifest.py index 2fb294c..c8649d8 100644 --- a/src/govoplan_search/backend/manifest.py +++ b/src/govoplan_search/backend/manifest.py @@ -118,7 +118,8 @@ manifest = ModuleManifest( ), provides_interfaces=( ModuleInterfaceProvider(name="search.provider", version="1.0.0"), - ModuleInterfaceProvider(name="search.index_writer", version="1.0.0"), + ModuleInterfaceProvider(name="search.index_writer", version="1.1.0"), + ModuleInterfaceProvider(name="search.source", version="1.0.0"), ), permissions=PERMISSIONS, role_templates=ROLE_TEMPLATES, @@ -158,8 +159,10 @@ manifest = ModuleManifest( script_location=str(Path(__file__).with_name("migrations") / "versions"), retirement_supported=True, retirement_provider=drop_table_retirement_provider( + search_models.SearchIndexChangeQueue, search_models.SearchIndexAclToken, search_models.SearchIndexDocument, + search_models.SearchIndexState, label="Search", ), retirement_notes=( @@ -169,7 +172,9 @@ manifest = ModuleManifest( ), uninstall_guard_providers=( persistent_table_uninstall_guard( + search_models.SearchIndexChangeQueue, search_models.SearchIndexDocument, + search_models.SearchIndexState, label="Search index", ), ), diff --git a/src/govoplan_search/backend/migrations/versions/b2c3d4e5f607_v0114_search_lifecycle.py b/src/govoplan_search/backend/migrations/versions/b2c3d4e5f607_v0114_search_lifecycle.py new file mode 100644 index 0000000..97e05d0 --- /dev/null +++ b/src/govoplan_search/backend/migrations/versions/b2c3d4e5f607_v0114_search_lifecycle.py @@ -0,0 +1,311 @@ +"""Add versioned search indexing lifecycle. + +Revision ID: b2c3d4e5f607 +Revises: a1b2c3d4e5f6 +Create Date: 2026-07-29 +""" + +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + + +revision = "b2c3d4e5f607" +down_revision = "a1b2c3d4e5f6" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + for column in ( + sa.Column( + "provider_id", + sa.String(length=200), + nullable=False, + server_default="legacy.index", + ), + sa.Column( + "source_revision", + sa.String(length=255), + nullable=False, + server_default="1", + ), + sa.Column( + "change_cursor", + sa.String(length=500), + nullable=True, + ), + sa.Column( + "source_updated_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column( + "language", + sa.String(length=32), + nullable=False, + server_default="simple", + ), + sa.Column( + "index_version", + sa.Integer(), + nullable=False, + server_default="1", + ), + sa.Column( + "requires_authorization_recheck", + sa.Boolean(), + nullable=False, + server_default=sa.false(), + ), + sa.Column( + "rebuild_id", + sa.String(length=36), + nullable=True, + ), + sa.Column( + "active", + sa.Boolean(), + nullable=False, + server_default=sa.true(), + ), + ): + op.add_column("search_index_documents", column) + for column in ("provider_id", "rebuild_id", "active"): + op.create_index( + op.f(f"ix_search_index_documents_{column}"), + "search_index_documents", + [column], + ) + op.create_index( + "ix_search_document_provider", + "search_index_documents", + [ + "tenant_id", + "provider_id", + "resource_type", + "active", + ], + ) + bind = op.get_bind() + if ( + bind.dialect.name == "postgresql" + and bind.scalar( + sa.text( + "SELECT EXISTS (" + "SELECT 1 FROM pg_extension " + "WHERE extname = 'pg_trgm'" + ")" + ) + ) + ): + op.execute( + "CREATE INDEX ix_search_document_title_trgm " + "ON search_index_documents USING gin " + "(lower(title) gin_trgm_ops)" + ) + op.execute( + "CREATE INDEX ix_search_document_text_trgm " + "ON search_index_documents USING gin " + "(lower(search_text) gin_trgm_ops)" + ) + + op.create_table( + "search_index_states", + sa.Column("id", sa.String(length=36), nullable=False), + sa.Column("tenant_id", sa.String(length=36), nullable=False), + sa.Column("provider_id", sa.String(length=200), nullable=False), + sa.Column("module_id", sa.String(length=100), nullable=False), + sa.Column("resource_type", sa.String(length=100), nullable=False), + sa.Column("index_version", sa.Integer(), nullable=False), + sa.Column("status", sa.String(length=30), nullable=False), + sa.Column("rebuild_id", sa.String(length=36), nullable=True), + sa.Column( + "checkpoint_cursor", + sa.String(length=500), + nullable=True, + ), + sa.Column( + "high_watermark", + sa.String(length=500), + nullable=True, + ), + sa.Column( + "last_change_cursor", + sa.String(length=500), + nullable=True, + ), + sa.Column("indexed_documents", sa.Integer(), nullable=False), + sa.Column("rejected_documents", sa.Integer(), nullable=False), + sa.Column( + "rebuild_started_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column( + "rebuild_completed_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column( + "last_success_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column("last_error", sa.Text(), nullable=True), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + nullable=False, + ), + sa.Column( + "updated_at", + sa.DateTime(timezone=True), + nullable=False, + ), + sa.PrimaryKeyConstraint( + "id", + name=op.f("pk_search_index_states"), + ), + sa.UniqueConstraint( + "tenant_id", + "provider_id", + "resource_type", + name="uq_search_index_state_source", + ), + ) + for column in ( + "tenant_id", + "provider_id", + "module_id", + "resource_type", + "status", + ): + op.create_index( + op.f(f"ix_search_index_states_{column}"), + "search_index_states", + [column], + ) + op.create_index( + "ix_search_index_state_status", + "search_index_states", + ["tenant_id", "status"], + ) + + op.create_table( + "search_index_change_queue", + sa.Column("id", sa.String(length=36), nullable=False), + sa.Column("change_id", sa.String(length=255), nullable=False), + sa.Column("tenant_id", sa.String(length=36), nullable=False), + sa.Column("provider_id", sa.String(length=200), nullable=False), + sa.Column("module_id", sa.String(length=100), nullable=False), + sa.Column("resource_type", sa.String(length=100), nullable=False), + sa.Column("resource_id", sa.String(length=255), nullable=False), + sa.Column("kind", sa.String(length=20), nullable=False), + sa.Column( + "source_revision", + sa.String(length=255), + nullable=False, + ), + sa.Column( + "source_cursor", + sa.String(length=500), + nullable=False, + ), + sa.Column("document", sa.JSON(), nullable=True), + sa.Column( + "occurred_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column("status", sa.String(length=30), nullable=False), + sa.Column("attempts", sa.Integer(), nullable=False), + sa.Column( + "available_at", + sa.DateTime(timezone=True), + nullable=False, + ), + sa.Column( + "processed_at", + sa.DateTime(timezone=True), + nullable=True, + ), + sa.Column("error", sa.Text(), nullable=True), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + nullable=False, + ), + sa.Column( + "updated_at", + sa.DateTime(timezone=True), + nullable=False, + ), + sa.PrimaryKeyConstraint( + "id", + name=op.f("pk_search_index_change_queue"), + ), + sa.UniqueConstraint( + "change_id", + name="uq_search_index_change_queue_change", + ), + ) + for column in ( + "change_id", + "tenant_id", + "provider_id", + "module_id", + "resource_type", + "resource_id", + "status", + "available_at", + ): + op.create_index( + op.f(f"ix_search_index_change_queue_{column}"), + "search_index_change_queue", + [column], + ) + op.create_index( + "ix_search_index_change_queue_pending", + "search_index_change_queue", + ["status", "available_at", "created_at"], + ) + op.create_index( + "ix_search_index_change_queue_source", + "search_index_change_queue", + ["tenant_id", "provider_id", "resource_type"], + ) + + +def downgrade() -> None: + if op.get_bind().dialect.name == "postgresql": + op.execute( + "DROP INDEX IF EXISTS ix_search_document_text_trgm" + ) + op.execute( + "DROP INDEX IF EXISTS ix_search_document_title_trgm" + ) + op.drop_table("search_index_change_queue") + op.drop_table("search_index_states") + op.drop_index( + "ix_search_document_provider", + table_name="search_index_documents", + ) + for column in ("active", "rebuild_id", "provider_id"): + op.drop_index( + op.f(f"ix_search_index_documents_{column}"), + table_name="search_index_documents", + ) + for column in ( + "active", + "rebuild_id", + "requires_authorization_recheck", + "index_version", + "language", + "source_updated_at", + "change_cursor", + "source_revision", + "provider_id", + ): + op.drop_column("search_index_documents", column) diff --git a/src/govoplan_search/backend/router.py b/src/govoplan_search/backend/router.py index 6a964c3..d428048 100644 --- a/src/govoplan_search/backend/router.py +++ b/src/govoplan_search/backend/router.py @@ -4,17 +4,27 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from sqlalchemy.orm import Session from govoplan_core.auth import ApiPrincipal, get_api_principal, has_scope +from govoplan_core.audit.logging import audit_from_principal +from govoplan_core.core.search import CAPABILITY_SEARCH_INDEX_WRITER from govoplan_core.core.registry import PlatformRegistry from govoplan_core.core.search import SearchQuery from govoplan_core.db.session import get_session -from govoplan_search.backend.manifest import READ_SCOPE +from govoplan_search.backend.manifest import ADMIN_SCOPE, READ_SCOPE from govoplan_search.backend.schemas import ( + SearchChangeDispatchResponse, + SearchDiagnosticsResponse, + SearchIndexStateResponse, + SearchModuleReconcileResponse, SearchProviderListResponse, SearchProviderResponse, + SearchRebuildResponse, SearchResponse, SearchResultResponse, ) -from govoplan_search.backend.service import aggregate_search +from govoplan_search.backend.service import ( + SearchIndexService, + aggregate_search_page, +) router = APIRouter(prefix="/search", tags=["search"]) @@ -38,6 +48,26 @@ def _require_read(principal: ApiPrincipal) -> None: ) +def _require_admin(principal: ApiPrincipal) -> None: + if not has_scope(principal, ADMIN_SCOPE): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"Missing scope: {ADMIN_SCOPE}", + ) + + +def _service(registry: PlatformRegistry) -> SearchIndexService: + capability = registry.capability( + CAPABILITY_SEARCH_INDEX_WRITER + ) + if not isinstance(capability, SearchIndexService): + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Search index service is not available.", + ) + return capability + + @router.get("", response_model=SearchResponse) def api_search( request: Request, @@ -48,27 +78,41 @@ def api_search( context_id: str | None = Query(default=None, max_length=255), limit: int = Query(default=25, ge=1, le=100), offset: int = Query(default=0, ge=0, le=10_000), + cursor: str | None = Query(default=None, max_length=2000), + language: str = Query( + default="simple", + max_length=32, + pattern=r"^[A-Za-z_]+$", + ), session: Session = Depends(get_session), principal: ApiPrincipal = Depends(get_api_principal), ) -> SearchResponse: _require_read(principal) registry = _registry(request) - query = SearchQuery( - text=q, - tenant_id=principal.tenant_id, - module_ids=tuple(dict.fromkeys(module)), - resource_types=tuple(dict.fromkeys(resource_type)), - context_kind=context_kind, # type: ignore[arg-type] - context_id=context_id, - limit=limit, - offset=offset, - ) - results, diagnostics = aggregate_search( - registry, - session, - principal, - query=query, - ) + try: + query = SearchQuery( + text=q, + tenant_id=principal.tenant_id, + module_ids=tuple(dict.fromkeys(module)), + resource_types=tuple(dict.fromkeys(resource_type)), + context_kind=context_kind, # type: ignore[arg-type] + context_id=context_id, + limit=limit, + offset=offset, + cursor=cursor, + language=language.casefold(), + ) + page = aggregate_search_page( + registry, + session, + principal, + query=query, + ) + except ValueError as exc: + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, + detail=str(exc), + ) from exc return SearchResponse( query=query.text, results=[ @@ -90,11 +134,15 @@ def api_search( else None ), "metadata": dict(result.metadata), + "source_revision": result.source_revision, + "provenance": dict(result.provenance), } ) - for result in results + for result in page.results ], - diagnostics=list(diagnostics), + diagnostics=list(page.diagnostics), + next_cursor=page.next_cursor, + has_more=page.next_cursor is not None, ) @@ -118,4 +166,190 @@ def api_search_providers( ) +@router.get( + "/admin/diagnostics", + response_model=SearchDiagnosticsResponse, +) +def api_search_diagnostics( + request: Request, + session: Session = Depends(get_session), + principal: ApiPrincipal = Depends(get_api_principal), +) -> SearchDiagnosticsResponse: + _require_admin(principal) + registry = _registry(request) + return SearchDiagnosticsResponse.model_validate( + _service(registry).diagnostics(session, principal) + ) + + +@router.post( + "/admin/reconcile-modules", + response_model=SearchModuleReconcileResponse, +) +def api_reconcile_search_modules( + request: Request, + session: Session = Depends(get_session), + principal: ApiPrincipal = Depends(get_api_principal), +) -> SearchModuleReconcileResponse: + _require_admin(principal) + result = _service(_registry(request)).reconcile_active_modules( + session, + tenant_id=principal.tenant_id, + ) + audit_from_principal( + session, + principal, + action="search.modules.reconciled", + scope="tenant", + object_type="search_index", + details=result, + ) + session.commit() + return SearchModuleReconcileResponse(**result) + + +@router.post( + "/admin/changes/process", + response_model=SearchChangeDispatchResponse, +) +def api_process_search_changes( + request: Request, + limit: int = Query(default=100, ge=1, le=500), + session: Session = Depends(get_session), + principal: ApiPrincipal = Depends(get_api_principal), +) -> SearchChangeDispatchResponse: + _require_admin(principal) + result = _service(_registry(request)).process_changes( + session, + limit=limit, + tenant_id=principal.tenant_id, + ) + audit_from_principal( + session, + principal, + action="search.changes.processed", + scope="tenant", + object_type="search_index", + details=result, + ) + session.commit() + return SearchChangeDispatchResponse(**result) + + +@router.post( + "/admin/rebuilds/{provider_id}/{resource_type}/start", + response_model=SearchRebuildResponse, +) +def api_start_search_rebuild( + provider_id: str, + resource_type: str, + request: Request, + session: Session = Depends(get_session), + principal: ApiPrincipal = Depends(get_api_principal), +) -> SearchRebuildResponse: + _require_admin(principal) + try: + state = _service(_registry(request)).start_rebuild( + session, + principal, + provider_id=provider_id, + resource_type=resource_type, + ) + except ValueError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=str(exc), + ) from exc + audit_from_principal( + session, + principal, + action="search.rebuild.started", + scope="tenant", + object_type="search_index", + object_id=state.id, + details={ + "provider_id": provider_id, + "resource_type": resource_type, + "rebuild_id": state.rebuild_id, + }, + ) + session.commit() + return SearchRebuildResponse( + state=_search_state_response(state) + ) + + +@router.post( + "/admin/rebuilds/{provider_id}/{resource_type}/continue", + response_model=SearchRebuildResponse, +) +def api_continue_search_rebuild( + provider_id: str, + resource_type: str, + request: Request, + limit: int = Query(default=100, ge=1, le=500), + session: Session = Depends(get_session), + principal: ApiPrincipal = Depends(get_api_principal), +) -> SearchRebuildResponse: + _require_admin(principal) + service = _service(_registry(request)) + try: + state = service.continue_rebuild( + session, + principal, + provider_id=provider_id, + resource_type=resource_type, + limit=limit, + ) + except ValueError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=str(exc), + ) from exc + audit_from_principal( + session, + principal, + action="search.rebuild.continued", + scope="tenant", + object_type="search_index", + object_id=state.id, + details={ + "provider_id": provider_id, + "resource_type": resource_type, + "status": state.status, + "indexed_documents": state.indexed_documents, + }, + ) + session.commit() + return SearchRebuildResponse( + state=_search_state_response(state) + ) + + +def _search_state_response(state: object) -> SearchIndexStateResponse: + return SearchIndexStateResponse( + provider_id=str(getattr(state, "provider_id")), + module_id=str(getattr(state, "module_id")), + resource_type=str(getattr(state, "resource_type")), + index_version=int(getattr(state, "index_version")), + status=str(getattr(state, "status")), + checkpoint_cursor=getattr(state, "checkpoint_cursor"), + high_watermark=getattr(state, "high_watermark"), + last_change_cursor=getattr(state, "last_change_cursor"), + indexed_documents=int( + getattr(state, "indexed_documents") + ), + rejected_documents=int( + getattr(state, "rejected_documents") + ), + rebuild_started_at=getattr(state, "rebuild_started_at"), + rebuild_completed_at=getattr( + state, + "rebuild_completed_at", + ), + last_success_at=getattr(state, "last_success_at"), + last_error=getattr(state, "last_error"), + ) + + __all__ = ["router"] diff --git a/src/govoplan_search/backend/schemas.py b/src/govoplan_search/backend/schemas.py index 7fa1dc9..37db624 100644 --- a/src/govoplan_search/backend/schemas.py +++ b/src/govoplan_search/backend/schemas.py @@ -1,5 +1,6 @@ from __future__ import annotations +from datetime import datetime from typing import Any, Literal from pydantic import BaseModel, Field @@ -31,6 +32,8 @@ class SearchResultResponse(BaseModel): breadcrumbs: list[str] = Field(default_factory=list) external_reference: SearchExternalReferenceResponse | None = None metadata: dict[str, Any] = Field(default_factory=dict) + source_revision: str | None = None + provenance: dict[str, Any] = Field(default_factory=dict) class SearchProviderDiagnosticResponse(BaseModel): @@ -43,6 +46,8 @@ class SearchResponse(BaseModel): query: str results: list[SearchResultResponse] diagnostics: list[SearchProviderDiagnosticResponse] = Field(default_factory=list) + next_cursor: str | None = None + has_more: bool = False class SearchProviderResponse(BaseModel): @@ -56,10 +61,58 @@ class SearchProviderListResponse(BaseModel): providers: list[SearchProviderResponse] +class SearchIndexStateResponse(BaseModel): + provider_id: str + module_id: str + resource_type: str + index_version: int + status: str + checkpoint_cursor: str | None = None + high_watermark: str | None = None + last_change_cursor: str | None = None + indexed_documents: int = 0 + rejected_documents: int = 0 + rebuild_started_at: datetime | None = None + rebuild_completed_at: datetime | None = None + last_success_at: datetime | None = None + last_error: str | None = None + + +class SearchDiagnosticsResponse(BaseModel): + backend: str + trigram_available: bool + queue: dict[str, int] = Field(default_factory=dict) + queue_oldest_age_seconds: float | None = None + states: list[SearchIndexStateResponse] = Field( + default_factory=list + ) + + +class SearchRebuildResponse(BaseModel): + state: SearchIndexStateResponse + + +class SearchChangeDispatchResponse(BaseModel): + selected: int + applied: int + retrying: int + quarantined: int + + +class SearchModuleReconcileResponse(BaseModel): + disabled_documents: int + enabled_documents: int + + __all__ = [ "SearchProviderDiagnosticResponse", "SearchProviderListResponse", "SearchProviderResponse", + "SearchChangeDispatchResponse", + "SearchDiagnosticsResponse", + "SearchIndexStateResponse", + "SearchModuleReconcileResponse", + "SearchRebuildResponse", "SearchResponse", "SearchResultResponse", ] diff --git a/src/govoplan_search/backend/service.py b/src/govoplan_search/backend/service.py index 41d2733..fcf31e1 100644 --- a/src/govoplan_search/backend/service.py +++ b/src/govoplan_search/backend/service.py @@ -1,29 +1,96 @@ from __future__ import annotations +import base64 +from collections import defaultdict +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, replace +from datetime import datetime, timedelta, timezone import hashlib import json import logging -from dataclasses import replace -from datetime import datetime, timezone +import uuid from typing import Any -from sqlalchemy import delete, exists, func, literal, or_, select +from sqlalchemy import ( + cast, + delete, + exists, + func, + literal, + or_, + select, + text, + update, +) +from sqlalchemy.dialects.postgresql import REGCONFIG +from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from govoplan_core.core.external_references import ExternalObjectReference -from govoplan_core.core.search import SearchDocument, SearchQuery, SearchResult +from govoplan_core.core.search import ( + SearchAuthorizationRequest, + SearchBackfillRequest, + SearchDocument, + SearchIndexChange, + SearchQuery, + SearchResourceReference, + SearchResourceType, + SearchResult, + SearchSourceProvider, +) +from govoplan_core.db.base import utcnow from govoplan_search.backend.db.models import ( SearchIndexAclToken, + SearchIndexChangeQueue, SearchIndexDocument, + SearchIndexState, ) LOGGER = logging.getLogger(__name__) +MAX_QUERY_CANDIDATES = 500 +MAX_AGGREGATE_WINDOW = 200 +MAX_CHANGE_ATTEMPTS = 8 +SUPPORTED_SEARCH_CONFIGS = frozenset( + { + "simple", + "danish", + "dutch", + "english", + "finnish", + "french", + "german", + "hungarian", + "italian", + "norwegian", + "portuguese", + "romanian", + "russian", + "spanish", + "swedish", + "turkish", + } +) + + +@dataclass(frozen=True, slots=True) +class SearchPage: + results: tuple[SearchResult, ...] + next_cursor: str | None = None + scanned_candidates: int = 0 + + +@dataclass(frozen=True, slots=True) +class AggregatedSearchPage: + results: tuple[SearchResult, ...] + diagnostics: tuple[dict[str, str], ...] + next_cursor: str | None class SearchIndexService: def __init__(self, registry: object | None = None) -> None: self.registry = registry + self._trigram_available: bool | None = None def upsert_document( self, @@ -34,11 +101,17 @@ class SearchIndexService: ) -> None: db = _session(session) _require_tenant(principal, document.tenant_id) - normalized_tokens = tuple( - dict.fromkeys(token.strip() for token in document.acl_tokens if token.strip()) - ) - if document.visibility == "restricted" and not normalized_tokens: - raise ValueError("Restricted search documents require ACL tokens.") + self._upsert_document(db, document=document) + + def _upsert_document( + self, + db: Session, + *, + document: SearchDocument, + rebuild_id: str | None = None, + ) -> SearchIndexDocument: + _assert_safe_index_metadata(document) + normalized_tokens = _normalized_acl_tokens(document) identity = ( document.tenant_id, document.module_id, @@ -53,7 +126,10 @@ class SearchIndexService: SearchIndexDocument.resource_id == identity[3], ) ) - payload = _document_payload(document) + payload = _document_payload( + document, + rebuild_id=rebuild_id, + ) if model is None: model = SearchIndexDocument(**payload) db.add(model) @@ -67,7 +143,13 @@ class SearchIndexService: ) ) for token in normalized_tokens: - db.add(SearchIndexAclToken(document_id=model.id, token=token)) + db.add( + SearchIndexAclToken( + document_id=model.id, + token=token, + ) + ) + return model def delete_document( self, @@ -81,12 +163,32 @@ class SearchIndexService: ) -> bool: db = _session(session) _require_tenant(principal, tenant_id) + return self._delete_reference( + db, + reference=SearchResourceReference( + tenant_id=tenant_id, + module_id=module_id, + resource_type=resource_type, + resource_id=resource_id, + ), + ) + + def _delete_reference( + self, + db: Session, + *, + reference: SearchResourceReference, + ) -> bool: model = db.scalar( select(SearchIndexDocument).where( - SearchIndexDocument.tenant_id == tenant_id, - SearchIndexDocument.module_id == module_id, - SearchIndexDocument.resource_type == resource_type, - SearchIndexDocument.resource_id == resource_id, + SearchIndexDocument.tenant_id + == reference.tenant_id, + SearchIndexDocument.module_id + == reference.module_id, + SearchIndexDocument.resource_type + == reference.resource_type, + SearchIndexDocument.resource_id + == reference.resource_id, ) ) if model is None: @@ -94,6 +196,429 @@ class SearchIndexService: db.delete(model) return True + def enqueue_change( + self, + session: object, + *, + change: SearchIndexChange, + ) -> bool: + db = _session(session) + queued = SearchIndexChangeQueue( + change_id=change.change_id, + tenant_id=change.reference.tenant_id, + provider_id=change.provider_id, + module_id=change.reference.module_id, + resource_type=change.reference.resource_type, + resource_id=change.reference.resource_id, + kind=change.kind, + source_revision=change.source_revision, + source_cursor=change.cursor, + document_=( + _document_contract_payload(change.document) + if change.document is not None + else None + ), + occurred_at=change.occurred_at, + status="queued", + attempts=0, + available_at=utcnow(), + ) + try: + with db.begin_nested(): + db.add(queued) + db.flush() + except IntegrityError: + return False + return True + + def process_changes( + self, + session: object, + *, + limit: int = 100, + now: datetime | None = None, + tenant_id: str | None = None, + ) -> dict[str, int]: + db = _session(session) + current = _as_utc(now or utcnow()) + statement = select(SearchIndexChangeQueue).where( + SearchIndexChangeQueue.status.in_( + ("queued", "retrying") + ), + SearchIndexChangeQueue.available_at <= current, + ) + if tenant_id is not None: + statement = statement.where( + SearchIndexChangeQueue.tenant_id == tenant_id + ) + rows = list( + db.scalars( + statement + .order_by( + SearchIndexChangeQueue.created_at, + SearchIndexChangeQueue.id, + ) + .limit(max(1, min(int(limit), 500))) + .with_for_update(skip_locked=True) + ) + ) + result = { + "selected": len(rows), + "applied": 0, + "retrying": 0, + "quarantined": 0, + } + for row in rows: + state = self._state( + db, + tenant_id=row.tenant_id, + provider_id=row.provider_id, + resource_type=row.resource_type, + ) + if state is not None and state.status == "backfilling": + row.status = "retrying" + row.available_at = current + timedelta(seconds=30) + result["retrying"] += 1 + continue + row.attempts += 1 + try: + with db.begin_nested(): + self._apply_queued_change(db, row) + db.flush() + except Exception as exc: + LOGGER.exception( + "Search index change failed: %s", + row.change_id, + ) + row.error = _bounded_error(exc) + if row.attempts >= MAX_CHANGE_ATTEMPTS: + row.status = "quarantined" + row.processed_at = current + result["quarantined"] += 1 + else: + row.status = "retrying" + row.available_at = current + timedelta( + seconds=min(3600, 2 ** row.attempts) + ) + result["retrying"] += 1 + continue + row.status = "applied" + row.error = None + row.processed_at = current + result["applied"] += 1 + db.flush() + return result + + def _apply_queued_change( + self, + db: Session, + row: SearchIndexChangeQueue, + ) -> None: + reference = SearchResourceReference( + tenant_id=row.tenant_id, + module_id=row.module_id, + resource_type=row.resource_type, + resource_id=row.resource_id, + ) + if row.kind == "upsert": + if row.document_ is None: + raise ValueError( + "Queued search upsert has no document." + ) + document = _document_from_contract_payload( + row.document_ + ) + if ( + document.reference != reference + or document.provider_id != row.provider_id + or document.source_revision != row.source_revision + ): + raise ValueError( + "Queued search document does not match its envelope." + ) + self._upsert_document(db, document=document) + elif row.kind == "delete": + self._delete_reference(db, reference=reference) + else: + raise ValueError( + f"Unsupported search change kind: {row.kind}" + ) + state = self._state( + db, + tenant_id=row.tenant_id, + provider_id=row.provider_id, + resource_type=row.resource_type, + ) + if state is not None: + state.last_change_cursor = row.source_cursor + state.last_success_at = utcnow() + state.last_error = None + + def start_rebuild( + self, + session: object, + principal: object, + *, + provider_id: str, + resource_type: str, + ) -> SearchIndexState: + db = _session(session) + tenant_id = _principal_tenant(principal) + descriptor, _provider = self._source( + provider_id, + resource_type, + ) + state = self._state( + db, + tenant_id=tenant_id, + provider_id=provider_id, + resource_type=resource_type, + lock=True, + ) + if state is None: + state = SearchIndexState( + tenant_id=tenant_id, + provider_id=provider_id, + module_id=descriptor.module_id, + resource_type=resource_type, + index_version=descriptor.index_version, + ) + db.add(state) + state.module_id = descriptor.module_id + state.index_version = descriptor.index_version + state.status = "backfilling" + state.rebuild_id = str(uuid.uuid4()) + state.checkpoint_cursor = None + state.high_watermark = None + state.indexed_documents = 0 + state.rejected_documents = 0 + state.rebuild_started_at = utcnow() + state.rebuild_completed_at = None + state.last_error = None + db.flush() + return state + + def continue_rebuild( + self, + session: object, + principal: object, + *, + provider_id: str, + resource_type: str, + limit: int = 100, + ) -> SearchIndexState: + db = _session(session) + tenant_id = _principal_tenant(principal) + descriptor, provider = self._source( + provider_id, + resource_type, + ) + state = self._state( + db, + tenant_id=tenant_id, + provider_id=provider_id, + resource_type=resource_type, + lock=True, + ) + if state is None or state.status not in { + "backfilling", + "failed", + }: + raise ValueError( + "Search rebuild has not been started." + ) + if not state.rebuild_id: + raise ValueError("Search rebuild generation is missing.") + state.status = "backfilling" + try: + page = provider.backfill( + db, + request=SearchBackfillRequest( + tenant_id=tenant_id, + provider_id=provider_id, + resource_type=resource_type, + rebuild_id=state.rebuild_id, + cursor=state.checkpoint_cursor, + limit=max(1, min(int(limit), 500)), + ), + ) + accepted = 0 + for document in page.documents: + try: + _validate_backfill_document( + document, + descriptor=descriptor, + tenant_id=tenant_id, + ) + except ValueError: + state.rejected_documents += 1 + raise + self._upsert_document( + db, + document=document, + rebuild_id=state.rebuild_id, + ) + accepted += 1 + state.indexed_documents += accepted + state.checkpoint_cursor = page.next_cursor + state.high_watermark = ( + page.high_watermark or state.high_watermark + ) + state.last_success_at = utcnow() + state.last_error = None + if page.complete: + db.execute( + delete(SearchIndexDocument).where( + SearchIndexDocument.tenant_id == tenant_id, + SearchIndexDocument.provider_id == provider_id, + SearchIndexDocument.resource_type + == resource_type, + or_( + SearchIndexDocument.rebuild_id.is_(None), + SearchIndexDocument.rebuild_id + != state.rebuild_id, + ), + ) + ) + state.status = "ready" + state.rebuild_completed_at = utcnow() + state.checkpoint_cursor = None + except Exception as exc: + state.status = "failed" + state.last_error = _bounded_error(exc) + db.flush() + raise + db.flush() + return state + + def reconcile_active_modules( + self, + session: object, + *, + tenant_id: str | None = None, + ) -> dict[str, int]: + db = _session(session) + active_modules = set(_active_module_ids(self.registry)) + if not active_modules: + return {"disabled_documents": 0, "enabled_documents": 0} + document_scope = ( + SearchIndexDocument.tenant_id == tenant_id + if tenant_id is not None + else literal(True) + ) + state_scope = ( + SearchIndexState.tenant_id == tenant_id + if tenant_id is not None + else literal(True) + ) + disabled = db.execute( + update(SearchIndexDocument) + .where( + document_scope, + SearchIndexDocument.module_id.not_in(active_modules), + SearchIndexDocument.active.is_(True), + ) + .values(active=False) + ).rowcount + enabled = db.execute( + update(SearchIndexDocument) + .where( + document_scope, + SearchIndexDocument.module_id.in_(active_modules), + SearchIndexDocument.active.is_(False), + ) + .values(active=True) + ).rowcount + db.execute( + update(SearchIndexState) + .where( + state_scope, + SearchIndexState.module_id.not_in(active_modules), + ) + .values(status="disabled") + ) + db.execute( + update(SearchIndexState) + .where( + state_scope, + SearchIndexState.module_id.in_(active_modules), + SearchIndexState.status == "disabled", + ) + .values(status="idle") + ) + return { + "disabled_documents": int(disabled or 0), + "enabled_documents": int(enabled or 0), + } + + def diagnostics( + self, + session: object, + principal: object, + ) -> dict[str, object]: + db = _session(session) + tenant_id = _principal_tenant(principal) + states = list( + db.scalars( + select(SearchIndexState) + .where(SearchIndexState.tenant_id == tenant_id) + .order_by( + SearchIndexState.provider_id, + SearchIndexState.resource_type, + ) + ) + ) + queue_counts = dict( + db.execute( + select( + SearchIndexChangeQueue.status, + func.count(SearchIndexChangeQueue.id), + ) + .where( + SearchIndexChangeQueue.tenant_id == tenant_id + ) + .group_by(SearchIndexChangeQueue.status) + ).all() + ) + oldest_pending = db.scalar( + select( + func.min(SearchIndexChangeQueue.created_at) + ).where( + SearchIndexChangeQueue.tenant_id == tenant_id, + SearchIndexChangeQueue.status.in_( + ("queued", "retrying") + ), + ) + ) + oldest_age = ( + max( + 0.0, + ( + _as_utc(utcnow()) + - _as_utc(oldest_pending) + ).total_seconds(), + ) + if oldest_pending is not None + else None + ) + return { + "backend": db.get_bind().dialect.name, + "trigram_available": ( + self._postgres_trigram_available(db) + if db.get_bind().dialect.name == "postgresql" + else False + ), + "queue": { + str(status): int(count) + for status, count in queue_counts.items() + }, + "queue_oldest_age_seconds": oldest_age, + "states": [ + _state_payload(state) + for state in states + ], + } + def search( self, session: object, @@ -101,23 +626,40 @@ class SearchIndexService: *, query: SearchQuery, ) -> tuple[SearchResult, ...]: + return self.search_page( + session, + principal, + query=query, + ).results + + def search_page( + self, + session: object, + principal: object, + *, + query: SearchQuery, + ) -> SearchPage: db = _session(session) _require_tenant(principal, query.tenant_id) if not query.text: - return () + return SearchPage(results=()) active_modules = _active_module_ids(self.registry) if active_modules == (): - return () + return SearchPage(results=()) tokens = principal_acl_tokens(principal) acl_match = exists( select(SearchIndexAclToken.id).where( - SearchIndexAclToken.document_id == SearchIndexDocument.id, - SearchIndexAclToken.token.in_(tokens or ("__no_acl_token__",)), + SearchIndexAclToken.document_id + == SearchIndexDocument.id, + SearchIndexAclToken.token.in_( + tokens or ("__no_acl_token__",) + ), ) ) filters = [ SearchIndexDocument.tenant_id == query.tenant_id, + SearchIndexDocument.active.is_(True), SearchIndexDocument.module_id.in_(active_modules), or_( SearchIndexDocument.visibility == "tenant", @@ -125,52 +667,270 @@ class SearchIndexService: ), ] if query.module_ids: - filters.append(SearchIndexDocument.module_id.in_(query.module_ids)) + filters.append( + SearchIndexDocument.module_id.in_( + query.module_ids + ) + ) if query.resource_types: filters.append( - SearchIndexDocument.resource_type.in_(query.resource_types) + SearchIndexDocument.resource_type.in_( + query.resource_types + ) ) dialect = db.get_bind().dialect.name if dialect == "postgresql": - search_query = func.websearch_to_tsquery("simple", query.text) - vector = func.to_tsvector("simple", SearchIndexDocument.search_text) - rank = func.ts_rank_cd(vector, search_query) - statement = ( - select(SearchIndexDocument, rank.label("search_rank")) - .where(*filters, vector.op("@@")(search_query)) - .order_by(rank.desc(), SearchIndexDocument.title.asc()) + statement = self._postgres_statement( + db, + query=query, + filters=filters, ) else: - lowered = func.lower(SearchIndexDocument.search_text) + lowered = func.lower( + SearchIndexDocument.search_text + ) text_filters = [ lowered.contains(token.casefold()) for token in query.text.split() if token.strip() ] statement = ( - select(SearchIndexDocument, literal(1.0).label("search_rank")) + select( + SearchIndexDocument, + literal(1.0).label("search_rank"), + ) .where(*filters, *text_filters) - .order_by(SearchIndexDocument.title.asc()) + .order_by( + func.lower(SearchIndexDocument.title), + SearchIndexDocument.module_id, + SearchIndexDocument.resource_type, + SearchIndexDocument.resource_id, + ) ) + candidate_limit = min( + MAX_QUERY_CANDIDATES, + max(50, (query.limit + query.offset + 1) * 5), + ) rows = db.execute( - statement.limit(query.limit).offset(query.offset) + statement.limit(candidate_limit) ).all() - return tuple( - _search_result(model, float(rank or 0.0)) - for model, rank in rows + authorized = self._authorized_results( + db, + principal, + rows, + ) + start = query.offset + selected = authorized[start : start + query.limit] + return SearchPage( + results=tuple(selected), + scanned_candidates=len(rows), ) + def _postgres_statement( + self, + db: Session, + *, + query: SearchQuery, + filters: list[object], + ): + config = ( + query.language + if query.language in SUPPORTED_SEARCH_CONFIGS + else "simple" + ) + search_config = cast(literal(config), REGCONFIG) + search_query = func.websearch_to_tsquery( + search_config, + query.text, + ) + vector = func.to_tsvector( + search_config, + SearchIndexDocument.search_text, + ) + fts_rank = func.ts_rank_cd(vector, search_query) + if self._postgres_trigram_available(db): + similarity = func.greatest( + func.similarity( + func.lower(SearchIndexDocument.title), + query.text.casefold(), + ), + func.similarity( + func.lower(SearchIndexDocument.search_text), + query.text.casefold(), + ), + ) + rank = fts_rank + similarity * 0.2 + text_filter = or_( + vector.op("@@")(search_query), + func.lower(SearchIndexDocument.title).op("%")( + query.text.casefold() + ), + func.lower(SearchIndexDocument.search_text).op("%")( + query.text.casefold() + ), + ) + else: + rank = fts_rank + text_filter = vector.op("@@")(search_query) + return ( + select( + SearchIndexDocument, + rank.label("search_rank"), + ) + .where(*filters, text_filter) + .order_by( + rank.desc(), + func.lower(SearchIndexDocument.title), + SearchIndexDocument.module_id, + SearchIndexDocument.resource_type, + SearchIndexDocument.resource_id, + ) + ) -def aggregate_search( + def _postgres_trigram_available( + self, + db: Session, + ) -> bool: + if self._trigram_available is None: + try: + self._trigram_available = bool( + db.scalar( + text( + "SELECT EXISTS (" + "SELECT 1 FROM pg_extension " + "WHERE extname = 'pg_trgm'" + ")" + ) + ) + ) + except Exception: + LOGGER.exception( + "Could not inspect PostgreSQL trigram support" + ) + self._trigram_available = False + return self._trigram_available + + def _authorized_results( + self, + db: Session, + principal: object, + rows: Sequence[tuple[SearchIndexDocument, object]], + ) -> list[SearchResult]: + decisions: dict[str, bool] = { + document.id: not document.requires_authorization_recheck + for document, _rank in rows + } + grouped: dict[ + str, + list[tuple[SearchIndexDocument, SearchAuthorizationRequest]], + ] = defaultdict(list) + for document, _rank in rows: + if not document.requires_authorization_recheck: + continue + request = SearchAuthorizationRequest( + reference=SearchResourceReference( + tenant_id=document.tenant_id, + module_id=document.module_id, + resource_type=document.resource_type, + resource_id=document.resource_id, + ), + source_revision=document.source_revision, + ) + grouped[document.provider_id].append( + (document, request) + ) + sources = { + registered.registration.id: provider + for registered, provider in _search_sources(self.registry) + } + for provider_id, items in grouped.items(): + provider = sources.get(provider_id) + if provider is None: + continue + requests = tuple(request for _document, request in items) + try: + provider_decisions = provider.authorize( + db, + principal, + requests=requests, + ) + except Exception: + LOGGER.exception( + "Search authorization recheck failed: %s", + provider_id, + ) + continue + for document, request in items: + decisions[document.id] = ( + provider_decisions.get( + request.reference.key + ) + is True + ) + return [ + _search_result(document, float(rank or 0.0)) + for document, rank in rows + if decisions.get(document.id) is True + ] + + def _source( + self, + provider_id: str, + resource_type: str, + ) -> tuple[SearchResourceType, SearchSourceProvider]: + for registered, provider in _search_sources(self.registry): + if registered.registration.id != provider_id: + continue + descriptors = [ + item + for item in provider.resource_types() + if item.provider_id == provider_id + and item.resource_type == resource_type + ] + if len(descriptors) != 1: + raise ValueError( + "Search source did not declare exactly one matching " + "resource type." + ) + return descriptors[0], provider + raise ValueError( + f"Search source is unavailable: {provider_id}" + ) + + @staticmethod + def _state( + db: Session, + *, + tenant_id: str, + provider_id: str, + resource_type: str, + lock: bool = False, + ) -> SearchIndexState | None: + statement = select(SearchIndexState).where( + SearchIndexState.tenant_id == tenant_id, + SearchIndexState.provider_id == provider_id, + SearchIndexState.resource_type == resource_type, + ) + if lock: + statement = statement.with_for_update() + return db.scalar(statement) + + +def aggregate_search_page( registry: object, session: Session, principal: object, *, query: SearchQuery, -) -> tuple[tuple[SearchResult, ...], tuple[dict[str, str], ...]]: - provider_limit = min(200, query.limit + query.offset) - provider_query = replace(query, limit=provider_limit, offset=0) +) -> AggregatedSearchPage: + provider_limit = MAX_AGGREGATE_WINDOW + provider_query = replace( + query, + limit=provider_limit, + offset=0, + cursor=None, + ) results: dict[tuple[str, str, str], SearchResult] = {} diagnostics: list[dict[str, str]] = [] for registered, provider in registry.search_providers(): @@ -182,12 +942,17 @@ def aggregate_search( query=provider_query, ) except Exception: - LOGGER.exception("Search provider failed: %s", provider_id) + LOGGER.exception( + "Search provider failed: %s", + provider_id, + ) diagnostics.append( { "provider_id": provider_id, "status": "unavailable", - "message": "This search provider is temporarily unavailable.", + "message": ( + "This search provider is temporarily unavailable." + ), } ) continue @@ -201,16 +966,45 @@ def aggregate_search( current = results.get(key) if current is None or normalized.score > current.score: results[key] = normalized - ordered = sorted( - results.values(), - key=lambda item: (-item.score, item.title.casefold(), item.resource_id), + ordered = sorted(results.values(), key=_result_sort_key) + if query.cursor: + cursor_key = _decode_cursor(query) + ordered = [ + item + for item in ordered + if _result_sort_key(item) > cursor_key + ] + elif query.offset: + ordered = ordered[query.offset :] + selected = ordered[: query.limit] + next_cursor = ( + _encode_cursor(query, selected[-1]) + if len(ordered) > query.limit and selected + else None ) - return ( - tuple(ordered[query.offset : query.offset + query.limit]), - tuple(diagnostics), + return AggregatedSearchPage( + results=tuple(selected), + diagnostics=tuple(diagnostics), + next_cursor=next_cursor, ) +def aggregate_search( + registry: object, + session: Session, + principal: object, + *, + query: SearchQuery, +) -> tuple[tuple[SearchResult, ...], tuple[dict[str, str], ...]]: + page = aggregate_search_page( + registry, + session, + principal, + query=query, + ) + return page.results, page.diagnostics + + def principal_acl_tokens(principal: object) -> tuple[str, ...]: token_groups: tuple[tuple[str, Any], ...] = ( ("account", getattr(principal, "account_id", None)), @@ -233,16 +1027,105 @@ def principal_acl_tokens(principal: object) -> tuple[str, ...]: for value in getattr(principal, attribute, ()) if value ) - return tuple(dict.fromkeys(tokens)) + return tuple(dict.fromkeys(tokens))[:500] -def _active_module_ids(registry: object | None) -> tuple[str, ...]: +def _active_module_ids( + registry: object | None, +) -> tuple[str, ...]: if registry is None or not hasattr(registry, "manifests"): return () return tuple(manifest.id for manifest in registry.manifests()) -def _document_payload(document: SearchDocument) -> dict[str, Any]: +def _search_sources( + registry: object | None, +) -> tuple[tuple[object, SearchSourceProvider], ...]: + if registry is None or not hasattr(registry, "search_sources"): + return () + return tuple(registry.search_sources()) + + +def _normalized_acl_tokens( + document: SearchDocument, +) -> tuple[str, ...]: + normalized = tuple( + dict.fromkeys( + token.strip() + for token in document.acl_tokens + if token.strip() + ) + ) + if document.visibility == "restricted" and not normalized: + raise ValueError( + "Restricted search documents require ACL tokens." + ) + if len(normalized) > 500: + raise ValueError( + "Search documents support at most 500 ACL tokens." + ) + return normalized + + +def _assert_safe_index_metadata( + document: SearchDocument, +) -> None: + _assert_no_secret_keys( + document.metadata, + location="metadata", + ) + if document.external_reference is not None: + _assert_no_secret_keys( + document.external_reference.metadata, + location="external_reference.metadata", + ) + + +def _assert_no_secret_keys( + value: object, + *, + location: str, +) -> None: + sensitive = { + "authorization", + "cookie", + "credential", + "credentials", + "password", + "private_key", + "refresh_token", + "secret", + "token", + "access_token", + } + if isinstance(value, Mapping): + for key, item in value.items(): + normalized = str(key).casefold().replace("-", "_") + if normalized in sensitive: + raise ValueError( + f"Search {location} must not contain secret field " + f"{key!r}." + ) + _assert_no_secret_keys( + item, + location=f"{location}.{key}", + ) + elif isinstance(value, Sequence) and not isinstance( + value, + (str, bytes), + ): + for item in value: + _assert_no_secret_keys( + item, + location=location, + ) + + +def _document_payload( + document: SearchDocument, + *, + rebuild_id: str | None, +) -> dict[str, Any]: search_text = " ".join( part for part in ( @@ -267,13 +1150,23 @@ def _document_payload(document: SearchDocument) -> dict[str, Any]: "visibility": document.visibility, "external_reference": external_reference, "metadata": dict(document.metadata), + "source_revision": document.source_revision, + "index_version": document.index_version, } content_hash = hashlib.sha256( - json.dumps(content, sort_keys=True, default=str).encode("utf-8") + json.dumps( + content, + sort_keys=True, + default=str, + ).encode("utf-8") ).hexdigest() return { "tenant_id": document.tenant_id, "module_id": document.module_id, + "provider_id": ( + document.provider_id + or f"{document.module_id}.index" + ), "resource_type": document.resource_type, "resource_id": document.resource_id, "title": document.title, @@ -286,10 +1179,105 @@ def _document_payload(document: SearchDocument) -> dict[str, Any]: "external_reference": external_reference, "metadata_": dict(document.metadata), "content_hash": content_hash, + "source_revision": document.source_revision, + "change_cursor": document.change_cursor, + "source_updated_at": document.source_updated_at, + "language": document.language, + "index_version": document.index_version, + "requires_authorization_recheck": ( + document.requires_authorization_recheck + ), + "rebuild_id": rebuild_id, + "active": True, "indexed_at": datetime.now(timezone.utc), } +def _document_contract_payload( + document: SearchDocument, +) -> dict[str, object]: + return { + "tenant_id": document.tenant_id, + "module_id": document.module_id, + "resource_type": document.resource_type, + "resource_id": document.resource_id, + "title": document.title, + "url": document.url, + "summary": document.summary, + "body": document.body, + "keywords": list(document.keywords), + "visibility": document.visibility, + "acl_tokens": list(document.acl_tokens), + "external_reference": ( + document.external_reference.to_dict() + if document.external_reference is not None + else None + ), + "metadata": dict(document.metadata), + "provider_id": document.provider_id, + "source_revision": document.source_revision, + "change_cursor": document.change_cursor, + "source_updated_at": ( + document.source_updated_at.isoformat() + if document.source_updated_at is not None + else None + ), + "language": document.language, + "index_version": document.index_version, + "requires_authorization_recheck": ( + document.requires_authorization_recheck + ), + } + + +def _document_from_contract_payload( + payload: Mapping[str, object], +) -> SearchDocument: + values = dict(payload) + external = values.get("external_reference") + if isinstance(external, Mapping): + external_values = dict(external) + observed = external_values.get("observed_at") + if isinstance(observed, str) and observed: + external_values["observed_at"] = datetime.fromisoformat( + observed + ) + values["external_reference"] = ExternalObjectReference( + **external_values + ) + source_updated_at = values.get("source_updated_at") + if isinstance(source_updated_at, str) and source_updated_at: + values["source_updated_at"] = datetime.fromisoformat( + source_updated_at + ) + values["keywords"] = tuple(values.get("keywords") or ()) + values["acl_tokens"] = tuple( + values.get("acl_tokens") or () + ) + return SearchDocument(**values) + + +def _validate_backfill_document( + document: SearchDocument, + *, + descriptor: SearchResourceType, + tenant_id: str, +) -> None: + if ( + document.tenant_id != tenant_id + or document.provider_id != descriptor.provider_id + or document.module_id != descriptor.module_id + or document.resource_type != descriptor.resource_type + or document.index_version != descriptor.index_version + or document.requires_authorization_recheck + != descriptor.requires_authorization_recheck + ): + raise ValueError( + "Search source returned a document outside its declared " + "tenant or resource contract." + ) + + def _search_result( document: SearchIndexDocument, rank: float, @@ -299,7 +1287,9 @@ def _search_result( payload = dict(document.external_reference) observed_at = payload.get("observed_at") if isinstance(observed_at, str) and observed_at: - payload["observed_at"] = datetime.fromisoformat(observed_at) + payload["observed_at"] = datetime.fromisoformat( + observed_at + ) reference = ExternalObjectReference(**payload) return SearchResult( provider_id="search.index", @@ -312,23 +1302,164 @@ def _search_result( score=rank, external_reference=reference, metadata=dict(document.metadata_), + source_revision=document.source_revision, + provenance={ + "backend": "database", + "source_provider": document.provider_id, + "source_revision": document.source_revision, + "index_version": document.index_version, + "indexed_at": document.indexed_at.isoformat(), + "acl": ( + "source_rechecked" + if document.requires_authorization_recheck + else "indexed_acl" + ), + }, ) +def _result_sort_key( + item: SearchResult, +) -> tuple[float, str, str, str, str]: + return ( + -round(float(item.score), 12), + item.title.casefold(), + item.module_id, + item.resource_type, + item.resource_id, + ) + + +def _query_fingerprint(query: SearchQuery) -> str: + payload = { + "text": query.text, + "tenant_id": query.tenant_id, + "module_ids": query.module_ids, + "resource_types": query.resource_types, + "context_kind": query.context_kind, + "context_id": query.context_id, + "language": query.language, + } + return hashlib.sha256( + json.dumps( + payload, + sort_keys=True, + separators=(",", ":"), + ).encode("utf-8") + ).hexdigest() + + +def _encode_cursor( + query: SearchQuery, + item: SearchResult, +) -> str: + payload = { + "v": 1, + "q": _query_fingerprint(query), + "k": list(_result_sort_key(item)), + } + encoded = base64.urlsafe_b64encode( + json.dumps( + payload, + separators=(",", ":"), + ).encode("utf-8") + ).decode("ascii") + return encoded.rstrip("=") + + +def _decode_cursor( + query: SearchQuery, +) -> tuple[float, str, str, str, str]: + value = query.cursor or "" + try: + padded = value + "=" * (-len(value) % 4) + payload = json.loads( + base64.urlsafe_b64decode( + padded.encode("ascii") + ).decode("utf-8") + ) + key = payload["k"] + if ( + payload.get("v") != 1 + or payload.get("q") != _query_fingerprint(query) + or not isinstance(key, list) + or len(key) != 5 + ): + raise ValueError + return ( + float(key[0]), + str(key[1]), + str(key[2]), + str(key[3]), + str(key[4]), + ) + except Exception as exc: + raise ValueError( + "Search cursor is invalid for this query." + ) from exc + + +def _state_payload(state: SearchIndexState) -> dict[str, object]: + return { + "provider_id": state.provider_id, + "module_id": state.module_id, + "resource_type": state.resource_type, + "index_version": state.index_version, + "status": state.status, + "checkpoint_cursor": state.checkpoint_cursor, + "high_watermark": state.high_watermark, + "last_change_cursor": state.last_change_cursor, + "indexed_documents": state.indexed_documents, + "rejected_documents": state.rejected_documents, + "rebuild_started_at": state.rebuild_started_at, + "rebuild_completed_at": state.rebuild_completed_at, + "last_success_at": state.last_success_at, + "last_error": state.last_error, + } + + def _require_tenant(principal: object, tenant_id: str) -> None: - principal_tenant = str(getattr(principal, "tenant_id", "") or "") - if principal_tenant != tenant_id: - raise PermissionError("Search operations cannot cross tenant boundaries.") + if _principal_tenant(principal) != tenant_id: + raise PermissionError( + "Search operations cannot cross tenant boundaries." + ) + + +def _principal_tenant(principal: object) -> str: + tenant_id = str( + getattr(principal, "tenant_id", "") or "" + ) + if not tenant_id: + raise PermissionError( + "Search operations require a tenant principal." + ) + return tenant_id def _session(value: object) -> Session: if not isinstance(value, Session): - raise TypeError("Search index operations require a SQLAlchemy session.") + raise TypeError( + "Search index operations require a SQLAlchemy session." + ) return value +def _as_utc(value: datetime) -> datetime: + if value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) + + +def _bounded_error(exc: Exception) -> str: + clean = " ".join(str(exc).split()) + return (clean or type(exc).__name__)[:2000] + + __all__ = [ + "AggregatedSearchPage", "SearchIndexService", + "SearchPage", "aggregate_search", + "aggregate_search_page", "principal_acl_tokens", ] diff --git a/tests/test_migrations.py b/tests/test_migrations.py index 88fba7b..284367b 100644 --- a/tests/test_migrations.py +++ b/tests/test_migrations.py @@ -26,13 +26,15 @@ class SearchMigrationTests(unittest.TestCase): try: with engine.connect() as connection: self.assertIn( - "a1b2c3d4e5f6", + "b2c3d4e5f607", set(MigrationContext.configure(connection).get_current_heads()), ) self.assertEqual( { "search_index_acl_tokens", + "search_index_change_queue", "search_index_documents", + "search_index_states", }, { name diff --git a/tests/test_search_service.py b/tests/test_search_service.py index 4b2a61e..7403882 100644 --- a/tests/test_search_service.py +++ b/tests/test_search_service.py @@ -6,20 +6,108 @@ from types import SimpleNamespace from sqlalchemy import create_engine from sqlalchemy.orm import Session -from govoplan_core.core.search import SearchDocument, SearchQuery +from govoplan_core.core.search import ( + SearchBackfillPage, + SearchDocument, + SearchIndexChange, + SearchQuery, + SearchResourceType, + SearchResult, +) from govoplan_core.db.base import Base from govoplan_search.backend.db.models import ( SearchIndexAclToken, + SearchIndexChangeQueue, SearchIndexDocument, + SearchIndexState, +) +from govoplan_search.backend.service import ( + SearchIndexService, + aggregate_search_page, ) -from govoplan_search.backend.service import SearchIndexService class _Registry: + def __init__(self, source=None, active_modules=("search", "cases")): + self.source = source + self.active_modules = active_modules + def manifests(self): + return tuple( + SimpleNamespace(id=module_id) + for module_id in self.active_modules + ) + + def search_sources(self): + if self.source is None: + return () return ( - SimpleNamespace(id="search"), - SimpleNamespace(id="cases"), + ( + SimpleNamespace( + registration=SimpleNamespace(id="cases.records") + ), + self.source, + ), + ) + + +class _Source: + def __init__(self, pages=()): + self.pages = list(pages) + self.authorized_ids = {"case-1"} + + def resource_types(self): + return ( + SearchResourceType( + provider_id="cases.records", + module_id="cases", + resource_type="case", + label="Cases", + requires_authorization_recheck=True, + ), + ) + + def backfill(self, session, *, request): + del session, request + return self.pages.pop(0) + + def authorize(self, session, principal, *, requests): + del session, principal + return { + request.reference.key: ( + request.reference.resource_id + in self.authorized_ids + ) + for request in requests + } + + +class _ResultProvider: + def search(self, session, principal, *, query): + del session, principal + return tuple( + SearchResult( + provider_id="ignored", + module_id="cases", + resource_type="case", + resource_id=f"case-{number}", + title=f"Permit {number}", + url=f"/cases/case-{number}", + score=float(10 - number), + ) + for number in range(1, 7) + )[: query.limit] + + +class _AggregateRegistry: + def search_providers(self): + return ( + ( + SimpleNamespace( + registration=SimpleNamespace(id="cases.live") + ), + _ResultProvider(), + ), ) @@ -31,6 +119,8 @@ class SearchServiceTests(unittest.TestCase): tables=( SearchIndexDocument.__table__, SearchIndexAclToken.__table__, + SearchIndexState.__table__, + SearchIndexChangeQueue.__table__, ), ) self.session = Session(self.engine) @@ -164,6 +254,267 @@ class SearchServiceTests(unittest.TestCase): ), ) + with self.assertRaisesRegex(ValueError, "secret field"): + self.service.upsert_document( + self.session, + self.principal, + document=SearchDocument( + tenant_id="tenant-1", + module_id="cases", + resource_type="case", + resource_id="case-secret", + title="Permit", + url="/cases/case-secret", + acl_tokens=("account:account-1",), + metadata={"access_token": "must-not-index"}, + ), + ) + + def test_query_rejects_cross_tenant_principal(self) -> None: + with self.assertRaises(PermissionError): + self.service.search( + self.session, + self.principal, + query=SearchQuery( + text="permit", + tenant_id="tenant-2", + ), + ) + + def test_source_authorization_recheck_is_fail_closed(self) -> None: + source = _Source() + service = SearchIndexService(_Registry(source)) + for resource_id in ("case-1", "case-2"): + service.upsert_document( + self.session, + self.principal, + document=_source_document(resource_id), + ) + self.session.flush() + + results = service.search( + self.session, + self.principal, + query=SearchQuery( + text="permit", + tenant_id="tenant-1", + ), + ) + self.assertEqual( + ["case-1"], + [item.resource_id for item in results], + ) + unavailable_results = SearchIndexService( + _Registry() + ).search( + self.session, + self.principal, + query=SearchQuery( + text="permit", + tenant_id="tenant-1", + ), + ) + self.assertEqual((), unavailable_results) + + def test_durable_change_queue_is_idempotent(self) -> None: + document = _source_document("case-1") + change = SearchIndexChange( + change_id="change-1", + provider_id="cases.records", + kind="upsert", + reference=document.reference, + source_revision=document.source_revision, + cursor=document.change_cursor or "", + document=document, + ) + + self.assertTrue( + self.service.enqueue_change( + self.session, + change=change, + ) + ) + self.assertFalse( + self.service.enqueue_change( + self.session, + change=change, + ) + ) + result = self.service.process_changes(self.session) + self.assertEqual(1, result["applied"]) + self.assertEqual( + "case-1", + self.session.query(SearchIndexDocument).one().resource_id, + ) + self.assertEqual( + "applied", + self.session.query(SearchIndexChangeQueue).one().status, + ) + + delete_change = SearchIndexChange( + change_id="change-2", + provider_id="cases.records", + kind="delete", + reference=document.reference, + source_revision="3", + cursor="cursor-2", + ) + self.service.enqueue_change( + self.session, + change=delete_change, + ) + self.service.process_changes(self.session) + self.assertEqual( + 0, + self.session.query(SearchIndexDocument).count(), + ) + + def test_rebuild_resumes_and_removes_stale_documents(self) -> None: + stale = _source_document("stale-case") + self.service.upsert_document( + self.session, + self.principal, + document=stale, + ) + source = _Source( + pages=( + SearchBackfillPage( + documents=(_source_document("case-1"),), + next_cursor="page-2", + complete=False, + high_watermark="changes-9", + ), + SearchBackfillPage( + documents=(_source_document("case-2"),), + next_cursor=None, + complete=True, + high_watermark="changes-9", + ), + ) + ) + service = SearchIndexService(_Registry(source)) + + started = service.start_rebuild( + self.session, + self.principal, + provider_id="cases.records", + resource_type="case", + ) + first_page = service.continue_rebuild( + self.session, + self.principal, + provider_id="cases.records", + resource_type="case", + ) + self.assertEqual(started.id, first_page.id) + self.assertEqual("backfilling", first_page.status) + self.assertEqual("page-2", first_page.checkpoint_cursor) + + complete = service.continue_rebuild( + self.session, + self.principal, + provider_id="cases.records", + resource_type="case", + ) + self.assertEqual("ready", complete.status) + self.assertEqual("changes-9", complete.high_watermark) + self.assertEqual( + ["case-1", "case-2"], + sorted( + item.resource_id + for item in self.session.query( + SearchIndexDocument + ).all() + ), + ) + + def test_module_reconciliation_disables_derived_rows(self) -> None: + self.service.upsert_document( + self.session, + self.principal, + document=SearchDocument( + tenant_id="tenant-1", + module_id="cases", + resource_type="case", + resource_id="case-1", + title="Permit", + url="/cases/case-1", + acl_tokens=("account:account-1",), + ), + ) + disabled_service = SearchIndexService( + _Registry(active_modules=("search",)) + ) + result = disabled_service.reconcile_active_modules( + self.session + ) + self.assertEqual(1, result["disabled_documents"]) + self.assertFalse( + self.session.query(SearchIndexDocument).one().active + ) + + def test_aggregate_cursor_is_stable_and_non_overlapping(self) -> None: + registry = _AggregateRegistry() + first = aggregate_search_page( + registry, + self.session, + self.principal, + query=SearchQuery( + text="permit", + tenant_id="tenant-1", + limit=2, + ), + ) + second = aggregate_search_page( + registry, + self.session, + self.principal, + query=SearchQuery( + text="permit", + tenant_id="tenant-1", + limit=2, + cursor=first.next_cursor, + ), + ) + + self.assertIsNotNone(first.next_cursor) + self.assertEqual( + {"case-1", "case-2"}, + {item.resource_id for item in first.results}, + ) + self.assertEqual( + {"case-3", "case-4"}, + {item.resource_id for item in second.results}, + ) + with self.assertRaisesRegex(ValueError, "cursor"): + aggregate_search_page( + registry, + self.session, + self.principal, + query=SearchQuery( + text="different", + tenant_id="tenant-1", + limit=2, + cursor=first.next_cursor, + ), + ) + + +def _source_document(resource_id: str) -> SearchDocument: + return SearchDocument( + tenant_id="tenant-1", + module_id="cases", + resource_type="case", + resource_id=resource_id, + title=f"Permit {resource_id}", + url=f"/cases/{resource_id}", + acl_tokens=("account:account-1",), + provider_id="cases.records", + source_revision="2", + change_cursor=f"cursor-{resource_id}", + requires_authorization_recheck=True, + ) + if __name__ == "__main__": unittest.main() diff --git a/webui/src/api/search.ts b/webui/src/api/search.ts index 755d54c..61e78ad 100644 --- a/webui/src/api/search.ts +++ b/webui/src/api/search.ts @@ -27,6 +27,8 @@ export type SearchResult = { breadcrumbs: string[]; external_reference?: SearchExternalReference | null; metadata: Record; + source_revision?: string | null; + provenance: Record; }; export type SearchResponse = { @@ -37,6 +39,8 @@ export type SearchResponse = { status: "unavailable"; message: string; }>; + next_cursor?: string | null; + has_more: boolean; }; export type SearchRequest = { @@ -47,6 +51,8 @@ export type SearchRequest = { contextId?: string; limit?: number; offset?: number; + cursor?: string; + language?: string; }; export function search( @@ -63,7 +69,9 @@ export function search( context_kind: request.contextKind, context_id: request.contextId, limit: request.limit, - offset: request.offset + offset: request.offset, + cursor: request.cursor, + language: request.language }), { signal } ); diff --git a/webui/src/features/search/SearchPage.tsx b/webui/src/features/search/SearchPage.tsx index 95fa669..0fe69dd 100644 --- a/webui/src/features/search/SearchPage.tsx +++ b/webui/src/features/search/SearchPage.tsx @@ -1,5 +1,5 @@ -import { ExternalLink, Search } from "lucide-react"; -import { useEffect, useMemo, useState, type FormEvent } from "react"; +import { ChevronDown, ExternalLink, Search } from "lucide-react"; +import { useCallback, useEffect, useMemo, useState, type FormEvent } from "react"; import { useSearchParams } from "react-router-dom"; import { DismissibleAlert, @@ -15,26 +15,31 @@ export default function SearchPage({ settings }: PlatformRouteContext) { const navigate = useGuardedNavigate(); const [params, setParams] = useSearchParams(); const query = params.get("q") ?? ""; - const modules = params.getAll("module"); - const resourceTypes = params.getAll("resource_type"); + const moduleKey = params.getAll("module").join("\u001f"); + const resourceTypeKey = params.getAll("resource_type").join("\u001f"); + const contextId = params.get("context") ?? undefined; + const modules = useMemo( + () => moduleKey ? moduleKey.split("\u001f") : [], + [moduleKey] + ); + const resourceTypes = useMemo( + () => resourceTypeKey ? resourceTypeKey.split("\u001f") : [], + [resourceTypeKey] + ); const [draft, setDraft] = useState(query); const [response, setResponse] = useState(null); const [loading, setLoading] = useState(false); const [error, setError] = useState(""); const requestKey = useMemo( - () => JSON.stringify([query, modules, resourceTypes]), - [query, modules, resourceTypes] + () => JSON.stringify([query, moduleKey, resourceTypeKey, contextId]), + [contextId, moduleKey, query, resourceTypeKey] ); useEffect(() => { setDraft(query); }, [query]); - useEffect(() => { - if (!query.trim()) { - setResponse(null); - return; - } + const loadResults = useCallback((cursor?: string) => { const controller = new AbortController(); setLoading(true); setError(""); @@ -45,20 +50,48 @@ export default function SearchPage({ settings }: PlatformRouteContext) { modules, resourceTypes, contextKind: modules.length ? "module" : "global", - contextId: params.get("context") ?? undefined, - limit: 100 + contextId, + limit: 50, + cursor }, controller.signal ). - then(setResponse). + then((next) => { + setResponse((current) => + cursor && current ? + { + ...next, + results: [...current.results, ...next.results], + diagnostics: [ + ...current.diagnostics, + ...next.diagnostics.filter((item) => + !current.diagnostics.some( + (currentItem) => currentItem.provider_id === item.provider_id + ) + ) + ] + } : + next + ); + }). catch((reason) => { if ((reason as Error).name !== "AbortError") { setError(reason instanceof Error ? reason.message : "Search failed."); } }). finally(() => setLoading(false)); + return controller; + }, [contextId, modules, query, resourceTypes, settings]); + + useEffect(() => { + if (!query.trim()) { + setResponse(null); + return; + } + setResponse(null); + const controller = loadResults(); return () => controller.abort(); - }, [requestKey, settings]); + }, [loadResults, requestKey]); function submit(event: FormEvent) { event.preventDefault(); @@ -123,6 +156,16 @@ export default function SearchPage({ settings }: PlatformRouteContext) { )} + {response?.next_cursor && + + } ); diff --git a/webui/src/styles/search.css b/webui/src/styles/search.css index ec4b416..d3d95ee 100644 --- a/webui/src/styles/search.css +++ b/webui/src/styles/search.css @@ -222,6 +222,36 @@ font-size: 12px; } +.search-load-more { + display: inline-flex; + align-items: center; + gap: 7px; + margin-top: 14px; + border: 1px solid var(--control-border); + border-radius: 4px; + background: linear-gradient( + var(--control-gradient-start), + var(--control-gradient-end) + ); + color: var(--control-text); + cursor: pointer; + padding: 7px 12px; + font: inherit; + font-weight: 700; +} + +.search-load-more:hover:not(:disabled) { + background: linear-gradient( + var(--control-gradient-start), + var(--control-gradient-end-hover) + ); +} + +.search-load-more:disabled { + cursor: default; + opacity: 0.55; +} + @media (max-width: 900px) { .global-search { width: 34px;