feat(persistence): PR6 T6.2 — ThreadMetaRepository.get filters by workspace_id
`get()` accepts `workspace_id: str | None | _AutoSentinel = AUTO` and moves the workspace check into the SQL WHERE so cross-workspace lookups short-circuit without loading the row. `MemoryThreadMetaStore._get_owned_record` gets the same filter for parity. The user_id check stays as a post-load comparison (preserves the existing shared-row semantics where row.user_id IS NULL means "everyone in this workspace"). Cross-workspace always returns None, never the row. 3 new tests cover the three states: in-workspace get returns the row, out-of-workspace get returns None even when user_id matches, explicit workspace_id=None bypasses (migration / CLI). Existing 22 thread_meta tests stay green.
This commit is contained in:
@@ -33,15 +33,19 @@ class MemoryThreadMetaStore(ThreadMetaStore):
|
|||||||
self,
|
self,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
user_id: str | None | _AutoSentinel,
|
user_id: str | None | _AutoSentinel,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel,
|
||||||
method_name: str,
|
method_name: str,
|
||||||
) -> dict | None:
|
) -> dict | None:
|
||||||
"""Fetch a record and verify ownership. Returns a mutable copy, or None."""
|
"""Fetch a record and verify workspace + ownership. Returns a mutable copy, or None."""
|
||||||
resolved = resolve_user_id(user_id, method_name=method_name)
|
resolved_user = resolve_user_id(user_id, method_name=method_name)
|
||||||
|
resolved_workspace = resolve_workspace_id(workspace_id, method_name=method_name)
|
||||||
item = await self._store.aget(THREADS_NS, thread_id)
|
item = await self._store.aget(THREADS_NS, thread_id)
|
||||||
if item is None:
|
if item is None:
|
||||||
return None
|
return None
|
||||||
record = dict(item.value)
|
record = dict(item.value)
|
||||||
if resolved is not None and record.get("user_id") != resolved:
|
if resolved_workspace is not None and record.get("workspace_id") != resolved_workspace:
|
||||||
|
return None
|
||||||
|
if resolved_user is not None and record.get("user_id") != resolved_user:
|
||||||
return None
|
return None
|
||||||
return record
|
return record
|
||||||
|
|
||||||
@@ -73,8 +77,14 @@ class MemoryThreadMetaStore(ThreadMetaStore):
|
|||||||
await self._store.aput(THREADS_NS, thread_id, record)
|
await self._store.aput(THREADS_NS, thread_id, record)
|
||||||
return record
|
return record
|
||||||
|
|
||||||
async def get(self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO) -> dict | None:
|
async def get(
|
||||||
return await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.get")
|
self,
|
||||||
|
thread_id: str,
|
||||||
|
*,
|
||||||
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
|
) -> dict | None:
|
||||||
|
return await self._get_owned_record(thread_id, user_id, workspace_id, "MemoryThreadMetaStore.get")
|
||||||
|
|
||||||
async def search(
|
async def search(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -71,13 +71,18 @@ 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,
|
||||||
) -> dict | None:
|
) -> dict | None:
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.get")
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.get")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.get")
|
||||||
|
stmt = select(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(ThreadMetaRow.workspace_id == resolved_workspace_id)
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
row = await session.get(ThreadMetaRow, thread_id)
|
row = (await session.execute(stmt)).scalar_one_or_none()
|
||||||
if row is None:
|
if row is None:
|
||||||
return None
|
return None
|
||||||
# Enforce owner filter unless explicitly bypassed (user_id=None).
|
# Owner filter still applies inside the workspace scope.
|
||||||
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 None
|
return None
|
||||||
return self._row_to_dict(row)
|
return self._row_to_dict(row)
|
||||||
|
|||||||
@@ -102,3 +102,53 @@ class TestCreateWorkspace:
|
|||||||
record = await repo.create("t1", workspace_id=None)
|
record = await repo.create("t1", workspace_id=None)
|
||||||
assert record["workspace_id"] is None
|
assert record["workspace_id"] is None
|
||||||
await _cleanup()
|
await _cleanup()
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetWorkspace:
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_filters_by_workspace(self, tmp_path):
|
||||||
|
"""Cross-workspace get returns None even when user_id matches."""
|
||||||
|
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:
|
||||||
|
assert await repo.get("t1", user_id="alice") is None
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_returns_row_in_same_workspace(self, tmp_path):
|
||||||
|
repo = await _make_repo(tmp_path, workspaces=("ws-alpha",))
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.create("t1", user_id="alice")
|
||||||
|
record = await repo.get("t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert record is not None
|
||||||
|
assert record["thread_id"] == "t1"
|
||||||
|
assert record["workspace_id"] == "ws-alpha"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_workspace_none_bypasses_filter(self, tmp_path):
|
||||||
|
"""Explicit workspace_id=None lets migration scripts see any row."""
|
||||||
|
repo = await _make_repo(tmp_path, workspaces=("ws-alpha",))
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.create("t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
|
||||||
|
token = _use_workspace("ws-beta")
|
||||||
|
try:
|
||||||
|
assert await repo.get("t1", user_id=None, workspace_id=None) is not None
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
|||||||
Reference in New Issue
Block a user