mirror of
https://github.com/rennf93/roboco.git
synced 2026-08-03 07:23:24 +02:00
[F056] messaging: lock group row before session check-then-create to prevent active-session orphan race
This commit is contained in:
@@ -386,6 +386,27 @@ class MessagingService(BaseService):
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def _lock_group(self, group_id: UUID) -> GroupTable | None:
|
||||
"""Lock the group row for the rest of this transaction and re-read its
|
||||
``active_session_id`` under the lock.
|
||||
|
||||
``SELECT ... FOR UPDATE`` serializes concurrent session creation for the
|
||||
same group: a check-then-create on ``active_session_id`` otherwise races
|
||||
(two concurrent posts both miss the active session, both INSERT a new
|
||||
ACTIVE session, and the second ``flush`` overwrites the group's
|
||||
``active_session_id`` — orphaning the first session as a forever-ACTIVE,
|
||||
unreferenced row). ``populate_existing`` refreshes the identity-map-cached
|
||||
instance's columns under the lock so the re-check sees the winner's link,
|
||||
not the stale pre-lock value.
|
||||
"""
|
||||
result = await self.session.execute(
|
||||
select(GroupTable)
|
||||
.where(GroupTable.id == group_id)
|
||||
.with_for_update()
|
||||
.execution_options(populate_existing=True)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_groups_in_channel(self, channel_id: UUID) -> list[GroupTable]:
|
||||
"""List all groups in a channel."""
|
||||
result = await self.session.execute(
|
||||
@@ -429,6 +450,16 @@ class MessagingService(BaseService):
|
||||
# sessions (the smoke run showed ~one session per message, and the CEO
|
||||
# could not hold a conversation). Only open a fresh session when none is
|
||||
# currently active (the prior one closed via timeout / boundary / merge).
|
||||
#
|
||||
# Serialize the check-then-create per group: lock the row and re-read
|
||||
# ``active_session_id`` under the lock so a concurrent creator that won
|
||||
# the race is visible here (we reuse its session) instead of both
|
||||
# inserting and orphaning the loser's session. ``get_or_create_active_
|
||||
# session`` and the channel-post adapter (L1868) route through here, so
|
||||
# the lock covers every active-session creation entry point.
|
||||
group = await self._lock_group(req.group_id)
|
||||
if group is None:
|
||||
raise ValueError(f"Group {req.group_id} not found")
|
||||
if group.active_session_id:
|
||||
existing = await self.get_session(cast("UUID", group.active_session_id))
|
||||
if existing is not None and existing.status == SessionStatus.ACTIVE:
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
"""F056: ``create_session`` (and its delegate ``get_or_create_active_session``,
|
||||
plus the L1868 channel-post adapter that routes through it) must not orphan an
|
||||
ACTIVE session under concurrent posts.
|
||||
|
||||
``create_session`` does a plain check-then-create: read ``group.active_session_id``,
|
||||
reuse if ACTIVE, else INSERT a new ACTIVE session and point the group at it.
|
||||
Two concurrent posts can both miss the active session, both INSERT, and the
|
||||
second ``flush`` overwrites ``group.active_session_id`` — the first session
|
||||
stays ACTIVE but unreferenced (orphaned) forever. There is no DB uniqueness on
|
||||
``(group_id, status='active')`` (tables.py:1121-1125 only carries indexes), so
|
||||
nothing stops the double-insert.
|
||||
|
||||
The fix: lock the group row (``SELECT ... FOR UPDATE``) and re-read
|
||||
``active_session_id`` under the lock before deciding to create, so concurrent
|
||||
callers serialize per group and the loser reuses the winner's session.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from roboco.db.tables import SessionTable
|
||||
from roboco.models.base import SessionStatus
|
||||
from roboco.models.messaging import SessionCreateRequest
|
||||
from roboco.services.messaging import MessagingService
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
_GROUP_ID = MagicMock(name="group-id")
|
||||
_WINNER_SESSION_ID = MagicMock(name="winner-session-id")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lock_group_emits_for_update() -> None:
|
||||
"""``_lock_group`` must issue ``SELECT ... FOR UPDATE`` (the row lock that
|
||||
serializes concurrent session creation per group)."""
|
||||
session = AsyncMock()
|
||||
captured: list[Any] = []
|
||||
result_mock = MagicMock()
|
||||
result_mock.scalar_one_or_none.return_value = MagicMock(active_session_id=None)
|
||||
|
||||
async def _exec(stmt: Any) -> Any:
|
||||
captured.append(stmt)
|
||||
return result_mock
|
||||
|
||||
session.execute = AsyncMock(side_effect=_exec)
|
||||
svc = MessagingService(session)
|
||||
|
||||
await svc._lock_group(_GROUP_ID)
|
||||
|
||||
sql = str(
|
||||
captured[0].compile(
|
||||
dialect=postgresql.dialect(), compile_kwargs={"literal_binds": True}
|
||||
)
|
||||
)
|
||||
assert "FOR UPDATE" in sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_session_race_loser_reuses_winner_under_lock() -> None:
|
||||
"""Concurrent posts: caller A wins the race and links its session while
|
||||
caller B is between the check and the create. Caller B locks the group,
|
||||
re-reads ``active_session_id`` (now A's session), and reuses it — no second
|
||||
ACTIVE session is created (no orphan)."""
|
||||
session = AsyncMock()
|
||||
session.add = MagicMock()
|
||||
session.flush = AsyncMock()
|
||||
svc = MessagingService(session)
|
||||
|
||||
winner = MagicMock(name="winner-session", status=SessionStatus.ACTIVE)
|
||||
svc.get_group = AsyncMock(return_value=MagicMock(active_session_id=None))
|
||||
# Under the lock, the group now reflects the winner's link.
|
||||
svc._lock_group = AsyncMock(
|
||||
return_value=MagicMock(active_session_id=_WINNER_SESSION_ID)
|
||||
)
|
||||
svc.get_session = AsyncMock(return_value=winner)
|
||||
|
||||
result = await svc.create_session(SessionCreateRequest(group_id=_GROUP_ID))
|
||||
|
||||
assert result is winner
|
||||
svc._lock_group.assert_awaited_once()
|
||||
svc.get_session.assert_awaited_once()
|
||||
session.add.assert_not_called() # no orphaning INSERT
|
||||
session.flush.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_session_creates_when_no_active_under_lock() -> None:
|
||||
"""No race: under the lock there is still no active session, so create a
|
||||
new ACTIVE session and link it on the group (regression guard — the lock
|
||||
must not break the happy path)."""
|
||||
session = AsyncMock()
|
||||
session.add = MagicMock()
|
||||
session.flush = AsyncMock()
|
||||
svc = MessagingService(session)
|
||||
|
||||
locked_group = MagicMock(active_session_id=None)
|
||||
svc.get_group = AsyncMock(return_value=MagicMock(active_session_id=None))
|
||||
svc._lock_group = AsyncMock(return_value=locked_group)
|
||||
svc.get_session = AsyncMock() # should NOT be called (no active id)
|
||||
|
||||
result = await svc.create_session(SessionCreateRequest(group_id=_GROUP_ID))
|
||||
|
||||
assert isinstance(result, SessionTable)
|
||||
assert result.status == SessionStatus.ACTIVE
|
||||
svc._lock_group.assert_awaited_once()
|
||||
svc.get_session.assert_not_awaited()
|
||||
session.add.assert_called_once()
|
||||
assert session.flush.await_count >= 1
|
||||
Reference in New Issue
Block a user