from __future__ import annotations import unittest from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from govoplan_core.auth import ApiPrincipal from govoplan_core.core.access import PrincipalRef from govoplan_core.core.datasources import DatasourceOriginReadRequest from govoplan_core.core.tabular_sources import TabularSnapshotInput from govoplan_core.db.base import Base from govoplan_connectors.backend.datasource_origins import ( ConnectorDatasourceOriginProvider, ) from govoplan_connectors.backend.db.models import ConnectorTabularSource from govoplan_connectors.backend.tabular_sources import ( WRITE_SCOPE, SqlTabularSourceProvider, ) def principal(*, scopes: tuple[str, ...]) -> ApiPrincipal: return ApiPrincipal( principal=PrincipalRef( account_id="account-1", membership_id="membership-1", tenant_id="tenant-1", scopes=frozenset(scopes), ), account=object(), user=object(), ) class ConnectorDatasourceOriginTests(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite:///:memory:") Base.metadata.create_all( self.engine, tables=[ConnectorTabularSource.__table__], ) self.Session = sessionmaker(bind=self.engine) self.session = self.Session() provider = SqlTabularSourceProvider() self.source = provider.create_snapshot( self.session, principal(scopes=(WRITE_SCOPE,)), snapshot=TabularSnapshotInput( name="Imported cases", source_name="imported_cases", rows=({"id": 1, "name": "Ada"},), ), ) self.session.commit() self.origins = ConnectorDatasourceOriginProvider(provider) def tearDown(self) -> None: self.session.close() Base.metadata.drop_all( self.engine, tables=[ConnectorTabularSource.__table__], ) self.engine.dispose() def test_datasource_reader_can_discover_and_read_connector_origin(self) -> None: datasource_principal = principal( scopes=("datasources:catalogue:read",), ) origins = self.origins.list_origins( self.session, datasource_principal, ) result = self.origins.read_origin( self.session, datasource_principal, request=DatasourceOriginReadRequest(origin_ref=self.source.ref), ) self.assertEqual((self.source.ref,), tuple(item.ref for item in origins)) self.assertEqual(("live", "cached"), origins[0].supported_modes) self.assertEqual(({"id": 1, "name": "Ada"},), result.rows) if __name__ == "__main__": unittest.main()