diff --git a/backend/packages/harness/deerflow/persistence/thread_meta/memory.py b/backend/packages/harness/deerflow/persistence/thread_meta/memory.py index ca43c7ae..769e047e 100644 --- a/backend/packages/harness/deerflow/persistence/thread_meta/memory.py +++ b/backend/packages/harness/deerflow/persistence/thread_meta/memory.py @@ -33,15 +33,19 @@ class MemoryThreadMetaStore(ThreadMetaStore): self, thread_id: str, user_id: str | None | _AutoSentinel, + workspace_id: str | None | _WorkspaceAutoSentinel, method_name: str, ) -> dict | None: - """Fetch a record and verify ownership. Returns a mutable copy, or None.""" - resolved = resolve_user_id(user_id, method_name=method_name) + """Fetch a record and verify workspace + ownership. Returns a mutable copy, or None.""" + 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) if item is None: return None 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 record @@ -73,8 +77,14 @@ class MemoryThreadMetaStore(ThreadMetaStore): await self._store.aput(THREADS_NS, thread_id, record) return record - async def get(self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO) -> dict | None: - return await self._get_owned_record(thread_id, user_id, "MemoryThreadMetaStore.get") + async def 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( self, diff --git a/backend/packages/harness/deerflow/persistence/thread_meta/sql.py b/backend/packages/harness/deerflow/persistence/thread_meta/sql.py index 8f4ee978..0798e79c 100644 --- a/backend/packages/harness/deerflow/persistence/thread_meta/sql.py +++ b/backend/packages/harness/deerflow/persistence/thread_meta/sql.py @@ -71,13 +71,18 @@ class ThreadMetaRepository(ThreadMetaStore): thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO, + workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO, ) -> dict | None: 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: - row = await session.get(ThreadMetaRow, thread_id) + row = (await session.execute(stmt)).scalar_one_or_none() if row is 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: return None return self._row_to_dict(row) diff --git a/backend/tests/test_thread_meta_workspace_filter.py b/backend/tests/test_thread_meta_workspace_filter.py index b76e870a..be9ff564 100644 --- a/backend/tests/test_thread_meta_workspace_filter.py +++ b/backend/tests/test_thread_meta_workspace_filter.py @@ -102,3 +102,53 @@ class TestCreateWorkspace: record = await repo.create("t1", workspace_id=None) assert record["workspace_id"] is None 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()