mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-19 19:16:17 +00:00
* fix(gateway): scope runs read endpoints by data identity, not authorization identity Trusted internal callers are authorized as a synthetic internal user (id "default", or the make_safe_user_id-normalized owner with an owner header), while start_run stamps run rows with the raw trusted-owner value. list_runs / get_run / list_runs_page filtered by the authorization identity, so the store-side user filter never matched and internal callers always saw an empty runs list (or 404) for threads they are authorized to read. The three read endpoints now resolve their filter id through _run_scope_user_id: internal-role callers skip the per-user filter (thread visibility is already authorized by owner_check=True), and browser/API sessions keep the existing per-user filter unchanged. Regression tests cover both identities across the three endpoints (list, keyset page, single get) with a MemoryRunStore seeded with mixed-owner rows; without the fix the four internal-caller cases fail while the browser-session isolation case passes. * fix(gateway): route the message read endpoints through the same data-identity scoping Review follow-up on #5448: list_thread_messages and list_thread_messages_page resolved get_current_user and passed it as the data filter to the event-store scan, hidden-run lookups, turn-duration injection and the feedback queries — the same authorization-vs-data identity conflation fixed for the runs endpoints, leaving the #5437 empty-read symptom in place for lossy owner values. Both endpoints now resolve their filter id through _run_scope_user_id as well. Regression tests extend to the two message endpoints, asserting the resolved filter identity at the runs-store and feedback-repo boundaries (None for internal callers, the session user id for browser sessions). * fix(feedback): deterministic per-run collapse for unfiltered feedback reads Review follow-up on #5448: with _run_scope_user_id returning None for internal callers, the feedback lookups now receive an explicit-None user id, which skips the user_id WHERE in FeedbackRepository. On shared/NULL-owner threads several browser users can hold feedback on the same run, and list_by_thread_grouped / list_by_run_ids collapsed rows per run_id via a dict comprehension over unordered results — the feedback attached to the last AI message would be an arbitrary user's row. Both methods now order by created_at ASC with feedback_id as the tie-break, so the collapse deterministically keeps the most recently created feedback. _run_scope_user_id's docstring now documents that the resolved id also scopes feedback and event-store reads, not just run rows. Regression test seeds multi-user feedback on one run and asserts the collapse outcome is stable across repeated unfiltered reads. * docs(feedback): the collapse keeps the most recently written feedback created_at is refreshed on upsert, so the surviving row per run is the most recently written (created or updated), not the most recently created — align both docstrings with the ordering key's actual semantics.
256 lines
9.9 KiB
Python
256 lines
9.9 KiB
Python
"""SQLAlchemy-backed feedback storage.
|
|
|
|
Each method acquires its own short-lived session.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import uuid
|
|
from datetime import UTC, datetime
|
|
|
|
from sqlalchemy import case, func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.feedback.model import FeedbackRow
|
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
|
from deerflow.utils.time import coerce_iso
|
|
|
|
|
|
class FeedbackRepository:
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._sf = session_factory
|
|
|
|
@staticmethod
|
|
def _row_to_dict(row: FeedbackRow) -> dict:
|
|
d = row.to_dict()
|
|
val = d.get("created_at")
|
|
if isinstance(val, datetime):
|
|
# SQLite drops tzinfo on read; normalize via ``coerce_iso`` so output is always tz-aware.
|
|
d["created_at"] = coerce_iso(val)
|
|
return d
|
|
|
|
async def create(
|
|
self,
|
|
*,
|
|
run_id: str,
|
|
thread_id: str,
|
|
rating: int,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
message_id: str | None = None,
|
|
comment: str | None = None,
|
|
) -> dict:
|
|
"""Create a feedback record. rating must be +1 or -1."""
|
|
if rating not in (1, -1):
|
|
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.create")
|
|
row = FeedbackRow(
|
|
feedback_id=str(uuid.uuid4()),
|
|
run_id=run_id,
|
|
thread_id=thread_id,
|
|
user_id=resolved_user_id,
|
|
message_id=message_id,
|
|
rating=rating,
|
|
comment=comment,
|
|
created_at=datetime.now(UTC),
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return self._row_to_dict(row)
|
|
|
|
async def get(
|
|
self,
|
|
feedback_id: str,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> dict | None:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.get")
|
|
async with self._sf() as session:
|
|
row = await session.get(FeedbackRow, feedback_id)
|
|
if row is None:
|
|
return None
|
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
|
return None
|
|
return self._row_to_dict(row)
|
|
|
|
async def list_by_run(
|
|
self,
|
|
thread_id: str,
|
|
run_id: str,
|
|
*,
|
|
limit: int = 100,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> list[dict]:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_run")
|
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id, FeedbackRow.run_id == run_id)
|
|
if resolved_user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
|
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return [self._row_to_dict(r) for r in result.scalars()]
|
|
|
|
async def list_by_thread(
|
|
self,
|
|
thread_id: str,
|
|
*,
|
|
limit: int = 100,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> list[dict]:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_thread")
|
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
|
if resolved_user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
|
stmt = stmt.order_by(FeedbackRow.created_at.asc()).limit(limit)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return [self._row_to_dict(r) for r in result.scalars()]
|
|
|
|
async def delete(
|
|
self,
|
|
feedback_id: str,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> bool:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.delete")
|
|
async with self._sf() as session:
|
|
row = await session.get(FeedbackRow, feedback_id)
|
|
if row is None:
|
|
return False
|
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def upsert(
|
|
self,
|
|
*,
|
|
run_id: str,
|
|
thread_id: str,
|
|
rating: int,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
comment: str | None = None,
|
|
) -> dict:
|
|
"""Create or update feedback for (thread_id, run_id, user_id). rating must be +1 or -1."""
|
|
if rating not in (1, -1):
|
|
raise ValueError(f"rating must be +1 or -1, got {rating}")
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.upsert")
|
|
async with self._sf() as session:
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == thread_id,
|
|
FeedbackRow.run_id == run_id,
|
|
FeedbackRow.user_id == resolved_user_id,
|
|
)
|
|
result = await session.execute(stmt)
|
|
row = result.scalar_one_or_none()
|
|
if row is not None:
|
|
row.rating = rating
|
|
row.comment = comment
|
|
row.created_at = datetime.now(UTC)
|
|
else:
|
|
row = FeedbackRow(
|
|
feedback_id=str(uuid.uuid4()),
|
|
run_id=run_id,
|
|
thread_id=thread_id,
|
|
user_id=resolved_user_id,
|
|
rating=rating,
|
|
comment=comment,
|
|
created_at=datetime.now(UTC),
|
|
)
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return self._row_to_dict(row)
|
|
|
|
async def delete_by_run(
|
|
self,
|
|
*,
|
|
thread_id: str,
|
|
run_id: str,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> bool:
|
|
"""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")
|
|
async with self._sf() as session:
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == thread_id,
|
|
FeedbackRow.run_id == run_id,
|
|
FeedbackRow.user_id == resolved_user_id,
|
|
)
|
|
result = await session.execute(stmt)
|
|
row = result.scalar_one_or_none()
|
|
if row is None:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def list_by_thread_grouped(
|
|
self,
|
|
thread_id: str,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> dict[str, dict]:
|
|
"""Return feedback grouped by run_id for a thread: {run_id: feedback_dict}.
|
|
|
|
With an explicit ``None`` user id (unfiltered reads) several users may
|
|
hold feedback on the same run, so order deterministically — the
|
|
per-run collapse below keeps the last row per ``run_id``, i.e. the
|
|
most recently written feedback (``created_at`` is refreshed on
|
|
update), with ``feedback_id`` breaking ties.
|
|
"""
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_thread_grouped")
|
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
|
if resolved_user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
|
stmt = stmt.order_by(FeedbackRow.created_at.asc(), FeedbackRow.feedback_id.asc())
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return {row.run_id: self._row_to_dict(row) for row in result.scalars()}
|
|
|
|
async def list_by_run_ids(
|
|
self,
|
|
thread_id: str,
|
|
run_ids: set[str],
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
) -> dict[str, dict]:
|
|
"""Return feedback for only the selected runs in one thread.
|
|
|
|
Same deterministic ordering as :meth:`list_by_thread_grouped`: with an
|
|
explicit ``None`` user id the per-run collapse keeps the most recently
|
|
written feedback (``created_at`` is refreshed on update), ties broken
|
|
by ``feedback_id``.
|
|
"""
|
|
if not run_ids:
|
|
return {}
|
|
resolved_user_id = resolve_user_id(user_id, method_name="FeedbackRepository.list_by_run_ids")
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == thread_id,
|
|
FeedbackRow.run_id.in_(run_ids),
|
|
)
|
|
if resolved_user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == resolved_user_id)
|
|
stmt = stmt.order_by(FeedbackRow.created_at.asc(), FeedbackRow.feedback_id.asc())
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return {row.run_id: self._row_to_dict(row) for row in result.scalars()}
|
|
|
|
async def aggregate_by_run(self, thread_id: str, run_id: str) -> dict:
|
|
"""Aggregate feedback stats for a run using database-side counting."""
|
|
stmt = select(
|
|
func.count().label("total"),
|
|
func.coalesce(func.sum(case((FeedbackRow.rating == 1, 1), else_=0)), 0).label("positive"),
|
|
func.coalesce(func.sum(case((FeedbackRow.rating == -1, 1), else_=0)), 0).label("negative"),
|
|
).where(FeedbackRow.thread_id == thread_id, FeedbackRow.run_id == run_id)
|
|
async with self._sf() as session:
|
|
row = (await session.execute(stmt)).one()
|
|
return {
|
|
"run_id": run_id,
|
|
"total": row.total,
|
|
"positive": row.positive,
|
|
"negative": row.negative,
|
|
}
|