from __future__ import annotations import os import threading import unittest import uuid from concurrent.futures import ThreadPoolExecutor from sqlalchemy import create_engine, text from sqlalchemy.orm import Session from govoplan_core.auth import ApiPrincipal from govoplan_core.core.access import PrincipalRef from govoplan_core.core.change_sequence import ChangeSequenceEntry from govoplan_core.core.datasources import DatasourcePublicationRequest from govoplan_core.db.base import Base from govoplan_datasources.backend.db.models import ( DatasourceMaterializationRecord, DatasourcePayloadRecord, DatasourcePayloadRowRecord, DatasourcePublicationRecord, DatasourceRecord, ) from govoplan_datasources.backend.service import ( SOURCE_WRITE_SCOPE, SqlDatasourceProvider, _append_materialization, ) from govoplan_datasources.backend.tabular import ( encoded_size, field_payload, fingerprint_rows, infer_schema, ) @unittest.skipUnless( os.environ.get("GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL"), "set GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL for PostgreSQL concurrency checks", ) class DatasourceMaterializationPostgresTests(unittest.TestCase): def setUp(self) -> None: database_url = os.environ["GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL"] self.schema = f"datasource_revision_race_{uuid.uuid4().hex}" self.admin_engine = create_engine(database_url) with self.admin_engine.begin() as connection: connection.execute(text(f'CREATE SCHEMA "{self.schema}"')) self.engine = create_engine( database_url, connect_args={"options": f"-c search_path={self.schema}"}, ) self.tables = [ DatasourceRecord.__table__, DatasourcePayloadRecord.__table__, DatasourcePayloadRowRecord.__table__, DatasourceMaterializationRecord.__table__, DatasourcePublicationRecord.__table__, ChangeSequenceEntry.__table__, ] Base.metadata.create_all(self.engine, tables=self.tables) with Session(self.engine) as session: session.add( DatasourceRecord( id="datasource-1", tenant_id="tenant-1", source_name="concurrent", name="Concurrent", kind="custom", mode="static", shape="tabular", ) ) session.commit() def tearDown(self) -> None: try: Base.metadata.drop_all( self.engine, tables=list(reversed(self.tables)), ) finally: self.engine.dispose() with self.admin_engine.begin() as connection: connection.execute(text(f'DROP SCHEMA "{self.schema}"')) self.admin_engine.dispose() def test_competing_publishers_receive_distinct_monotonic_revisions(self) -> None: barrier = threading.Barrier(2) def publish(value: int) -> int: rows = ({"id": value},) schema = infer_schema(rows) with Session(self.engine) as session: datasource = session.get(DatasourceRecord, "datasource-1") self.assertIsNotNone(datasource) barrier.wait(timeout=10) materialization = _append_materialization( session, datasource=datasource, rows=rows, schema=[field_payload(field) for field in schema], fingerprint=fingerprint_rows(rows, schema), byte_count=encoded_size(rows), actor_id=f"worker-{value}", set_current=True, ) revision = materialization.revision session.commit() return revision with ThreadPoolExecutor(max_workers=2) as executor: revisions = tuple(executor.map(publish, (1, 2))) self.assertEqual((1, 2), tuple(sorted(revisions))) with Session(self.engine) as session: self.assertEqual( [1, 2], list( session.scalars( DatasourceMaterializationRecord.__table__.select() .with_only_columns( DatasourceMaterializationRecord.revision ) .order_by( DatasourceMaterializationRecord.revision.asc() ) ) ), ) def test_same_publication_identity_is_serialized_and_replayed(self) -> None: barrier = threading.Barrier(2) provider = SqlDatasourceProvider() request = DatasourcePublicationRequest( producer_module="dataflow", producer_run_ref="dataflow-run:concurrent", idempotency_key="concurrent-publication", name="Concurrent publication", source_name="concurrent_publication", rows=({"id": 1, "result": "match"},), ) def publish(worker: int) -> tuple[str, str, bool]: api_principal = ApiPrincipal( principal=PrincipalRef( account_id=f"worker-{worker}", membership_id=f"membership-{worker}", tenant_id="tenant-1", scopes=frozenset((SOURCE_WRITE_SCOPE,)), ), account=object(), user=object(), ) with Session(self.engine) as session: barrier.wait(timeout=10) result = provider.publish_rows( session, api_principal, request=request, ) session.commit() return result.ref, result.materialization.ref, result.replayed with ThreadPoolExecutor(max_workers=2) as executor: results = tuple(executor.map(publish, (1, 2))) self.assertEqual(1, len({item[0] for item in results})) self.assertEqual(1, len({item[1] for item in results})) self.assertEqual([False, True], sorted(item[2] for item in results)) with Session(self.engine) as session: self.assertEqual( 1, session.query(DatasourcePublicationRecord).count(), ) self.assertEqual( 1, session.query(DatasourceMaterializationRecord).count(), ) if __name__ == "__main__": unittest.main()