from __future__ import annotations import os import sqlite3 import threading import unittest import uuid from concurrent.futures import ThreadPoolExecutor from unittest.mock import MagicMock, patch from sqlalchemy import create_engine, text from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session, sessionmaker from govoplan_core.db.base import Base, utcnow from govoplan_poll.backend import service as poll_service from govoplan_poll.backend.db.models import Poll, PollOption, PollResponse from govoplan_poll.backend.schemas import ( PollAnswerInput, PollCreateRequest, PollOptionInput, PollSubmitResponseRequest, ) from govoplan_poll.backend.service import ( PollError, _share_lock_poll_for_response, create_poll, submit_poll_response, ) class PollResponseUniquenessTests(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite:///:memory:") self.tables = [Poll.__table__, PollOption.__table__, PollResponse.__table__] Base.metadata.create_all(self.engine, tables=self.tables) self.Session = sessionmaker(bind=self.engine) self.session: Session = self.Session() def tearDown(self) -> None: self.session.close() Base.metadata.drop_all(self.engine, tables=list(reversed(self.tables))) self.engine.dispose() def _poll( self, *, allow_anonymous: bool = False, allow_response_update: bool = True, ) -> Poll: return create_poll( self.session, tenant_id="tenant-1", user_id="owner-1", payload=PollCreateRequest( title="One response each", kind="single_choice", status="open", allow_anonymous=allow_anonymous, allow_response_update=allow_response_update, options=[ PollOptionInput(key="yes", label="Yes"), PollOptionInput(key="no", label="No"), ], ), ) @staticmethod def _payload( respondent_id: str | None, *, option_key: str = "yes", ) -> PollSubmitResponseRequest: return PollSubmitResponseRequest( respondent_id=respondent_id, respondent_label=respondent_id, answers=[PollAnswerInput(option_key=option_key)], ) def _hide_first_lookup(self): original = poll_service._existing_response calls = 0 def hide_once(session, *, poll, respondent_id): nonlocal calls calls += 1 if calls == 1: return None return original( session, poll=poll, respondent_id=respondent_id, ) return patch.object(poll_service, "_existing_response", side_effect=hide_once) def test_sqlite_invariant_reconciles_a_racing_insert_to_normal_update(self) -> None: poll = self._poll() winner = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1"), ) with self._hide_first_lookup(): reconciled = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1", option_key="no"), ) self.assertEqual(reconciled.id, winner.id) self.assertEqual(reconciled.answers[0]["option_key"], "no") self.assertEqual( self.session.query(PollResponse) .filter(PollResponse.deleted_at.is_(None)) .count(), 1, ) def test_racing_insert_obeys_update_disabled_policy(self) -> None: poll = self._poll() winner = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1"), ) poll.allow_response_update = False self.session.flush() with self._hide_first_lookup(): with self.assertRaisesRegex( PollError, "Response updates are not allowed", ): submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1", option_key="no"), ) self.assertEqual(winner.answers[0]["option_key"], "yes") self.assertEqual(self.session.query(PollResponse).count(), 1) def test_anonymous_and_tombstoned_responses_do_not_conflict(self) -> None: poll = self._poll(allow_anonymous=True) anonymous_one = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload(None), ) anonymous_two = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload(None, option_key="no"), ) identified_one = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1"), ) identified_one.deleted_at = utcnow() self.session.flush() identified_two = submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("person-1", option_key="no"), ) self.assertNotEqual(anonymous_one.id, anonymous_two.id) self.assertNotEqual(identified_one.id, identified_two.id) self.assertEqual(self.session.query(PollResponse).count(), 4) def test_unrelated_integrity_error_is_not_reconciled(self) -> None: poll = self._poll() self.session.execute( text( """ CREATE TRIGGER reject_blocked_poll_response BEFORE INSERT ON poll_responses WHEN NEW.respondent_id = 'blocked' BEGIN SELECT RAISE(ABORT, 'blocked by unrelated invariant'); END """ ) ) with self.assertRaises(IntegrityError) as caught: submit_poll_response( self.session, tenant_id=poll.tenant_id, poll_id=poll.id, payload=self._payload("blocked"), ) self.assertIsInstance(caught.exception.orig, sqlite3.IntegrityError) self.assertIn("blocked by unrelated invariant", str(caught.exception.orig)) def test_response_read_lock_is_shared(self) -> None: session = MagicMock() query = MagicMock() poll = MagicMock() session.query.return_value = query query.filter.return_value = query query.populate_existing.return_value = query query.with_for_update.return_value = query query.first.return_value = poll result = _share_lock_poll_for_response( session, tenant_id="tenant-1", poll_id="poll-1", ) self.assertIs(result, poll) query.with_for_update.assert_called_once_with(read=True) def test_orm_mirrors_both_partial_index_predicates(self) -> None: index = next( index for index in PollResponse.__table__.indexes if index.name == "uq_poll_responses_active_respondent" ) self.assertTrue(index.unique) self.assertEqual( str(index.dialect_options["sqlite"]["where"]), "deleted_at IS NULL AND respondent_id IS NOT NULL", ) self.assertEqual( str(index.dialect_options["postgresql"]["where"]), "deleted_at IS NULL AND respondent_id IS NOT NULL", ) @unittest.skipUnless( os.environ.get("GOVOPLAN_POLL_TEST_POSTGRES_URL"), "set GOVOPLAN_POLL_TEST_POSTGRES_URL for the two-session PostgreSQL check", ) class PollResponsePostgresConcurrencyTests(unittest.TestCase): def setUp(self) -> None: database_url = os.environ["GOVOPLAN_POLL_TEST_POSTGRES_URL"] self.schema = f"poll_response_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 = [Poll.__table__, PollOption.__table__, PollResponse.__table__] Base.metadata.create_all(self.engine, tables=self.tables) 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_two_sessions_converge_on_one_active_response(self) -> None: with Session(self.engine) as session: poll = create_poll( session, tenant_id="tenant-1", user_id="owner-1", payload=PollCreateRequest( title="Concurrent response", kind="single_choice", status="open", options=[ PollOptionInput(key="yes", label="Yes"), PollOptionInput(key="no", label="No"), ], ), ) poll_id = poll.id session.commit() original = poll_service._existing_response barrier = threading.Barrier(2) thread_state = threading.local() def synchronize_first_lookup(session, *, poll, respondent_id): response = original( session, poll=poll, respondent_id=respondent_id, ) lookup_count = getattr(thread_state, "lookup_count", 0) + 1 thread_state.lookup_count = lookup_count if lookup_count == 1: self.assertIsNone(response) barrier.wait(timeout=10) return response def submit(option_key: str) -> str: with Session(self.engine) as session: response = submit_poll_response( session, tenant_id="tenant-1", poll_id=poll_id, payload=PollSubmitResponseRequest( respondent_id="person-1", respondent_label="Person One", answers=[PollAnswerInput(option_key=option_key)], ), ) response_id = response.id session.commit() return response_id with patch.object( poll_service, "_existing_response", side_effect=synchronize_first_lookup, ): with ThreadPoolExecutor(max_workers=2) as executor: response_ids = tuple( executor.map(submit, ("yes", "no")) ) self.assertEqual(len(set(response_ids)), 1) with Session(self.engine) as session: active = ( session.query(PollResponse) .filter( PollResponse.poll_id == poll_id, PollResponse.respondent_id == "person-1", PollResponse.deleted_at.is_(None), ) .all() ) self.assertEqual(len(active), 1) self.assertIn(active[0].answers[0]["option_key"], {"yes", "no"}) if __name__ == "__main__": unittest.main()