feat(persistence): PR6 T6.5 — Run/Feedback/RunEvent repos workspace_id
`RunRepository` (put / get / list_by_thread / delete), `FeedbackRepository` (create / get / list_by_run / list_by_thread / list_by_thread_grouped / upsert / delete / delete_by_run), and `DbRunEventStore` (put / put_batch / list_messages / list_events / list_messages_by_run / count_messages / delete_by_thread / delete_by_run) all accept `workspace_id: str | None | _AutoSentinel = AUTO`. - Write paths stamp `workspace_id` from the contextvar (same shape as the existing `user_id` stamping). For event writes the soft-read `_workspace_id_from_context()` mirrors `_user_id_from_context()` so background worker writes without a contextvar leave the column NULL — consistent with PR5's nullable-during-backfill stance. - Read paths get an extra `WHERE workspace_id = :wid` clause when the resolved value is not None. 7 new tests (3 RunRepo + 2 Feedback + 2 RunEvent) prove cross-workspace reads see zero rows. 92 existing run / feedback / event tests stay green.
This commit is contained in:
@@ -13,6 +13,13 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|||||||
|
|
||||||
from deerflow.persistence.feedback.model import FeedbackRow
|
from deerflow.persistence.feedback.model import FeedbackRow
|
||||||
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
||||||
|
from deerflow.runtime.workspace_context import AUTO as WORKSPACE_AUTO
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
_AutoSentinel as _WorkspaceAutoSentinel,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
resolve_workspace_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class FeedbackRepository:
|
class FeedbackRepository:
|
||||||
@@ -34,6 +41,7 @@ class FeedbackRepository:
|
|||||||
thread_id: str,
|
thread_id: str,
|
||||||
rating: int,
|
rating: int,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
message_id: str | None = None,
|
message_id: str | None = None,
|
||||||
comment: str | None = None,
|
comment: str | None = None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
@@ -41,11 +49,13 @@ class FeedbackRepository:
|
|||||||
if rating not in (1, -1):
|
if rating not in (1, -1):
|
||||||
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.create")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.create")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.create")
|
||||||
row = FeedbackRow(
|
row = FeedbackRow(
|
||||||
feedback_id=str(uuid.uuid4()),
|
feedback_id=str(uuid.uuid4()),
|
||||||
run_id=run_id,
|
run_id=run_id,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
user_id=resolved_user_id,
|
user_id=resolved_user_id,
|
||||||
|
workspace_id=resolved_workspace_id,
|
||||||
message_id=message_id,
|
message_id=message_id,
|
||||||
rating=rating,
|
rating=rating,
|
||||||
comment=comment,
|
comment=comment,
|
||||||
@@ -62,12 +72,16 @@ class FeedbackRepository:
|
|||||||
feedback_id: str,
|
feedback_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="FeedbackRepository.get")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.get")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.get")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
row = await session.get(FeedbackRow, feedback_id)
|
row = await session.get(FeedbackRow, feedback_id)
|
||||||
if row is None:
|
if row is None:
|
||||||
return None
|
return None
|
||||||
|
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
|
||||||
|
return None
|
||||||
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)
|
||||||
@@ -79,9 +93,13 @@ class FeedbackRepository:
|
|||||||
*,
|
*,
|
||||||
limit: int = 100,
|
limit: int = 100,
|
||||||
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="FeedbackRepository.list_by_run")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_run")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.list_by_run")
|
||||||
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id, FeedbackRow.run_id == run_id)
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id, FeedbackRow.run_id == run_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(FeedbackRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
||||||
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
||||||
@@ -95,9 +113,13 @@ class FeedbackRepository:
|
|||||||
*,
|
*,
|
||||||
limit: int = 100,
|
limit: int = 100,
|
||||||
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="FeedbackRepository.list_by_thread")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_thread")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.list_by_thread")
|
||||||
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(FeedbackRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
||||||
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
||||||
@@ -110,12 +132,16 @@ class FeedbackRepository:
|
|||||||
feedback_id: str,
|
feedback_id: str,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.delete")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.delete")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.delete")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
row = await session.get(FeedbackRow, feedback_id)
|
row = await session.get(FeedbackRow, feedback_id)
|
||||||
if row is None:
|
if row is None:
|
||||||
return False
|
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:
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
||||||
return False
|
return False
|
||||||
await session.delete(row)
|
await session.delete(row)
|
||||||
@@ -129,18 +155,22 @@ class FeedbackRepository:
|
|||||||
thread_id: str,
|
thread_id: str,
|
||||||
rating: int,
|
rating: int,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
comment: str | None = None,
|
comment: str | None = None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Create or update feedback for (thread_id, run_id, user_id). rating must be +1 or -1."""
|
"""Create or update feedback for (thread_id, run_id, user_id). rating must be +1 or -1."""
|
||||||
if rating not in (1, -1):
|
if rating not in (1, -1):
|
||||||
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.upsert")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.upsert")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.upsert")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
stmt = select(FeedbackRow).where(
|
stmt = select(FeedbackRow).where(
|
||||||
FeedbackRow.thread_id == thread_id,
|
FeedbackRow.thread_id == thread_id,
|
||||||
FeedbackRow.run_id == run_id,
|
FeedbackRow.run_id == run_id,
|
||||||
FeedbackRow.user_id == resolved_user_id,
|
FeedbackRow.user_id == resolved_user_id,
|
||||||
)
|
)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(FeedbackRow.workspace_id == resolved_workspace_id)
|
||||||
result = await session.execute(stmt)
|
result = await session.execute(stmt)
|
||||||
row = result.scalar_one_or_none()
|
row = result.scalar_one_or_none()
|
||||||
if row is not None:
|
if row is not None:
|
||||||
@@ -153,6 +183,7 @@ class FeedbackRepository:
|
|||||||
run_id=run_id,
|
run_id=run_id,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
user_id=resolved_user_id,
|
user_id=resolved_user_id,
|
||||||
|
workspace_id=resolved_workspace_id,
|
||||||
rating=rating,
|
rating=rating,
|
||||||
comment=comment,
|
comment=comment,
|
||||||
created_at=datetime.now(UTC),
|
created_at=datetime.now(UTC),
|
||||||
@@ -168,15 +199,19 @@ class FeedbackRepository:
|
|||||||
thread_id: str,
|
thread_id: str,
|
||||||
run_id: str,
|
run_id: str,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Delete the current user's feedback for a run. Returns True if a record was deleted."""
|
"""Delete the current user's feedback for a run. Returns True if a record was deleted."""
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.delete_by_run")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.delete_by_run")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.delete_by_run")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
stmt = select(FeedbackRow).where(
|
stmt = select(FeedbackRow).where(
|
||||||
FeedbackRow.thread_id == thread_id,
|
FeedbackRow.thread_id == thread_id,
|
||||||
FeedbackRow.run_id == run_id,
|
FeedbackRow.run_id == run_id,
|
||||||
FeedbackRow.user_id == resolved_user_id,
|
FeedbackRow.user_id == resolved_user_id,
|
||||||
)
|
)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(FeedbackRow.workspace_id == resolved_workspace_id)
|
||||||
result = await session.execute(stmt)
|
result = await session.execute(stmt)
|
||||||
row = result.scalar_one_or_none()
|
row = result.scalar_one_or_none()
|
||||||
if row is None:
|
if row is None:
|
||||||
@@ -190,10 +225,14 @@ class FeedbackRepository:
|
|||||||
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[str, dict]:
|
) -> dict[str, dict]:
|
||||||
"""Return feedback grouped by run_id for a thread: {run_id: feedback_dict}."""
|
"""Return feedback grouped by run_id for a thread: {run_id: feedback_dict}."""
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_thread_grouped")
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_thread_grouped")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="FeedbackRepository.list_by_thread_grouped")
|
||||||
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(FeedbackRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
|
|||||||
@@ -17,6 +17,13 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|||||||
from deerflow.persistence.run.model import RunRow
|
from deerflow.persistence.run.model import RunRow
|
||||||
from deerflow.runtime.runs.store.base import RunStore
|
from deerflow.runtime.runs.store.base import RunStore
|
||||||
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
||||||
|
from deerflow.runtime.workspace_context import AUTO as WORKSPACE_AUTO
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
_AutoSentinel as _WorkspaceAutoSentinel,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
resolve_workspace_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RunRepository(RunStore):
|
class RunRepository(RunStore):
|
||||||
@@ -70,6 +77,7 @@ class RunRepository(RunStore):
|
|||||||
thread_id,
|
thread_id,
|
||||||
assistant_id=None,
|
assistant_id=None,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
status="pending",
|
status="pending",
|
||||||
multitask_strategy="reject",
|
multitask_strategy="reject",
|
||||||
metadata=None,
|
metadata=None,
|
||||||
@@ -79,12 +87,14 @@ class RunRepository(RunStore):
|
|||||||
follow_up_to_run_id=None,
|
follow_up_to_run_id=None,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.put")
|
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.put")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="RunRepository.put")
|
||||||
now = datetime.now(UTC)
|
now = datetime.now(UTC)
|
||||||
row = RunRow(
|
row = RunRow(
|
||||||
run_id=run_id,
|
run_id=run_id,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
assistant_id=assistant_id,
|
assistant_id=assistant_id,
|
||||||
user_id=resolved_user_id,
|
user_id=resolved_user_id,
|
||||||
|
workspace_id=resolved_workspace_id,
|
||||||
status=status,
|
status=status,
|
||||||
multitask_strategy=multitask_strategy,
|
multitask_strategy=multitask_strategy,
|
||||||
metadata_json=self._safe_json(metadata) or {},
|
metadata_json=self._safe_json(metadata) or {},
|
||||||
@@ -103,12 +113,16 @@ class RunRepository(RunStore):
|
|||||||
run_id,
|
run_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.get")
|
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.get")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="RunRepository.get")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
row = await session.get(RunRow, run_id)
|
row = await session.get(RunRow, run_id)
|
||||||
if row is None:
|
if row is None:
|
||||||
return None
|
return None
|
||||||
|
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
|
||||||
|
return None
|
||||||
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)
|
||||||
@@ -118,10 +132,14 @@ class RunRepository(RunStore):
|
|||||||
thread_id,
|
thread_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
limit=100,
|
limit=100,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.list_by_thread")
|
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.list_by_thread")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="RunRepository.list_by_thread")
|
||||||
stmt = select(RunRow).where(RunRow.thread_id == thread_id)
|
stmt = select(RunRow).where(RunRow.thread_id == thread_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(RunRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(RunRow.user_id == resolved_user_id)
|
stmt = stmt.where(RunRow.user_id == resolved_user_id)
|
||||||
stmt = stmt.order_by(RunRow.created_at.desc()).limit(limit)
|
stmt = stmt.order_by(RunRow.created_at.desc()).limit(limit)
|
||||||
@@ -142,12 +160,16 @@ class RunRepository(RunStore):
|
|||||||
run_id,
|
run_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.delete")
|
resolved_user_id = resolve_user_id(user_id, method_name="RunRepository.delete")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="RunRepository.delete")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
row = await session.get(RunRow, run_id)
|
row = await session.get(RunRow, run_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)
|
||||||
|
|||||||
@@ -17,6 +17,14 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|||||||
from deerflow.persistence.models.run_event import RunEventRow
|
from deerflow.persistence.models.run_event import RunEventRow
|
||||||
from deerflow.runtime.events.store.base import RunEventStore
|
from deerflow.runtime.events.store.base import RunEventStore
|
||||||
from deerflow.runtime.user_context import AUTO, _AutoSentinel, get_current_user, resolve_user_id
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel, get_current_user, resolve_user_id
|
||||||
|
from deerflow.runtime.workspace_context import AUTO as WORKSPACE_AUTO
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
_AutoSentinel as _WorkspaceAutoSentinel,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
get_current_workspace,
|
||||||
|
resolve_workspace_id,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -86,6 +94,19 @@ class DbRunEventStore(RunEventStore):
|
|||||||
user = get_current_user()
|
user = get_current_user()
|
||||||
return str(user.id) if user is not None else None
|
return str(user.id) if user is not None else None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _workspace_id_from_context() -> str | None:
|
||||||
|
"""Soft read of workspace_id from contextvar for write paths.
|
||||||
|
|
||||||
|
Mirrors :meth:`_user_id_from_context`. Returns ``None`` (no stamp)
|
||||||
|
when no workspace is in context — typical for background worker
|
||||||
|
writes that fire outside an HTTP request. The DB column is
|
||||||
|
nullable through PR5 and becomes NOT NULL only after the alembic
|
||||||
|
0003 migration runs (verified by the backfill path).
|
||||||
|
"""
|
||||||
|
workspace = get_current_workspace()
|
||||||
|
return str(workspace.id) if workspace is not None else None
|
||||||
|
|
||||||
async def put(self, *, thread_id, run_id, event_type, category, content="", metadata=None, created_at=None): # noqa: D401
|
async def put(self, *, thread_id, run_id, event_type, category, content="", metadata=None, created_at=None): # noqa: D401
|
||||||
"""Write a single event — low-frequency path only.
|
"""Write a single event — low-frequency path only.
|
||||||
|
|
||||||
@@ -98,6 +119,7 @@ class DbRunEventStore(RunEventStore):
|
|||||||
content, metadata = self._truncate_trace(category, content, metadata)
|
content, metadata = self._truncate_trace(category, content, metadata)
|
||||||
db_content, metadata = self._content_to_db(content, metadata)
|
db_content, metadata = self._content_to_db(content, metadata)
|
||||||
user_id = self._user_id_from_context()
|
user_id = self._user_id_from_context()
|
||||||
|
workspace_id = self._workspace_id_from_context()
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
async with session.begin():
|
async with session.begin():
|
||||||
# Use FOR UPDATE to serialize seq assignment within a thread.
|
# Use FOR UPDATE to serialize seq assignment within a thread.
|
||||||
@@ -109,6 +131,7 @@ class DbRunEventStore(RunEventStore):
|
|||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
run_id=run_id,
|
run_id=run_id,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
|
workspace_id=workspace_id,
|
||||||
event_type=event_type,
|
event_type=event_type,
|
||||||
category=category,
|
category=category,
|
||||||
content=db_content,
|
content=db_content,
|
||||||
@@ -123,6 +146,7 @@ class DbRunEventStore(RunEventStore):
|
|||||||
if not events:
|
if not events:
|
||||||
return []
|
return []
|
||||||
user_id = self._user_id_from_context()
|
user_id = self._user_id_from_context()
|
||||||
|
workspace_id = self._workspace_id_from_context()
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
async with session.begin():
|
async with session.begin():
|
||||||
# Get max seq for the thread (assume all events in batch belong to same thread).
|
# Get max seq for the thread (assume all events in batch belong to same thread).
|
||||||
@@ -143,6 +167,7 @@ class DbRunEventStore(RunEventStore):
|
|||||||
thread_id=e["thread_id"],
|
thread_id=e["thread_id"],
|
||||||
run_id=e["run_id"],
|
run_id=e["run_id"],
|
||||||
user_id=e.get("user_id", user_id),
|
user_id=e.get("user_id", user_id),
|
||||||
|
workspace_id=e.get("workspace_id", workspace_id),
|
||||||
event_type=e["event_type"],
|
event_type=e["event_type"],
|
||||||
category=category,
|
category=category,
|
||||||
content=db_content,
|
content=db_content,
|
||||||
@@ -162,9 +187,13 @@ class DbRunEventStore(RunEventStore):
|
|||||||
before_seq=None,
|
before_seq=None,
|
||||||
after_seq=None,
|
after_seq=None,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_messages")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_messages")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.list_messages")
|
||||||
stmt = select(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.category == "message")
|
stmt = select(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.category == "message")
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
||||||
if before_seq is not None:
|
if before_seq is not None:
|
||||||
@@ -194,9 +223,13 @@ class DbRunEventStore(RunEventStore):
|
|||||||
event_types=None,
|
event_types=None,
|
||||||
limit=500,
|
limit=500,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_events")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_events")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.list_events")
|
||||||
stmt = select(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.run_id == run_id)
|
stmt = select(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.run_id == run_id)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
||||||
if event_types:
|
if event_types:
|
||||||
@@ -215,13 +248,17 @@ class DbRunEventStore(RunEventStore):
|
|||||||
before_seq=None,
|
before_seq=None,
|
||||||
after_seq=None,
|
after_seq=None,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_messages_by_run")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.list_messages_by_run")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.list_messages_by_run")
|
||||||
stmt = select(RunEventRow).where(
|
stmt = select(RunEventRow).where(
|
||||||
RunEventRow.thread_id == thread_id,
|
RunEventRow.thread_id == thread_id,
|
||||||
RunEventRow.run_id == run_id,
|
RunEventRow.run_id == run_id,
|
||||||
RunEventRow.category == "message",
|
RunEventRow.category == "message",
|
||||||
)
|
)
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
||||||
if before_seq is not None:
|
if before_seq is not None:
|
||||||
@@ -246,9 +283,13 @@ class DbRunEventStore(RunEventStore):
|
|||||||
thread_id,
|
thread_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.count_messages")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.count_messages")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.count_messages")
|
||||||
stmt = select(func.count()).select_from(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.category == "message")
|
stmt = select(func.count()).select_from(RunEventRow).where(RunEventRow.thread_id == thread_id, RunEventRow.category == "message")
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
stmt = stmt.where(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
stmt = stmt.where(RunEventRow.user_id == resolved_user_id)
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
@@ -259,10 +300,14 @@ class DbRunEventStore(RunEventStore):
|
|||||||
thread_id,
|
thread_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.delete_by_thread")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.delete_by_thread")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.delete_by_thread")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
count_conditions = [RunEventRow.thread_id == thread_id]
|
count_conditions = [RunEventRow.thread_id == thread_id]
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
count_conditions.append(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
count_conditions.append(RunEventRow.user_id == resolved_user_id)
|
count_conditions.append(RunEventRow.user_id == resolved_user_id)
|
||||||
count_stmt = select(func.count()).select_from(RunEventRow).where(*count_conditions)
|
count_stmt = select(func.count()).select_from(RunEventRow).where(*count_conditions)
|
||||||
@@ -278,10 +323,14 @@ class DbRunEventStore(RunEventStore):
|
|||||||
run_id,
|
run_id,
|
||||||
*,
|
*,
|
||||||
user_id: str | None | _AutoSentinel = AUTO,
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
||||||
):
|
):
|
||||||
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.delete_by_run")
|
resolved_user_id = resolve_user_id(user_id, method_name="DbRunEventStore.delete_by_run")
|
||||||
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="DbRunEventStore.delete_by_run")
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
count_conditions = [RunEventRow.thread_id == thread_id, RunEventRow.run_id == run_id]
|
count_conditions = [RunEventRow.thread_id == thread_id, RunEventRow.run_id == run_id]
|
||||||
|
if resolved_workspace_id is not None:
|
||||||
|
count_conditions.append(RunEventRow.workspace_id == resolved_workspace_id)
|
||||||
if resolved_user_id is not None:
|
if resolved_user_id is not None:
|
||||||
count_conditions.append(RunEventRow.user_id == resolved_user_id)
|
count_conditions.append(RunEventRow.user_id == resolved_user_id)
|
||||||
count_stmt = select(func.count()).select_from(RunEventRow).where(*count_conditions)
|
count_stmt = select(func.count()).select_from(RunEventRow).where(*count_conditions)
|
||||||
|
|||||||
@@ -0,0 +1,181 @@
|
|||||||
|
"""Tests for Run/Feedback/RunEvent repository workspace_id filtering (PR6 T6.5)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deerflow.runtime.workspace_context import (
|
||||||
|
reset_current_workspace,
|
||||||
|
set_current_workspace,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _init_engine(tmp_path, *, workspaces: tuple[str, ...] = ()):
|
||||||
|
from deerflow.persistence.engine import get_session_factory, init_engine
|
||||||
|
|
||||||
|
url = f"sqlite+aiosqlite:///{tmp_path / 'test.db'}"
|
||||||
|
await init_engine("sqlite", url=url, sqlite_dir=str(tmp_path))
|
||||||
|
for wid in workspaces:
|
||||||
|
await _seed_workspace(wid)
|
||||||
|
return get_session_factory()
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_workspace(wid: str) -> None:
|
||||||
|
from deerflow.persistence.engine import get_session_factory
|
||||||
|
from deerflow.persistence.workspace.model import WorkspaceRow
|
||||||
|
|
||||||
|
factory = get_session_factory()
|
||||||
|
async with factory() as session:
|
||||||
|
if await session.get(WorkspaceRow, wid) is not None:
|
||||||
|
return
|
||||||
|
now = datetime.now(UTC)
|
||||||
|
session.add(
|
||||||
|
WorkspaceRow(
|
||||||
|
id=wid,
|
||||||
|
name=f"WS {wid}",
|
||||||
|
slug=wid.replace("_", "-")[:32],
|
||||||
|
status="active",
|
||||||
|
owner_id="test-user-autouse",
|
||||||
|
created_at=now,
|
||||||
|
updated_at=now,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _cleanup():
|
||||||
|
from deerflow.persistence.engine import close_engine
|
||||||
|
|
||||||
|
await close_engine()
|
||||||
|
|
||||||
|
|
||||||
|
def _use_workspace(wid: str):
|
||||||
|
return set_current_workspace(SimpleNamespace(id=wid, role="owner"))
|
||||||
|
|
||||||
|
|
||||||
|
class TestRunRepositoryWorkspace:
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_put_records_workspace_id(self, tmp_path):
|
||||||
|
from deerflow.persistence.run import RunRepository
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha",))
|
||||||
|
repo = RunRepository(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.put("r1", thread_id="t1", user_id="alice")
|
||||||
|
record = await repo.get("r1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert record["workspace_id"] == "ws-alpha"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_filters_cross_workspace(self, tmp_path):
|
||||||
|
from deerflow.persistence.run import RunRepository
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha", "ws-beta"))
|
||||||
|
repo = RunRepository(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.put("r1", thread_id="t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
token = _use_workspace("ws-beta")
|
||||||
|
try:
|
||||||
|
assert await repo.get("r1", user_id="alice") is None
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_by_thread_filters_workspace(self, tmp_path):
|
||||||
|
from deerflow.persistence.run import RunRepository
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha", "ws-beta"))
|
||||||
|
repo = RunRepository(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.put("r1", thread_id="t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
token = _use_workspace("ws-beta")
|
||||||
|
try:
|
||||||
|
await repo.put("r2", thread_id="t1", user_id="alice")
|
||||||
|
rows = await repo.list_by_thread("t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert [r["run_id"] for r in rows] == ["r2"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFeedbackRepositoryWorkspace:
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_create_records_workspace_id(self, tmp_path):
|
||||||
|
from deerflow.persistence.feedback.sql import FeedbackRepository
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha",))
|
||||||
|
repo = FeedbackRepository(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
row = await repo.create(run_id="r1", thread_id="t1", rating=1, user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert row["workspace_id"] == "ws-alpha"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_by_thread_filters_workspace(self, tmp_path):
|
||||||
|
from deerflow.persistence.feedback.sql import FeedbackRepository
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha", "ws-beta"))
|
||||||
|
repo = FeedbackRepository(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await repo.create(run_id="r1", thread_id="t1", rating=1, user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
token = _use_workspace("ws-beta")
|
||||||
|
try:
|
||||||
|
rows = await repo.list_by_thread("t1", user_id="alice")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert rows == []
|
||||||
|
|
||||||
|
|
||||||
|
class TestRunEventStoreWorkspace:
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_put_records_workspace_id(self, tmp_path):
|
||||||
|
from deerflow.runtime.events.store.db import DbRunEventStore
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha",))
|
||||||
|
store = DbRunEventStore(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
row = await store.put(thread_id="t1", run_id="r1", event_type="msg", category="message", content="hi")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert row["workspace_id"] == "ws-alpha"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_messages_filters_cross_workspace(self, tmp_path):
|
||||||
|
from deerflow.runtime.events.store.db import DbRunEventStore
|
||||||
|
|
||||||
|
sf = await _init_engine(tmp_path, workspaces=("ws-alpha", "ws-beta"))
|
||||||
|
store = DbRunEventStore(sf)
|
||||||
|
token = _use_workspace("ws-alpha")
|
||||||
|
try:
|
||||||
|
await store.put(thread_id="t1", run_id="r1", event_type="msg", category="message", content="from-alpha")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
token = _use_workspace("ws-beta")
|
||||||
|
try:
|
||||||
|
rows = await store.list_messages("t1", user_id="test-user-autouse")
|
||||||
|
finally:
|
||||||
|
reset_current_workspace(token)
|
||||||
|
await _cleanup()
|
||||||
|
assert rows == []
|
||||||
Reference in New Issue
Block a user