from __future__ import annotations import tempfile import unittest from pathlib import Path from types import SimpleNamespace from sqlalchemy import Column, Integer, MetaData, String, Table, create_engine from sqlalchemy.orm import sessionmaker from govoplan_core.auth import ApiPrincipal from govoplan_core.core.access import PrincipalRef from govoplan_core.core.files import ( CAPABILITY_FILES_TABULAR_CONTENT, ManagedTabularFile, ManagedTabularFileContent, ) from govoplan_core.core.tabular_sources import ( TabularReadRequest, TabularSourceUnavailableError, TabularSourceValidationError, ) from govoplan_core.db.base import Base from govoplan_core.security.credential_envelopes import CredentialEnvelope from govoplan_connectors.backend.db.models import ( ConnectorConfiguration, ConnectorDefinition, ConnectorTabularSource, ) from govoplan_connectors.backend.tabular_adapters import PostgresqlTabularAdapter from govoplan_connectors.backend.tabular_sources import SqlTabularSourceProvider def principal(tenant_id: str = "tenant-1") -> ApiPrincipal: return ApiPrincipal( principal=PrincipalRef( account_id="account-1", membership_id="membership-1", tenant_id=tenant_id, scopes=frozenset( { "connectors:source:read", "connectors:source:write", "files:file:read", "files:file:download", } ), ), account=object(), user=SimpleNamespace(id="user-1"), ) class _ManagedFiles: def __init__(self) -> None: self.current = "version-1" self.payloads = { "version-1": b"id,name\n1,Ada\n2,Lin\n", "version-2": b"id,name,active\n1,Ada,true\n2,Lin,false\n", } def _metadata(self, version_id: str) -> ManagedTabularFile: payload = self.payloads[version_id] return ManagedTabularFile( file_asset_id="asset-1", file_version_id=version_id, filename="people.csv", display_path="Imports/people.csv", content_type="text/csv", size_bytes=len(payload), sha256=("a" if version_id == "version-1" else "b") * 64, current_version=version_id == self.current, ) def list_tabular_files(self, session, principal, *, query="", limit=100): del session, principal, query, limit return (self._metadata(self.current),) def get_tabular_file( self, session, principal, *, file_asset_id, file_version_id=None, ): del session, principal if file_asset_id != "asset-1": return None return self._metadata(file_version_id or self.current) def read_tabular_file( self, session, principal, *, file_asset_id, file_version_id, max_bytes, ): del session, principal, file_asset_id, max_bytes return ManagedTabularFileContent( file=self._metadata(file_version_id), payload=self.payloads[file_version_id], ) class _Registry: def __init__(self, files) -> None: self.files = files def has_capability(self, name): return name == CAPABILITY_FILES_TABULAR_CONTENT def require_capability(self, name): if not self.has_capability(name): raise KeyError(name) return self.files class ConnectorTabularOriginProviderTests(unittest.TestCase): def setUp(self) -> None: self.catalog_engine = create_engine("sqlite+pysqlite:///:memory:") Base.metadata.create_all( self.catalog_engine, tables=( ConnectorTabularSource.__table__, ConnectorDefinition.__table__, ConnectorConfiguration.__table__, CredentialEnvelope.__table__, ), ) self.session = sessionmaker(bind=self.catalog_engine)() self.files = _ManagedFiles() self.directory = tempfile.TemporaryDirectory( prefix="govoplan-connectors-origin-provider-" ) source_path = Path(self.directory.name) / "source.db" self.sql_url = f"sqlite+pysqlite:///{source_path}" source_engine = create_engine(self.sql_url) metadata = MetaData() source_table = Table( "monthly_cases", metadata, Column("case_id", String, nullable=False), Column("amount", Integer, nullable=True), ) metadata.create_all(source_engine) with source_engine.begin() as connection: connection.execute( source_table.insert(), ( {"case_id": "A-1", "amount": 12}, {"case_id": "A-2", "amount": None}, ), ) source_engine.dispose() definition = ConnectorDefinition( id="definition-1", tenant_id="tenant-1", definition_key="postgresql.reader", name="PostgreSQL reader", status="active", current_revision=1, local_definition=True, ) self.configuration = ConnectorConfiguration( id="configuration-1", tenant_id="tenant-1", definition_id=definition.id, name="Monthly SQL", status="active", endpoint_url=self.sql_url, credential_ref=None, base_definition_revision=1, local_overrides={}, protected_paths=[], effective_configuration={"provider": "sql", "protocol": "sql"}, effective_hash="configuration-hash-1", resource_revision=1, ambiguity_policy="manual_review", ) self.session.add_all((definition, self.configuration)) self.session.commit() self.provider = SqlTabularSourceProvider( registry=_Registry(self.files), sql_adapter=PostgresqlTabularAdapter(allow_sqlite_for_tests=True), ) def tearDown(self) -> None: self.session.close() self.catalog_engine.dispose() self.directory.cleanup() def test_managed_file_source_stays_pinned_until_explicit_refresh(self) -> None: created = self.provider.create_file_source( self.session, principal(), name="Managed people", source_name="managed_people", file_asset_id="asset-1", ) self.session.commit() self.files.current = "version-2" preview = self.provider.read_source( self.session, principal(), request=TabularReadRequest(source_ref=created.ref, limit=10), ) refreshed = self.provider.refresh_source( self.session, principal(), source_ref=created.ref, ) self.assertTrue(created.ref.startswith("file:")) self.assertEqual("file_backed", preview.source.source_mode) self.assertEqual("version-1", preview.source.metadata["file_version_id"]) self.assertEqual( "files.newer_version_available", preview.diagnostics[0].code, ) self.assertEqual("version-2", refreshed.metadata["file_version_id"]) self.assertEqual("2", refreshed.schema_version) self.assertEqual(3, len(refreshed.schema)) def test_sql_source_projects_and_blocks_changed_configuration_until_refresh(self) -> None: created = self.provider.create_sql_source( self.session, principal(), name="Monthly cases", source_name="monthly_cases", configuration_id=self.configuration.id, table_name="monthly_cases", ) self.session.commit() preview = self.provider.read_source( self.session, principal(), request=TabularReadRequest( source_ref=created.ref, columns=("case_id",), limit=1, ), ) self.assertTrue(created.ref.startswith("sql:")) self.assertEqual("live", preview.source.source_mode) self.assertEqual(({"case_id": "A-1"},), preview.rows) self.assertEqual("preview.row_limit_reached", preview.diagnostics[-1].code) self.assertIsNone( self.provider.get_source( self.session, principal("tenant-2"), source_ref=created.ref, ) ) self.configuration.effective_hash = "configuration-hash-2" self.configuration.resource_revision = 2 self.session.commit() with self.assertRaisesRegex( TabularSourceValidationError, "configuration changed", ): self.provider.read_source( self.session, principal(), request=TabularReadRequest(source_ref=created.ref), ) refreshed = self.provider.refresh_source( self.session, principal(), source_ref=created.ref, ) self.assertEqual("2", refreshed.schema_version) self.assertEqual("configuration-hash-2", refreshed.metadata["configuration_hash"]) def test_inactive_or_stale_sql_credentials_have_sanitized_diagnostics(self) -> None: self.configuration.endpoint_url = ( "postgresql+psycopg://db.example.invalid/govoplan" ) self.configuration.credential_ref = "missing-credential" self.session.commit() adapter = PostgresqlTabularAdapter() with self.assertRaisesRegex( TabularSourceUnavailableError, "credential is unavailable, inactive, or outside its allowed scope", ): adapter.inspect( self.session, principal(), configuration_id=self.configuration.id, table_name="monthly_cases", ) self.configuration.status = "disabled" self.session.commit() with self.assertRaisesRegex( TabularSourceUnavailableError, "configuration is not active", ): adapter.inspect( self.session, principal(), configuration_id=self.configuration.id, table_name="monthly_cases", ) if __name__ == "__main__": unittest.main()