feat(persistence): PR6 T6.3 — search/update_*/delete workspace_id sentinel

Every remaining ThreadMetaStore method gains a `workspace_id` keyword
mirroring `user_id`:

- `search()` adds WHERE workspace_id and (for the memory impl) folds it
  into the BaseStore filter dict.
- `update_display_name` / `update_status` / `update_metadata` / `delete`
  no-op if the row lives in a different workspace. The SQL helper
  `_check_ownership()` was widened to do both checks in one pass.
- `MemoryThreadMetaStore._get_owned_record()` likewise takes both ids.

5 new isolation tests prove writes from workspace B against workspace A's
thread are silently dropped (no row mutation when the caller re-reads
from workspace A). 38 existing thread_meta / owner / memory-store tests
still pass.
This commit is contained in:
1445043649
2026-05-13 17:20:40 +08:00
parent 296a4f1950
commit 28ad6c2b0b
3 changed files with 181 additions and 18 deletions
@@ -94,13 +94,17 @@ class MemoryThreadMetaStore(ThreadMetaStore):
limit: int = 100,
offset: int = 0,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> list[dict]:
resolved_user_id = resolve_user_id(user_id, method_name="MemoryThreadMetaStore.search")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="MemoryThreadMetaStore.search")
filter_dict: dict[str, Any] = {}
if metadata:
filter_dict.update(metadata)
if status:
filter_dict["status"] = status
if resolved_workspace_id is not None:
filter_dict["workspace_id"] = resolved_workspace_id
if resolved_user_id is not None:
filter_dict["user_id"] = resolved_user_id
@@ -121,24 +125,45 @@ class MemoryThreadMetaStore(ThreadMetaStore):
return True
return record_user_id == user_id
async def update_display_name(self, thread_id: str, display_name: str, *, user_id: str | None | _AutoSentinel = AUTO) -> None:
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.update_display_name")
async def update_display_name(
self,
thread_id: str,
display_name: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
record = await self._get_owned_record(thread_id, user_id, workspace_id, "MemoryThreadMetaStore.update_display_name")
if record is None:
return
record["display_name"] = display_name
record["updated_at"] = now_iso()
await self._store.aput(THREADS_NS, thread_id, record)
async def update_status(self, thread_id: str, status: str, *, user_id: str | None | _AutoSentinel = AUTO) -> None:
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.update_status")
async def update_status(
self,
thread_id: str,
status: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
record = await self._get_owned_record(thread_id, user_id, workspace_id, "MemoryThreadMetaStore.update_status")
if record is None:
return
record["status"] = status
record["updated_at"] = now_iso()
await self._store.aput(THREADS_NS, thread_id, record)
async def update_metadata(self, thread_id: str, metadata: dict, *, user_id: str | None | _AutoSentinel = AUTO) -> None:
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.update_metadata")
async def update_metadata(
self,
thread_id: str,
metadata: dict,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
record = await self._get_owned_record(thread_id, user_id, workspace_id, "MemoryThreadMetaStore.update_metadata")
if record is None:
return
merged = dict(record.get("metadata") or {})
@@ -147,8 +172,14 @@ class MemoryThreadMetaStore(ThreadMetaStore):
record["updated_at"] = now_iso()
await self._store.aput(THREADS_NS, thread_id, record)
async def delete(self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO) -> None:
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.delete")
async def delete(
self,
thread_id: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
record = await self._get_owned_record(thread_id, user_id, workspace_id, "MemoryThreadMetaStore.delete")
if record is None:
return
await self._store.adelete(THREADS_NS, thread_id)
@@ -125,14 +125,19 @@ class ThreadMetaRepository(ThreadMetaStore):
limit: int = 100,
offset: int = 0,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> list[dict]:
"""Search threads with optional metadata and status filters.
Owner filter is enforced by default: caller must be in a user
context. Pass ``user_id=None`` to bypass (migration/CLI).
Both workspace and owner filters are enforced by default. Pass
``workspace_id=None`` and / or ``user_id=None`` to bypass for
migration / CLI paths.
"""
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.search")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.search")
stmt = select(ThreadMetaRow).order_by(ThreadMetaRow.updated_at.desc())
if resolved_workspace_id is not None:
stmt = stmt.where(ThreadMetaRow.workspace_id == resolved_workspace_id)
if resolved_user_id is not None:
stmt = stmt.where(ThreadMetaRow.user_id == resolved_user_id)
if status:
@@ -154,12 +159,22 @@ class ThreadMetaRepository(ThreadMetaStore):
result = await session.execute(stmt)
return [self._row_to_dict(r) for r in result.scalars()]
async def _check_ownership(self, session: AsyncSession, thread_id: str, resolved_user_id: str | None) -> bool:
"""Return True if the row exists and is owned (or filter bypassed)."""
if resolved_user_id is None:
return True # explicit bypass
async def _check_ownership(
self,
session: AsyncSession,
thread_id: str,
resolved_user_id: str | None,
resolved_workspace_id: str | None,
) -> bool:
"""Return True if the row exists, is in scope, and is owned (or filter bypassed)."""
row = await session.get(ThreadMetaRow, thread_id)
return row is not None and row.user_id == resolved_user_id
if row is None:
return False
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
return False
if resolved_user_id is not None and row.user_id != resolved_user_id:
return False
return True
async def update_display_name(
self,
@@ -167,11 +182,13 @@ class ThreadMetaRepository(ThreadMetaStore):
display_name: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
"""Update the display_name (title) for a thread."""
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_display_name")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.update_display_name")
async with self._sf() as session:
if not await self._check_ownership(session, thread_id, resolved_user_id):
if not await self._check_ownership(session, thread_id, resolved_user_id, resolved_workspace_id):
return
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(display_name=display_name, updated_at=datetime.now(UTC)))
await session.commit()
@@ -182,10 +199,12 @@ class ThreadMetaRepository(ThreadMetaStore):
status: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_status")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.update_status")
async with self._sf() as session:
if not await self._check_ownership(session, thread_id, resolved_user_id):
if not await self._check_ownership(session, thread_id, resolved_user_id, resolved_workspace_id):
return
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(status=status, updated_at=datetime.now(UTC)))
await session.commit()
@@ -196,18 +215,22 @@ class ThreadMetaRepository(ThreadMetaStore):
metadata: dict,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
"""Merge ``metadata`` into ``metadata_json``.
Read-modify-write inside a single session/transaction so concurrent
callers see consistent state. No-op if the row does not exist or
the user_id check fails.
the workspace / user check fails.
"""
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_metadata")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.update_metadata")
async with self._sf() as session:
row = await session.get(ThreadMetaRow, thread_id)
if row is None:
return
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
return
if resolved_user_id is not None and row.user_id != resolved_user_id:
return
merged = dict(row.metadata_json or {})
@@ -221,12 +244,16 @@ class ThreadMetaRepository(ThreadMetaStore):
thread_id: str,
*,
user_id: str | None | _AutoSentinel = AUTO,
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
) -> None:
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.delete")
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.delete")
async with self._sf() as session:
row = await session.get(ThreadMetaRow, thread_id)
if row is None:
return
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
return
if resolved_user_id is not None and row.user_id != resolved_user_id:
return
await session.delete(row)
@@ -152,3 +152,108 @@ class TestGetWorkspace:
finally:
reset_current_workspace(token)
await _cleanup()
class TestSearchUpdateDeleteWorkspace:
@pytest.mark.anyio
async def test_search_only_returns_current_workspace(self, tmp_path):
repo = await _make_repo(tmp_path, workspaces=("ws-alpha", "ws-beta"))
token = _use_workspace("ws-alpha")
try:
await repo.create("t1", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-beta")
try:
await repo.create("t2", user_id="alice")
rows = await repo.search(user_id="alice")
finally:
reset_current_workspace(token)
await _cleanup()
ids = {r["thread_id"] for r in rows}
assert ids == {"t2"}
@pytest.mark.anyio
async def test_update_status_blocked_across_workspace(self, tmp_path):
repo = await _make_repo(tmp_path, workspaces=("ws-alpha", "ws-beta"))
token = _use_workspace("ws-alpha")
try:
await repo.create("t1", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-beta")
try:
await repo.update_status("t1", "busy", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-alpha")
try:
row = await repo.get("t1", user_id="alice")
finally:
reset_current_workspace(token)
await _cleanup()
assert row["status"] == "idle"
@pytest.mark.anyio
async def test_update_display_name_blocked_across_workspace(self, tmp_path):
repo = await _make_repo(tmp_path, workspaces=("ws-alpha", "ws-beta"))
token = _use_workspace("ws-alpha")
try:
await repo.create("t1", user_id="alice", display_name="A")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-beta")
try:
await repo.update_display_name("t1", "B", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-alpha")
try:
row = await repo.get("t1", user_id="alice")
finally:
reset_current_workspace(token)
await _cleanup()
assert row["display_name"] == "A"
@pytest.mark.anyio
async def test_update_metadata_blocked_across_workspace(self, tmp_path):
repo = await _make_repo(tmp_path, workspaces=("ws-alpha", "ws-beta"))
token = _use_workspace("ws-alpha")
try:
await repo.create("t1", user_id="alice", metadata={"k": "alpha"})
finally:
reset_current_workspace(token)
token = _use_workspace("ws-beta")
try:
await repo.update_metadata("t1", {"k": "beta"}, user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-alpha")
try:
row = await repo.get("t1", user_id="alice")
finally:
reset_current_workspace(token)
await _cleanup()
assert row["metadata"] == {"k": "alpha"}
@pytest.mark.anyio
async def test_delete_blocked_across_workspace(self, tmp_path):
repo = await _make_repo(tmp_path, workspaces=("ws-alpha", "ws-beta"))
token = _use_workspace("ws-alpha")
try:
await repo.create("t1", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-beta")
try:
await repo.delete("t1", user_id="alice")
finally:
reset_current_workspace(token)
token = _use_workspace("ws-alpha")
try:
row = await repo.get("t1", user_id="alice")
finally:
reset_current_workspace(token)
await _cleanup()
assert row is not None and row["thread_id"] == "t1"