Files
govoplan-connectors/tests/test_tabular_adapters.py
zemion 2c2b11f860
Module Package Release / publish-packages (push) Successful in 12s
feat(connectors): add managed file and PostgreSQL origins
2026-08-21 19:29:58 +02:00

342 lines
11 KiB
Python

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()