314 lines
10 KiB
Python
314 lines
10 KiB
Python
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()
|