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:
parent
436c39bbcc
commit
2467ea19f8
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Reference in New Issue