from __future__ import annotations import tempfile import unittest from io import BytesIO from pathlib import Path from types import SimpleNamespace from openpyxl import Workbook 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 ( TabularSourceUnavailableError, TabularSourceValidationError, ) from govoplan_core.db.base import Base from govoplan_connectors.backend.db.models import ( ConnectorConfiguration, ConnectorDefinition, ) from govoplan_connectors.backend.tabular_adapters import ( ManagedFileTabularAdapter, PostgresqlTabularAdapter, parse_managed_tabular_content, ) 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, payload: bytes, *, filename: str = "cases.csv") -> None: self.payload = payload self.filename = filename self.current_version_id = "version-2" def _file(self, version_id: str) -> ManagedTabularFile: return ManagedTabularFile( file_asset_id="asset-1", file_version_id=version_id, filename=self.filename, display_path=f"Imports/{self.filename}", content_type=( "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet" if self.filename.endswith(".xlsx") else "text/csv" ), size_bytes=len(self.payload), sha256=("a" if version_id == "version-1" else "b") * 64, current_version=version_id == self.current_version_id, ) def list_tabular_files(self, session, principal, *, query="", limit=100): del session, principal, query, limit return (self._file(self.current_version_id),) 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._file(file_version_id or self.current_version_id) def read_tabular_file( self, session, principal, *, file_asset_id, file_version_id, max_bytes, ): del session, principal, file_asset_id if len(self.payload) > max_bytes: raise AssertionError("test payload exceeded adapter limit") return ManagedTabularFileContent( file=self._file(file_version_id), payload=self.payload, ) class _Registry: def __init__(self, provider) -> None: self.provider = provider 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.provider class ManagedFileTabularAdapterTests(unittest.TestCase): def test_csv_is_exact_version_pinned_and_reports_newer_version(self) -> None: adapter = ManagedFileTabularAdapter( _Registry(_ManagedFiles(b"id,amount\n0012,12.5\n2,7\n")) ) result = adapter.inspect( object(), principal(), file_asset_id="asset-1", file_version_id="version-1", ) self.assertEqual("managed_file", result.provider) self.assertEqual("version-1", result.metadata["file_version_id"]) self.assertEqual("warning", result.health.status) self.assertEqual("files.newer_version_available", result.health.code) self.assertEqual("mixed", result.schema[0].data_type) self.assertEqual(2, result.row_count) self.assertTrue(result.pushdown.projections) self.assertEqual("files.newer_version_available", result.diagnostics[0].code) def test_xlsx_uses_requested_sheet_and_closed_typed_schema(self) -> None: workbook = Workbook() first = workbook.active first.title = "Ignore" first.append(["ignored"]) target = workbook.create_sheet("Monthly") target.append(["case_id", "amount", "active"]) target.append(["A-1", 12.5, True]) target.append(["A-2", None, False]) payload = BytesIO() workbook.save(payload) workbook.close() rows, sheet = parse_managed_tabular_content( payload.getvalue(), filename="monthly.xlsx", content_type=( "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet" ), delimiter=",", sheet_name="Monthly", ) self.assertEqual("Monthly", sheet) self.assertEqual("A-1", rows[0]["case_id"]) self.assertIsNone(rows[1]["amount"]) def test_missing_files_capability_is_explicitly_unavailable(self) -> None: with self.assertRaisesRegex( TabularSourceUnavailableError, "require the Files module", ): ManagedFileTabularAdapter(None).inspect( object(), principal(), file_asset_id="asset-1", file_version_id=None, ) class PostgresqlTabularAdapterTests(unittest.TestCase): def setUp(self) -> None: self.directory = tempfile.TemporaryDirectory( prefix="govoplan-connectors-sql-adapter-" ) source_path = Path(self.directory.name) / "source.db" self.source_url = f"sqlite+pysqlite:///{source_path}" source_engine = create_engine(self.source_url) metadata = MetaData() self.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( self.table.insert(), ( {"case_id": "A-1", "amount": 12}, {"case_id": "A-2", "amount": None}, ), ) source_engine.dispose() self.catalog_engine = create_engine("sqlite+pysqlite:///:memory:") Base.metadata.create_all( self.catalog_engine, tables=( ConnectorDefinition.__table__, ConnectorConfiguration.__table__, ), ) self.session = sessionmaker(bind=self.catalog_engine)() 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.source_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.adapter = PostgresqlTabularAdapter(allow_sqlite_for_tests=True) def tearDown(self) -> None: self.session.close() self.catalog_engine.dispose() self.directory.cleanup() def test_discovers_and_reads_projection_from_governed_sql_configuration(self) -> None: inspection = self.adapter.inspect( self.session, principal(), configuration_id=self.configuration.id, table_name="monthly_cases", ) metadata = { **dict(inspection.metadata), "discovery_fingerprint": inspection.fingerprint, } read = self.adapter.read( self.session, principal(), metadata=metadata, columns=("case_id",), offset=1, limit=10, timeout_ms=2_000, ) self.assertEqual("live", "live") self.assertEqual(["case_id", "amount"], [item.name for item in inspection.schema]) self.assertEqual(2, inspection.row_count) self.assertEqual(({"case_id": "A-2"},), read.rows) self.assertTrue(inspection.pushdown.projections) self.assertFalse(inspection.pushdown.filters) def test_schema_drift_and_tenant_isolation_fail_closed(self) -> None: inspection = self.adapter.inspect( self.session, principal(), configuration_id=self.configuration.id, table_name="monthly_cases", ) metadata = { **dict(inspection.metadata), "discovery_fingerprint": inspection.fingerprint, } engine = create_engine(self.source_url) with engine.begin() as connection: connection.exec_driver_sql( "ALTER TABLE monthly_cases ADD COLUMN category TEXT" ) engine.dispose() with self.assertRaisesRegex(TabularSourceValidationError, "schema drifted"): self.adapter.read( self.session, principal(), metadata=metadata, columns=(), offset=0, limit=10, timeout_ms=2_000, ) with self.assertRaisesRegex( TabularSourceUnavailableError, "configuration is unavailable", ): self.adapter.inspect( self.session, principal("tenant-2"), configuration_id=self.configuration.id, table_name="monthly_cases", ) def test_endpoint_query_credentials_are_rejected_before_connection(self) -> None: self.configuration.endpoint_url = f"{self.source_url}?password=not-allowed" self.session.commit() with self.assertRaisesRegex( TabularSourceValidationError, "query parameters must not contain credentials", ): self.adapter.inspect( self.session, principal(), configuration_id=self.configuration.id, table_name="monthly_cases", ) if __name__ == "__main__": unittest.main()