fix(deriver): lock scope membership across backfill chunk writes

SELECT ... FOR UPDATE on the active SessionPeer row so a concurrent
leave cannot commit between the membership check and the copy inserts.
Adds a concurrency test that asserts the leave blocks until commit.
This commit is contained in:
Aakash Kattelu 2026-08-31 16:31:07 -04:00
parent 436c39bbcc
commit 2467ea19f8
2 changed files with 139 additions and 8 deletions

View File

@ -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):

View File

@ -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
# ---------------------------------------------------------------------------