Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"))
Expand Down
5 changes: 4 additions & 1 deletion db.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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):
Expand Down
1 change: 1 addition & 0 deletions docs/workflow.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
17 changes: 16 additions & 1 deletion tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
7 changes: 7 additions & 0 deletions tests/test_debounce_dedup.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
#!/usr/bin/env python3
"""Tests for debounce and deduplication mechanisms."""

import asyncio
import datetime

from unittest.mock import AsyncMock
Expand Down Expand Up @@ -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"
Expand Down
129 changes: 129 additions & 0 deletions tests/test_e2e_race_conditions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
75 changes: 75 additions & 0 deletions tests/test_mr_lock.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading