xiaodu55 0745fb268f
fix(gateway): scope runs read endpoints by data identity, not authorization identity (#5448)
* 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.
2026-09-16 15:56:25 +08:00

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,
}