from __future__ import annotations from govoplan_core.core.datasources import ( DatasourceAccessError, DatasourceField, DatasourceNotFoundError, DatasourceOrigin, DatasourceOriginReadRequest, DatasourceOriginReadResult, DatasourceUnavailableError, DatasourceValidationError, ) from govoplan_core.core.tabular_sources import ( TabularReadRequest, TabularSource, TabularSourceAccessError, TabularSourceError, TabularSourceNotFoundError, TabularSourceUnavailableError, TabularSourceValidationError, ) from govoplan_connectors.backend.tabular_sources import SqlTabularSourceProvider class ConnectorDatasourceOriginProvider: """Expose connector-owned sources through the Datasources origin contract.""" def __init__(self, provider: SqlTabularSourceProvider | None = None) -> None: self._provider = provider or SqlTabularSourceProvider() def list_origins( self, session: object, principal: object, *, query: str = "", limit: int = 100, ): try: rows = self._provider.list_sources( session, principal, query=query, limit=limit, ) except TabularSourceError as exc: raise _datasource_error(exc) from exc return tuple(_origin(source) for source in rows) def get_origin( self, session: object, principal: object, *, origin_ref: str, ) -> DatasourceOrigin | None: try: source = self._provider.get_source( session, principal, source_ref=origin_ref, ) except TabularSourceError as exc: raise _datasource_error(exc) from exc return _origin(source) if source is not None else None def read_origin( self, session: object, principal: object, *, request: DatasourceOriginReadRequest, ) -> DatasourceOriginReadResult: try: result = self._provider.read_source( session, principal, request=TabularReadRequest( source_ref=request.origin_ref, limit=request.limit, offset=request.offset, columns=request.columns, expected_fingerprint=request.expected_fingerprint, max_bytes=request.max_bytes, timeout_ms=request.timeout_ms, ), ) except TabularSourceError as exc: raise _datasource_error(exc) from exc return DatasourceOriginReadResult( origin=_origin(result.source), rows=result.rows, total_rows=result.total_rows, truncated=result.truncated, returned_bytes=result.returned_bytes, elapsed_ms=result.elapsed_ms, effective_row_limit=result.effective_row_limit, effective_byte_limit=result.effective_byte_limit, effective_timeout_ms=result.effective_timeout_ms, diagnostics=result.diagnostics, ) def _origin(source: TabularSource) -> DatasourceOrigin: kind = { "managed_file": "file", "postgresql": "database", }.get(source.provider, "upload") return DatasourceOrigin( ref=source.ref, source_name=source.source_name, name=source.name, description=source.description, kind=kind, shape="tabular", supported_modes=("live", "cached"), provider=f"connectors.{source.provider}", schema=tuple( DatasourceField( name=column.name, data_type=column.data_type, nullable=column.nullable, ) for column in source.schema ), schema_version=source.schema_version, fingerprint=source.fingerprint, row_count=source.row_count, byte_count=source.byte_count, updated_at=source.updated_at, capabilities=source.capabilities, metadata=dict(source.metadata), source_mode=source.source_mode, pushdown=source.pushdown, health=source.health, ) def _datasource_error(exc: TabularSourceError): if isinstance(exc, TabularSourceAccessError): return DatasourceAccessError(str(exc)) if isinstance(exc, TabularSourceNotFoundError): return DatasourceNotFoundError(str(exc)) if isinstance(exc, TabularSourceUnavailableError): return DatasourceUnavailableError(str(exc)) if isinstance(exc, TabularSourceValidationError): return DatasourceValidationError(str(exc)) return DatasourceValidationError(str(exc)) __all__ = ["ConnectorDatasourceOriginProvider"]