354 lines
12 KiB
Python
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()
|