from __future__ import annotations import tempfile import unittest from pathlib import Path from unittest.mock import patch from sqlalchemy import Column, String, Table, create_engine from sqlalchemy.orm import Session, sessionmaker from govoplan_campaign.backend.db.models import ( Campaign, CampaignIssue, CampaignVersion, ) from govoplan_campaign.backend.persistence.versions import update_campaign_version from govoplan_core.core.change_sequence import ChangeSequenceEntry from govoplan_core.core.concurrency import RevisionConflictError from govoplan_core.db.base import Base class CampaignOptimisticConcurrencyTests(unittest.TestCase): def setUp(self) -> None: self.temp_dir = tempfile.TemporaryDirectory() database_path = Path(self.temp_dir.name) / "campaign.db" self.engine = create_engine(f"sqlite+pysqlite:///{database_path}") access_users = Base.metadata.tables.get("access_users") if access_users is None: access_users = Table( "access_users", Base.metadata, Column("id", String(36), primary_key=True), ) access_groups = Base.metadata.tables.get("access_groups") if access_groups is None: access_groups = Table( "access_groups", Base.metadata, Column("id", String(36), primary_key=True), ) Base.metadata.create_all( self.engine, tables=[ access_users, access_groups, ChangeSequenceEntry.__table__, Campaign.__table__, CampaignVersion.__table__, CampaignIssue.__table__, ], ) self.SessionLocal = sessionmaker( bind=self.engine, class_=Session, expire_on_commit=False, ) with self.SessionLocal() as session: campaign = Campaign( id="campaign-1", tenant_id="tenant-1", external_id="C-1", name="Campaign", current_version_id="version-1", ) version = CampaignVersion( id="version-1", campaign_id=campaign.id, version_number=1, raw_json={ "version": "1.0", "campaign": { "id": "C-1", "name": "Campaign", "description": "Base", }, }, ) session.add_all((campaign, version)) session.commit() self.addCleanup(self.engine.dispose) self.addCleanup(self.temp_dir.cleanup) @staticmethod def _update( session: Session, *, expected_revision: int, name: str, ) -> CampaignVersion: with ( patch( "govoplan_campaign.backend.persistence.versions._updated_runtime_json", side_effect=lambda _session, **kwargs: kwargs["raw_json"], ), patch( "govoplan_campaign.backend.persistence.versions._write_campaign_snapshot" ), ): return update_campaign_version( session, tenant_id="tenant-1", campaign_id="campaign-1", version_id="version-1", raw_json={ "version": "1.0", "campaign": { "id": "C-1", "name": name, "description": "Base", }, }, expected_revision=expected_revision, ) def test_only_one_of_two_writers_can_commit_the_same_revision(self) -> None: first = self.SessionLocal() second = self.SessionLocal() self.addCleanup(first.close) self.addCleanup(second.close) first.get(CampaignVersion, "version-1") second.get(CampaignVersion, "version-1") saved = self._update( first, expected_revision=1, name="First writer", ) self.assertEqual(saved.edit_revision, 2) with self.assertRaises(RevisionConflictError) as raised: self._update( second, expected_revision=1, name="Second writer", ) self.assertEqual(raised.exception.current_revision, 2) self.assertEqual(raised.exception.submitted_base_revision, 1) with self.SessionLocal() as verification: current = verification.get(CampaignVersion, "version-1") assert current is not None self.assertEqual(current.raw_json["campaign"]["name"], "First writer") self.assertEqual(current.edit_revision, 2) def test_stale_revision_is_rejected_before_mutation(self) -> None: with self.SessionLocal() as first: self._update( first, expected_revision=1, name="First writer", ) with self.SessionLocal() as stale: with self.assertRaises(RevisionConflictError): self._update( stale, expected_revision=1, name="Stale writer", ) stale.rollback() with self.SessionLocal() as verification: current = verification.get(CampaignVersion, "version-1") assert current is not None self.assertEqual(current.raw_json["campaign"]["name"], "First writer") if __name__ == "__main__": unittest.main()