Files
govoplan-poll/tests/test_response_uniqueness.py

354 lines
12 KiB
Python

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()