342 lines
11 KiB
Python
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()
|