diff --git a/backend/packages/harness/deerflow/persistence/thread_meta/memory.py b/backend/packages/harness/deerflow/persistence/thread_meta/memory.py index 769e047e..9afd1a75 100644 --- a/backend/packages/harness/deerflow/persistence/thread_meta/memory.py +++ b/backend/packages/harness/deerflow/persistence/thread_meta/memory.py @@ -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) diff --git a/backend/packages/harness/deerflow/persistence/thread_meta/sql.py b/backend/packages/harness/deerflow/persistence/thread_meta/sql.py index 0798e79c..bd845579 100644 --- a/backend/packages/harness/deerflow/persistence/thread_meta/sql.py +++ b/backend/packages/harness/deerflow/persistence/thread_meta/sql.py @@ -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) diff --git a/backend/tests/test_thread_meta_workspace_filter.py b/backend/tests/test_thread_meta_workspace_filter.py index be9ff564..511540fe 100644 --- a/backend/tests/test_thread_meta_workspace_filter.py +++ b/backend/tests/test_thread_meta_workspace_filter.py @@ -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"