diff --git a/config.py b/config.py index 5e6889f..7afc2c0 100644 --- a/config.py +++ b/config.py @@ -26,6 +26,7 @@ class DefaultConfig: DATABASE_URL = os.environ.get("DATABASE_URL", "") DATABASE_POOL_MIN_SIZE = int(os.environ.get("DATABASE_POOL_MIN_SIZE", "1")) DATABASE_POOL_MAX_SIZE = int(os.environ.get("DATABASE_POOL_MAX_SIZE", "10")) + DATABASE_ACQUIRE_TIMEOUT_SECONDS = float(os.environ.get("DATABASE_ACQUIRE_TIMEOUT_SECONDS", "30")) LOG_QUERIES = os.environ.get("LOG_QUERIES", "") VALID_X_GITLAB_TOKEN = os.environ.get("VALID_X_GITLAB_TOKEN", "") MESSAGE_DELETE_DELAY_SECONDS = int(os.environ.get("MESSAGE_DELETE_DELAY_SECONDS", "30")) diff --git a/db.py b/db.py index 5462b7c..de8dfb3 100644 --- a/db.py +++ b/db.py @@ -43,6 +43,9 @@ class DatabaseLifecycleHandler: def __init__(self, conf: DefaultConfig): self._pool: asyncpg.Pool | None = None self._config = conf + # A webhook handler pins its lock connection for its whole body and needs one more for + # the body itself: capping holders at pool // 2 - 1 leaves the rest to everyone else. + self.handler_slots = asyncio.Semaphore(max(1, conf.DATABASE_POOL_MAX_SIZE // 2 - 1)) async def connect(self): log.debug("creating database connection pool") @@ -89,7 +92,7 @@ async def disconnect(self): async def acquire(self) -> asyncpg.pool.PoolAcquireContext: assert self._pool is not None - return self._pool.acquire() + return self._pool.acquire(timeout=self._config.DATABASE_ACQUIRE_TIMEOUT_SECONDS) class GitlabUser(BaseModel): diff --git a/docs/workflow.md b/docs/workflow.md index 67886bb..029a0e7 100644 --- a/docs/workflow.md +++ b/docs/workflow.md @@ -41,6 +41,7 @@ When a merge request event occurs in GitLab: **Database Protection**: - Unique constraint on `(merge_request_ref_id, conversation_token)` prevents duplicate messages per channel - `FOR UPDATE` lock in `get_or_create_message_refs()` prevents concurrent corruption +- A per-MR advisory lock (`mr_lock` in `webhook/merge_request.py`) serialises the handlers of one merge request, so the 3-4 deliveries GitLab fans out for a single event patch each Teams card once instead of concurrently ### 2. MR Switches to Draft diff --git a/tests/conftest.py b/tests/conftest.py index 09d68cd..5146885 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,9 +1,12 @@ +import asyncio + #!/usr/bin/env python3 import time import uuid from collections.abc import AsyncGenerator from typing import Any +from unittest.mock import DEFAULT from unittest.mock import AsyncMock from unittest.mock import MagicMock @@ -31,10 +34,18 @@ def __init__(self): self.execute = AsyncMock() self.fetch = AsyncMock(return_value=[]) self.fetchrow = AsyncMock(return_value=None) - self.fetchval = AsyncMock(return_value=None) + # The per-MR advisory lock is always granted; every other fetchval keeps the answer a + # test configured through return_value (DEFAULT falls through to it). + self.fetchval = AsyncMock(return_value=None, side_effect=self._fetchval) self.prepare = AsyncMock() self._transaction = None + @staticmethod + def _fetchval(query, *args, **kwargs): + if "pg_try_advisory_xact_lock" in query: + return True + return DEFAULT + def transaction(self): if self._transaction is None: self._transaction = MockTransaction() @@ -63,6 +74,7 @@ async def __aexit__(self, exc_type, exc_val, exc_tb): class MockDatabase: def __init__(self): self.connection = MockConnection() + self.handler_slots = asyncio.Semaphore(4) async def acquire(self): return MockAcquireContext(self.connection) @@ -447,6 +459,9 @@ async def __aexit__(self, *args): async def request(self, method, url, **kwargs): requests.append({"method": method, "url": url, "kwargs": kwargs}) + # A real call yields to the loop while in flight; without this the handlers under + # test never interleave and no race can show up. + await asyncio.sleep(0.01) if responses: response_data = responses.pop(0) diff --git a/tests/test_debounce_dedup.py b/tests/test_debounce_dedup.py index 6e1e15a..20b73ef 100644 --- a/tests/test_debounce_dedup.py +++ b/tests/test_debounce_dedup.py @@ -1,6 +1,7 @@ #!/usr/bin/env python3 """Tests for debounce and deduplication mechanisms.""" +import asyncio import datetime from unittest.mock import AsyncMock @@ -203,9 +204,15 @@ async def test_merge_request_skips_stats_fetch_when_no_update_needed(self, sampl mock_conn = AsyncMock() mock_conn.fetchrow = AsyncMock(return_value=None) + # the per-MR lock opens a transaction and reads pg_try_advisory_xact_lock + mock_conn.transaction = MagicMock( + return_value=MagicMock(__aenter__=AsyncMock(), __aexit__=AsyncMock()) + ) + mock_conn.fetchval = AsyncMock(return_value=True) mock_database.acquire = AsyncMock( return_value=MagicMock(__aenter__=AsyncMock(return_value=mock_conn)) ) + mock_database.handler_slots = asyncio.Semaphore(1) mock_render.return_value = {"type": "AdaptiveCard"} mock_fingerprint.return_value = "test-fingerprint" diff --git a/tests/test_e2e_race_conditions.py b/tests/test_e2e_race_conditions.py index d35fd92..e7a5749 100644 --- a/tests/test_e2e_race_conditions.py +++ b/tests/test_e2e_race_conditions.py @@ -338,6 +338,135 @@ async def test_e2e_race_concurrent_updates_with_locking( assert len(exceptions) == 0, f"No exceptions should occur: {exceptions}" +@pytest.mark.asyncio +async def test_e2e_fanout_deliveries_patch_each_card_once( + db_connection, clean_database, base_mr_payload, mock_activity_api, db_lifecycle_handler +): + """ + RACE CONDITION: GitLab delivers one event through several hooks at once. + + Scenario: a card exists; the same `update` event arrives 4 times within milliseconds + Result: the Teams card is patched exactly once. The other deliveries wait for the first, + then find its fingerprint already stored and skip. + Impact: concurrent PUTs on one Teams activity are what earns the fixed ~10s penalty upstream + + Location: webhook/merge_request.py:mr_lock + """ + from db import DBHelper + from gitlab_model import MergeRequestPayload + from webhook.merge_request import merge_request + + dbh = DBHelper(db_lifecycle_handler) + conv_token = str(uuid.uuid4()) + + base_mr_payload["object_attributes"]["action"] = "open" + open_payload = MergeRequestPayload(**base_mr_payload) + + update = base_mr_payload.copy() + update["object_attributes"] = base_mr_payload["object_attributes"].copy() + update["object_attributes"]["action"] = "update" + update["object_attributes"]["updated_at"] = "2025-01-01 00:01:00 UTC" + update_payload = MergeRequestPayload(**update) + + with ( + patch("webhook.merge_request.database", db_lifecycle_handler), + patch("webhook.merge_request.dbh", dbh), + patch("webhook.messaging.database", db_lifecycle_handler), + patch("db.database", db_lifecycle_handler), + patch("webhook.merge_request.render") as mock_render, + ): + mock_render.return_value = {"type": "AdaptiveCard"} + + await merge_request( + mr=open_payload, + conversation_tokens=[conv_token], + participant_ids_filter=[], + new_commits_revoke_approvals=False, + ) + assert sum(r["method"] == "POST" for r in mock_activity_api["requests"]) == 1 + + results = await asyncio.gather( + *[ + merge_request( + mr=update_payload, + conversation_tokens=[conv_token], + participant_ids_filter=[], + new_commits_revoke_approvals=False, + ) + for _ in range(4) + ], + return_exceptions=True, + ) + + assert [r for r in results if isinstance(r, Exception)] == [] + patches = [r for r in mock_activity_api["requests"] if r["method"] == "PATCH"] + assert len(patches) == 1, f"one card, one event, expected one PATCH, got {len(patches)}" + + +@pytest.mark.asyncio +async def test_e2e_distinct_mrs_never_starve_the_pool( + db_connection, clean_database, base_mr_payload, mock_activity_api, test_database_url +): + """ + RESILIENCE: handlers of different MRs must not deadlock the connection pool. + + Scenario: a pool of 2 and two MRs arriving together; each handler pins its lock connection + and needs one more for its body + Result: both complete. Without the holder cap both wait on each other forever, and + /healthz shares the pool, so the pod would hang until liveness restarts it + + Location: db.py:handler_slots, webhook/merge_request.py:mr_lock + """ + from config import DefaultConfig + from db import DatabaseLifecycleHandler + from db import DBHelper + from gitlab_model import MergeRequestPayload + from webhook.merge_request import merge_request + + cfg = DefaultConfig() + cfg.DATABASE_URL = test_database_url + cfg.DATABASE_POOL_MIN_SIZE = 1 + cfg.DATABASE_POOL_MAX_SIZE = 2 + handler = DatabaseLifecycleHandler(cfg) + await handler.connect() + dbh = DBHelper(handler) + + payloads = [] + for iid in (1, 2): + p = base_mr_payload.copy() + p["object_attributes"] = base_mr_payload["object_attributes"].copy() + p["object_attributes"]["iid"] = iid + p["object_attributes"]["id"] = 1000 + iid + p["object_attributes"]["action"] = "open" + payloads.append(MergeRequestPayload(**p)) + + with ( + patch("webhook.merge_request.database", handler), + patch("webhook.merge_request.dbh", dbh), + patch("webhook.messaging.database", handler), + patch("db.database", handler), + patch("webhook.merge_request.render", return_value={"type": "AdaptiveCard"}), + ): + tasks = [ + asyncio.ensure_future( + merge_request( + mr=p, + conversation_tokens=[str(uuid.uuid4())], + participant_ids_filter=[], + new_commits_revoke_approvals=False, + ) + ) + for p in payloads + ] + done, pending = await asyncio.wait(tasks, timeout=15) + for t in pending: + t.cancel() + await handler.disconnect() + + assert not pending, "handlers of distinct MRs starved each other of pool connections" + assert [t.exception() for t in done if t.exception()] == [] + + @pytest.mark.asyncio async def test_e2e_race_update_during_close_deletion( db_connection, clean_database, base_mr_payload, mock_activity_api, db_lifecycle_handler diff --git a/tests/test_mr_lock.py b/tests/test_mr_lock.py new file mode 100644 index 0000000..f7c8125 --- /dev/null +++ b/tests/test_mr_lock.py @@ -0,0 +1,75 @@ +#!/usr/bin/env python3 +"""The per-MR lock: it polls without pinning a connection, and it degrades instead of failing.""" + +import asyncio + +from unittest.mock import AsyncMock +from unittest.mock import patch + +from webhook.merge_request import mr_lock + + +async def test_waits_by_polling_and_releases_the_connection_between_tries(mock_database): + mock_db = mock_database + mock_conn = mock_db.connection + mock_conn.fetchval = AsyncMock(side_effect=[False, False, True]) + acquires = 0 + real_acquire = mock_db.acquire + + async def counting_acquire(): + nonlocal acquires + acquires += 1 + return await real_acquire() + + mock_db.acquire = counting_acquire + ran = False + + with patch("webhook.merge_request.database", mock_db): + async with mr_lock(42, poll=0.001): + ran = True + + assert ran + assert mock_conn.fetchval.await_count == 3 + # one fresh acquire per attempt: a waiter never holds a connection while it sleeps + assert acquires == 3 + assert mock_conn.fetchval.await_args_list[0].args == ("SELECT pg_try_advisory_xact_lock(1, $1::int)", 42) + + +async def test_runs_unlocked_after_the_deadline_instead_of_failing(mock_database): + mock_db = mock_database + mock_conn = mock_db.connection + mock_conn.fetchval = AsyncMock(return_value=False) + ran = False + + with ( + patch("webhook.merge_request.database", mock_db), + patch("webhook.merge_request.logger") as mock_logger, + ): + async with mr_lock(42, timeout=0.05, poll=0.01): + ran = True + + assert ran + mock_logger.warning.assert_called_once() + assert mock_logger.warning.call_args.args[0] == "merge request lock not acquired, proceeding unlocked" + assert mock_logger.warning.call_args.kwargs["merge_request_ref_id"] == 42 + + +async def test_holders_are_capped_by_the_handler_slots(mock_database): + """More holders than slots must queue, not pile onto the pool.""" + mock_db = mock_database + mock_db.handler_slots = asyncio.Semaphore(1) + inside = 0 + peak = 0 + + async def hold(): + nonlocal inside, peak + async with mr_lock(7): + inside += 1 + peak = max(peak, inside) + await asyncio.sleep(0.01) + inside -= 1 + + # one patch around all three: concurrent patches of one attribute unpatch each other + with patch("webhook.merge_request.database", mock_db): + await asyncio.gather(hold(), hold(), hold()) + assert peak == 1 diff --git a/webhook/merge_request.py b/webhook/merge_request.py index f904b94..6e2ca5e 100644 --- a/webhook/merge_request.py +++ b/webhook/merge_request.py @@ -1,7 +1,12 @@ #!/usr/bin/env python3 +import asyncio +import contextlib import datetime import hashlib +from collections.abc import AsyncIterator + +import asyncpg import fastapi_structured_logging import httpx @@ -26,6 +31,39 @@ logger = fastapi_structured_logging.get_logger() +@contextlib.asynccontextmanager +async def mr_lock( + merge_request_ref_id: int, *, timeout: float = 30.0, poll: float = 0.05 +) -> AsyncIterator[None]: + """Serialise the handlers of one merge request across both replicas. + + Transaction-scoped: the pool never resets sessions (NoResetConnection), so a session lock + would leak. Polled: a waiter must not pin a pool connection while it waits. + """ + deadline = asyncio.get_running_loop().time() + timeout + async with database.handler_slots: + while True: + connection: asyncpg.Connection + async with await database.acquire() as connection: + async with connection.transaction(): + acquired = await connection.fetchval( + "SELECT pg_try_advisory_xact_lock(1, $1::int)", merge_request_ref_id + ) + if acquired: + yield + return + if asyncio.get_running_loop().time() >= deadline: + # Losing the update is worse than a rare concurrent patch: run unlocked. + logger.warning( + "merge request lock not acquired, proceeding unlocked", + merge_request_ref_id=merge_request_ref_id, + timeout_seconds=timeout, + ) + yield + return + await asyncio.sleep(poll) + + class PartialMessageUpdateError(Exception): """Raised when some but not all messages were updated successfully.""" @@ -64,6 +102,27 @@ async def merge_request( # This prevents payload corruption from OOO events merge_request_ref_id = await dbh.get_or_create_merge_request_ref_id(mr) + async with mr_lock(merge_request_ref_id): + return await _merge_request_locked( + mr, + merge_request_ref_id, + payload_updated_at, + is_closing_action, + conversation_tokens, + participant_ids_filter, + new_commits_revoke_approvals, + ) + + +async def _merge_request_locked( + mr: MergeRequestPayload, + merge_request_ref_id: int, + payload_updated_at: datetime.datetime, + is_closing_action: bool, + conversation_tokens: list[str], + participant_ids_filter: list[int], + new_commits_revoke_approvals: bool, +): # Early out-of-order check # Reopen bypasses OOO: recreates messages after close, or updates existing if close arrives late # Closing actions check OOO: late close must not delete a reopened MR's messages