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:
@@ -94,13 +94,17 @@ class MemoryThreadMetaStore(ThreadMetaStore):
|
|||||||
limit: int = 100,
|
limit: int = 100,
|
||||||
offset: int = 0,
|
offset: int = 0,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="MemoryThreadMetaStore.search")
|
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] = {}
|
filter_dict: dict[str, Any] = {}
|
||||||
if metadata:
|
if metadata:
|
||||||
filter_dict.update(metadata)
|
filter_dict.update(metadata)
|
||||||
if status:
|
if status:
|
||||||
filter_dict["status"] = 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:
|
if resolved_user_id is not None:
|
||||||
filter_dict["user_id"] = resolved_user_id
|
filter_dict["user_id"] = resolved_user_id
|
||||||
|
|
||||||
@@ -121,24 +125,45 @@ class MemoryThreadMetaStore(ThreadMetaStore):
|
|||||||
return True
|
return True
|
||||||
return record_user_id == user_id
|
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:
|
async def update_display_name(
|
||||||
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.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:
|
if record is None:
|
||||||
return
|
return
|
||||||
record["display_name"] = display_name
|
record["display_name"] = display_name
|
||||||
record["updated_at"] = now_iso()
|
record["updated_at"] = now_iso()
|
||||||
await self._store.aput(THREADS_NS, thread_id, record)
|
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:
|
async def update_status(
|
||||||
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.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:
|
if record is None:
|
||||||
return
|
return
|
||||||
record["status"] = status
|
record["status"] = status
|
||||||
record["updated_at"] = now_iso()
|
record["updated_at"] = now_iso()
|
||||||
await self._store.aput(THREADS_NS, thread_id, record)
|
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:
|
async def update_metadata(
|
||||||
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.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:
|
if record is None:
|
||||||
return
|
return
|
||||||
merged = dict(record.get("metadata") or {})
|
merged = dict(record.get("metadata") or {})
|
||||||
@@ -147,8 +172,14 @@ class MemoryThreadMetaStore(ThreadMetaStore):
|
|||||||
record["updated_at"] = now_iso()
|
record["updated_at"] = now_iso()
|
||||||
await self._store.aput(THREADS_NS, thread_id, record)
|
await self._store.aput(THREADS_NS, thread_id, record)
|
||||||
|
|
||||||
async def delete(self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO) -> None:
|
async def delete(
|
||||||
record = await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.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:
|
if record is None:
|
||||||
return
|
return
|
||||||
await self._store.adelete(THREADS_NS, thread_id)
|
await self._store.adelete(THREADS_NS, thread_id)
|
||||||
|
|||||||
@@ -125,14 +125,19 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
limit: int = 100,
|
limit: int = 100,
|
||||||
offset: int = 0,
|
offset: int = 0,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
"""Search threads with optional metadata and status filters.
|
"""Search threads with optional metadata and status filters.
|
||||||
|
|
||||||
Owner filter is enforced by default: caller must be in a user
|
Both workspace and owner filters are enforced by default. Pass
|
||||||
context. Pass ``user_id=None`` to bypass (migration/CLI).
|
``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_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())
|
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:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(ThreadMetaRow.user_id == resolved_user_id)
|
stmt = stmt.where(ThreadMetaRow.user_id == resolved_user_id)
|
||||||
if status:
|
if status:
|
||||||
@@ -154,12 +159,22 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
result = await session.execute(stmt)
|
result = await session.execute(stmt)
|
||||||
return [self._row_to_dict(r) for r in result.scalars()]
|
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:
|
async def _check_ownership(
|
||||||
"""Return True if the row exists and is owned (or filter bypassed)."""
|
self,
|
||||||
if resolved_user_id is None:
|
session: AsyncSession,
|
||||||
return True # explicit bypass
|
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)
|
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(
|
async def update_display_name(
|
||||||
self,
|
self,
|
||||||
@@ -167,11 +182,13 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
display_name: str,
|
display_name: str,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Update the display_name (title) for a thread."""
|
"""Update the display_name (title) for a thread."""
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_display_name")
|
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:
|
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
|
return
|
||||||
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(display_name=display_name, updated_at=datetime.now(UTC)))
|
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(display_name=display_name, updated_at=datetime.now(UTC)))
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -182,10 +199,12 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
status: str,
|
status: str,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> None:
|
) -> None:
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_status")
|
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:
|
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
|
return
|
||||||
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(status=status, updated_at=datetime.now(UTC)))
|
await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(status=status, updated_at=datetime.now(UTC)))
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -196,18 +215,22 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
metadata: dict,
|
metadata: dict,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Merge ``metadata`` into ``metadata_json``.
|
"""Merge ``metadata`` into ``metadata_json``.
|
||||||
|
|
||||||
Read-modify-write inside a single session/transaction so concurrent
|
Read-modify-write inside a single session/transaction so concurrent
|
||||||
callers see consistent state. No-op if the row does not exist or
|
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_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:
|
async with self._sf() as session:
|
||||||
row = await session.get(ThreadMetaRow, thread_id)
|
row = await session.get(ThreadMetaRow, thread_id)
|
||||||
if row is None:
|
if row is None:
|
||||||
return
|
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:
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
||||||
return
|
return
|
||||||
merged = dict(row.metadata_json or {})
|
merged = dict(row.metadata_json or {})
|
||||||
@@ -221,12 +244,16 @@ class ThreadMetaRepository(ThreadMetaStore):
|
|||||||
thread_id: str,
|
thread_id: str,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> None:
|
) -> None:
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.delete")
|
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:
|
async with self._sf() as session:
|
||||||
row = await session.get(ThreadMetaRow, thread_id)
|
row = await session.get(ThreadMetaRow, thread_id)
|
||||||
if row is None:
|
if row is None:
|
||||||
return
|
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:
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
||||||
return
|
return
|
||||||
await session.delete(row)
|
await session.delete(row)
|
||||||
|
|||||||
@@ -152,3 +152,108 @@ class TestGetWorkspace:
|
|||||||
finally:
|
finally:
|
||||||
reset_current_workspace(token)
|
reset_current_workspace(token)
|
||||||
await _cleanup()
|
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"
|
||||||
|
|||||||
Reference in New Issue
Block a user