diff --git a/src/deriver/scope_backfill.py b/src/deriver/scope_backfill.py index aef32c1d..9be082f9 100644 --- a/src/deriver/scope_backfill.py +++ b/src/deriver/scope_backfill.py @@ -36,7 +36,6 @@ from sqlalchemy.sql.functions import func from src import crud, models from src.config import settings from src.crud.scope import ScopeBackfillState -from src.crud.session import is_peer_in_session from src.dependencies import tracked_db from src.embedding_client import embedding_client from src.schemas import DreamType @@ -338,12 +337,18 @@ async def _copy_chunk( touched_observed = {spec.observed for spec in plans} new_rows: list[models.Document] = [] async with tracked_db("scope_backfill.write") as db: - # scope_backfill and scope_removal carry different work-unit keys, so - # nothing orders them: a removal enqueued right after the add (or one - # that landed while phase 2 was embedding) can sweep the scope before - # these copies exist. Re-checking membership here, in the transaction - # that inserts, keeps a removed session from being copied back in. - if not await is_peer_in_session(db, workspace_name, session_name, scope_peer): + # Row-lock active membership for this txn so a concurrent leave + # (``left_at``) cannot commit between the check and the inserts. + membership = await db.scalar( + select(models.SessionPeer.peer_name) + .where(models.SessionPeer.workspace_name == workspace_name) + .where(models.SessionPeer.session_name == session_name) + .where(models.SessionPeer.peer_name == scope_peer) + .where(models.SessionPeer.left_at.is_(None)) + .with_for_update() + .limit(1) + ) + if membership is None: return False for observed in sorted(touched_observed): diff --git a/tests/deriver/test_scope_backfill.py b/tests/deriver/test_scope_backfill.py index 426ba0fa..67cf744a 100644 --- a/tests/deriver/test_scope_backfill.py +++ b/tests/deriver/test_scope_backfill.py @@ -22,15 +22,19 @@ engine (see ``mock_tracked_db_context`` in conftest.py) — a different connection that cannot see another session's uncommitted writes. """ +import asyncio +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager from typing import Any import pytest from fastapi.testclient import TestClient from nanoid import generate as generate_nanoid from sqlalchemy import func, select, update -from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker from src import crud, models +from src.deriver import scope_backfill as scope_backfill_mod from src.deriver.scope_backfill import ( COPIED_FROM_KEY, process_scope_backfill, @@ -415,6 +419,128 @@ async def test_backfill_skips_a_session_that_left_the_scope( assert session_name not in peer.internal_metadata.get("backfill_status", {}) +@pytest.mark.asyncio +async def test_copy_chunk_membership_lock_blocks_leave_until_write_commits( + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + db_engine: AsyncEngine, + monkeypatch: pytest.MonkeyPatch, +): + """A concurrent leave cannot commit between membership check and inserts.""" + test_workspace, sender = sample_data + workspace_name = test_workspace.name + scope_name = str(generate_nanoid()) + scope_peer = await _create_scope_peer(db_session, workspace_name, scope_name) + session = await _create_session(db_session, workspace_name) + await _join_scope(db_session, workspace_name, session.name, scope_peer.name) + await _create_collection( + db_session, workspace_name, observer=sender.name, observed=sender.name + ) + await _create_collection( + db_session, workspace_name, observer=scope_peer.name, observed=sender.name + ) + source = await _create_document( + db_session, + workspace_name, + observer=sender.name, + observed=sender.name, + session_name=session.name, + content="locked membership fact", + ) + + factory = async_sessionmaker(bind=db_engine, expire_on_commit=False) + leave_finished = asyncio.Event() + leave_task_box: dict[str, asyncio.Task[None]] = {} + + async def concurrent_leave() -> None: + async with factory() as leave_db: + await leave_db.execute( + update(models.SessionPeer) + .where( + models.SessionPeer.workspace_name == workspace_name, + models.SessionPeer.session_name == session.name, + models.SessionPeer.peer_name == scope_peer.name, + models.SessionPeer.left_at.is_(None), + ) + .values(left_at=func.now()) + ) + await leave_db.commit() + leave_finished.set() + + original_tracked_db = scope_backfill_mod.tracked_db # pyright: ignore[reportPrivateLocalImportUsage] + + @asynccontextmanager + async def tracked_db_with_leave_race( + operation_name: str | None = None, *, read_only: bool = False + ) -> AsyncGenerator[AsyncSession]: + async with original_tracked_db(operation_name, read_only=read_only) as db: + if operation_name == "scope_backfill.write": + real_scalar = db.scalar + raced = False + + async def scalar_then_race(statement: Any, *args: Any, **kwargs: Any): + nonlocal raced + result = await real_scalar(statement, *args, **kwargs) + if not raced and result is not None: + raced = True + leave_task_box["task"] = asyncio.create_task(concurrent_leave()) + # Leave's UPDATE must block on this txn's row lock. + for _ in range(50): + await asyncio.sleep(0.01) + if leave_task_box["task"].done(): + break + assert not leave_task_box["task"].done() + return result + + db.scalar = scalar_then_race # type: ignore[method-assign] + yield db + + monkeypatch.setattr(scope_backfill_mod, "tracked_db", tracked_db_with_leave_race) + + ok = await scope_backfill_mod._copy_chunk( # pyright: ignore[reportPrivateUsage] + workspace_name, + scope_peer.name, + session.name, + [ + scope_backfill_mod._CopySpec( # pyright: ignore[reportPrivateUsage] + observed=sender.name, + source_id=source.id, + content=source.content, + embedding=None, + internal_metadata={}, + times_derived=1, + source_ids=None, + session_name=session.name, + ) + ], + store_in_postgres=True, + ) + assert ok is True + + leave_task = leave_task_box["task"] + await asyncio.wait_for(leave_task, timeout=2.0) + assert leave_finished.is_set() + + copies = await _get_docs( + db_session, + workspace_name, + observer=scope_peer.name, + observed=sender.name, + include_deleted=False, + ) + assert len(copies) == 1 + assert copies[0].internal_metadata.get(COPIED_FROM_KEY) == source.id + + membership = await db_session.scalar( + select(models.SessionPeer.left_at).where( + models.SessionPeer.workspace_name == workspace_name, + models.SessionPeer.session_name == session.name, + models.SessionPeer.peer_name == scope_peer.name, + ) + ) + assert membership is not None + + # --------------------------------------------------------------------------- # 3. Multi-peer session # ---------------------------------------------------------------------------