"""SearchProvider protocol and global registry.""" from __future__ import annotations import logging from typing import TYPE_CHECKING, Protocol, runtime_checkable from sqlalchemy import select if TYPE_CHECKING: import uuid from sqlalchemy.ext.asyncio import AsyncSession logger = logging.getLogger(__name__) @runtime_checkable class SearchProvider(Protocol): """Protocol for entity-specific search providers.""" entity_type: str async def search_fts( self, db: AsyncSession, tsquery: str, tenant_id: uuid.UUID, limit: int, ) -> list[dict]: """Full-text search using PostgreSQL tsquery.""" ... async def search_vector( self, db: AsyncSession, embedding: list[float], tenant_id: uuid.UUID, limit: int, ) -> list[dict]: """Semantic vector search using pgvector.""" ... async def get_embedding_text( self, db: AsyncSession, entity_id: uuid.UUID, tenant_id: uuid.UUID, ) -> str: """Get text representation for embedding generation.""" ... def to_search_result(self, entity: object) -> dict: """Convert an ORM entity to a search result dict.""" ... class SearchProviderRegistry: """In-memory registry of search providers.""" def __init__(self) -> None: self._providers: dict[str, SearchProvider] = {} def register(self, provider: SearchProvider) -> None: """Register a search provider.""" self._providers[provider.entity_type] = provider logger.debug("Registered search provider: %s", provider.entity_type) def unregister(self, entity_type: str) -> None: """Unregister a search provider by entity type.""" self._providers.pop(entity_type, None) def get(self, entity_type: str) -> SearchProvider | None: """Get a provider by entity type.""" return self._providers.get(entity_type) def get_all(self) -> list[SearchProvider]: """Get all registered providers.""" return list(self._providers.values()) def get_entity_types(self) -> list[str]: """Get all registered entity types.""" return list(self._providers.keys()) def clear(self) -> None: """Clear all registered providers.""" self._providers.clear() # Global singleton _registry = SearchProviderRegistry() def get_search_registry() -> SearchProviderRegistry: """Get the global search provider registry.""" return _registry async def auto_register_providers(db: AsyncSession) -> None: """Auto-register providers for active plugins. Checks which plugins are active and registers corresponding providers. """ from app.plugins.builtins.unified_search.providers.contact_provider import ( ContactSearchProvider, ) from app.plugins.builtins.unified_search.providers.company_provider import ( CompanySearchProvider, ) from app.plugins.builtins.unified_search.providers.mail_provider import ( MailSearchProvider, ) from app.plugins.builtins.unified_search.providers.file_provider import ( FileSearchProvider, ) from app.plugins.builtins.unified_search.providers.event_provider import ( EventSearchProvider, ) registry = get_search_registry() registry.clear() # Register all built-in providers for provider_cls in [ ContactSearchProvider, CompanySearchProvider, MailSearchProvider, FileSearchProvider, EventSearchProvider, ]: try: registry.register(provider_cls()) except Exception: logger.exception("Failed to register %s", provider_cls.__name__)