mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-27 00:17:53 +00:00
refactor(gateway): route feedback through the domain service
The composition root wires FeedbackService with its SQL adapter; routers
become protocol translation only (domain errors map to 400/404/409).
BREAKING CHANGE: removes five feedback endpoints with no frontend
consumers (POST, GET list, GET stats, DELETE by id under /threads, and
GET /runs/{id}/feedback). PUT/DELETE now follow the mainstream chat UX:
idempotent upsert plus retract, echoed through the message-list payload.
This commit is contained in:
parent
86e4baccf8
commit
70613f8996
@ -29,7 +29,7 @@ from langgraph.types import Checkpointer
|
||||
|
||||
from deerflow.community.browser_automation.session import browser_multi_worker_error
|
||||
from deerflow.config.app_config import AppConfig, get_app_config
|
||||
from deerflow.persistence.feedback import FeedbackRepository
|
||||
from deerflow.domain.feedback import FeedbackService
|
||||
from deerflow.runtime import RunContext, RunManager, StreamBridge
|
||||
from deerflow.runtime.events.store.base import RunEventStore
|
||||
from deerflow.runtime.runs.store.base import RunStore
|
||||
@ -322,16 +322,25 @@ async def langgraph_runtime(app: FastAPI, startup_config: AppConfig) -> AsyncGen
|
||||
# Initialize repositories — one get_session_factory() call for all.
|
||||
sf = get_session_factory()
|
||||
if sf is not None:
|
||||
from deerflow.persistence.feedback import FeedbackRepository
|
||||
from app.infra.persistence.feedback import (
|
||||
RunStoreRunLookup,
|
||||
SqlFeedbackRepository,
|
||||
)
|
||||
from deerflow.persistence.run import RunRepository
|
||||
|
||||
app.state.run_store = RunRepository(sf)
|
||||
app.state.feedback_repo = FeedbackRepository(sf)
|
||||
# Hexagonal feedback slice: the service (input port) is wired with
|
||||
# its SQL adapter here — the composition root is the only place
|
||||
# adapters are instantiated.
|
||||
app.state.feedback_service = FeedbackService(
|
||||
repository=SqlFeedbackRepository(sf),
|
||||
runs=RunStoreRunLookup(app.state.run_store),
|
||||
)
|
||||
else:
|
||||
from deerflow.runtime.runs.store.memory import MemoryRunStore
|
||||
|
||||
app.state.run_store = MemoryRunStore()
|
||||
app.state.feedback_repo = None
|
||||
app.state.feedback_service = None # memory backend → 503, as before
|
||||
|
||||
from deerflow.persistence.thread_meta import make_thread_store
|
||||
|
||||
@ -417,7 +426,7 @@ get_stream_bridge: Callable[[Request], StreamBridge] = _require("stream_bridge",
|
||||
get_run_manager: Callable[[Request], RunManager] = _require("run_manager", "Run manager")
|
||||
get_checkpointer: Callable[[Request], Checkpointer] = _require("checkpointer", "Checkpointer")
|
||||
get_run_event_store: Callable[[Request], RunEventStore] = _require("run_event_store", "Run event store")
|
||||
get_feedback_repo: Callable[[Request], FeedbackRepository] = _require("feedback_repo", "Feedback")
|
||||
get_feedback_service: Callable[[Request], FeedbackService] = _require("feedback_service", "Feedback")
|
||||
get_run_store: Callable[[Request], RunStore] = _require("run_store", "Run store")
|
||||
|
||||
|
||||
|
||||
@ -1,19 +1,32 @@
|
||||
"""Feedback endpoints — create, list, stats, delete.
|
||||
"""Feedback endpoints — rate a run, retract a rating.
|
||||
|
||||
Allows users to submit thumbs-up/down feedback on runs,
|
||||
optionally scoped to a specific message.
|
||||
Thumbs-up/down on a run, aligned with the mainstream chat UX: PUT is an
|
||||
idempotent upsert ("my current verdict is X"), DELETE retracts it
|
||||
(clicking the active button again). The current rating is echoed back to
|
||||
the UI inside the message-list payload, not through a separate read
|
||||
endpoint.
|
||||
|
||||
Primary adapter only: resolves the current user, delegates to
|
||||
``FeedbackService`` (input port), and maps domain errors to HTTP codes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.gateway.authz import require_permission
|
||||
from app.gateway.deps import get_current_user, get_feedback_repo, get_run_store
|
||||
from app.gateway.deps import get_current_user, get_feedback_service
|
||||
from deerflow.domain.feedback import (
|
||||
DuplicateFeedbackError,
|
||||
Feedback,
|
||||
InvalidRatingError,
|
||||
InvalidTagError,
|
||||
RunNotFoundError,
|
||||
)
|
||||
from deerflow.utils.time import coerce_iso
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/api/threads", tags=["feedback"])
|
||||
@ -24,15 +37,13 @@ router = APIRouter(prefix="/api/threads", tags=["feedback"])
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FeedbackCreateRequest(BaseModel):
|
||||
rating: int = Field(..., description="Feedback rating: +1 (positive) or -1 (negative)")
|
||||
comment: str | None = Field(default=None, description="Optional text feedback")
|
||||
message_id: str | None = Field(default=None, description="Optional: scope feedback to a specific message")
|
||||
|
||||
|
||||
class FeedbackUpsertRequest(BaseModel):
|
||||
rating: int = Field(..., description="Feedback rating: +1 (positive) or -1 (negative)")
|
||||
comment: str | None = Field(default=None, description="Optional text feedback")
|
||||
tags: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Optional thumbs-down reason slugs (e.g. 'incorrect', 'slow')",
|
||||
)
|
||||
|
||||
|
||||
class FeedbackResponse(BaseModel):
|
||||
@ -43,14 +54,24 @@ class FeedbackResponse(BaseModel):
|
||||
message_id: str | None = None
|
||||
rating: int
|
||||
comment: str | None = None
|
||||
tags: list[str] = []
|
||||
created_at: str = ""
|
||||
|
||||
|
||||
class FeedbackStatsResponse(BaseModel):
|
||||
run_id: str
|
||||
total: int = 0
|
||||
positive: int = 0
|
||||
negative: int = 0
|
||||
def _to_response(feedback: Feedback) -> FeedbackResponse:
|
||||
"""Domain object -> wire shape. ``created_at`` keeps the legacy
|
||||
``coerce_iso`` serialization so the API output is byte-identical."""
|
||||
return FeedbackResponse(
|
||||
feedback_id=feedback.feedback_id,
|
||||
run_id=feedback.run_id,
|
||||
thread_id=feedback.thread_id,
|
||||
user_id=feedback.user_id,
|
||||
message_id=feedback.message_id,
|
||||
rating=feedback.rating,
|
||||
comment=feedback.comment,
|
||||
tags=list(feedback.tags),
|
||||
created_at=coerce_iso(feedback.created_at),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@ -65,28 +86,30 @@ async def upsert_feedback(
|
||||
run_id: str,
|
||||
body: FeedbackUpsertRequest,
|
||||
request: Request,
|
||||
) -> dict[str, Any]:
|
||||
"""Create or update feedback for a run (idempotent)."""
|
||||
if body.rating not in (1, -1):
|
||||
raise HTTPException(status_code=400, detail="rating must be +1 or -1")
|
||||
|
||||
) -> FeedbackResponse:
|
||||
"""Set the current user's rating for a run (idempotent upsert)."""
|
||||
user_id = await get_current_user(request)
|
||||
|
||||
run_store = get_run_store(request)
|
||||
run = await run_store.get(run_id)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
||||
if run.get("thread_id") != thread_id:
|
||||
service = get_feedback_service(request)
|
||||
try:
|
||||
feedback = await service.rate_run(
|
||||
thread_id,
|
||||
run_id,
|
||||
rating=body.rating,
|
||||
comment=body.comment,
|
||||
user_id=user_id,
|
||||
tags=body.tags,
|
||||
)
|
||||
except InvalidRatingError:
|
||||
raise HTTPException(status_code=400, detail="rating must be +1 or -1")
|
||||
except InvalidTagError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc))
|
||||
except RunNotFoundError:
|
||||
raise HTTPException(status_code=404, detail=f"Run {run_id} not found in thread {thread_id}")
|
||||
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
return await feedback_repo.upsert(
|
||||
run_id=run_id,
|
||||
thread_id=thread_id,
|
||||
rating=body.rating,
|
||||
user_id=user_id,
|
||||
comment=body.comment,
|
||||
)
|
||||
except DuplicateFeedbackError:
|
||||
# Lost a concurrent-upsert race (legacy behavior was a 500); the
|
||||
# client can simply retry.
|
||||
raise HTTPException(status_code=409, detail="Concurrent feedback update, please retry")
|
||||
return _to_response(feedback)
|
||||
|
||||
|
||||
@router.delete("/{thread_id}/runs/{run_id}/feedback")
|
||||
@ -96,93 +119,10 @@ async def delete_run_feedback(
|
||||
run_id: str,
|
||||
request: Request,
|
||||
) -> dict[str, bool]:
|
||||
"""Delete the current user's feedback for a run."""
|
||||
"""Retract the current user's rating for a run."""
|
||||
user_id = await get_current_user(request)
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
deleted = await feedback_repo.delete_by_run(
|
||||
thread_id=thread_id,
|
||||
run_id=run_id,
|
||||
user_id=user_id,
|
||||
)
|
||||
if not deleted:
|
||||
service = get_feedback_service(request)
|
||||
retracted = await service.retract_run_rating(thread_id, run_id, user_id=user_id)
|
||||
if not retracted:
|
||||
raise HTTPException(status_code=404, detail="No feedback found for this run")
|
||||
return {"success": True}
|
||||
|
||||
|
||||
@router.post("/{thread_id}/runs/{run_id}/feedback", response_model=FeedbackResponse)
|
||||
@require_permission("threads", "write", owner_check=True, require_existing=True)
|
||||
async def create_feedback(
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
body: FeedbackCreateRequest,
|
||||
request: Request,
|
||||
) -> dict[str, Any]:
|
||||
"""Submit feedback (thumbs-up/down) for a run."""
|
||||
if body.rating not in (1, -1):
|
||||
raise HTTPException(status_code=400, detail="rating must be +1 or -1")
|
||||
|
||||
user_id = await get_current_user(request)
|
||||
|
||||
# Validate run exists and belongs to thread
|
||||
run_store = get_run_store(request)
|
||||
run = await run_store.get(run_id)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
||||
if run.get("thread_id") != thread_id:
|
||||
raise HTTPException(status_code=404, detail=f"Run {run_id} not found in thread {thread_id}")
|
||||
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
return await feedback_repo.create(
|
||||
run_id=run_id,
|
||||
thread_id=thread_id,
|
||||
rating=body.rating,
|
||||
user_id=user_id,
|
||||
message_id=body.message_id,
|
||||
comment=body.comment,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{thread_id}/runs/{run_id}/feedback", response_model=list[FeedbackResponse])
|
||||
@require_permission("threads", "read", owner_check=True)
|
||||
async def list_feedback(
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
request: Request,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all feedback for a run."""
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
return await feedback_repo.list_by_run(thread_id, run_id)
|
||||
|
||||
|
||||
@router.get("/{thread_id}/runs/{run_id}/feedback/stats", response_model=FeedbackStatsResponse)
|
||||
@require_permission("threads", "read", owner_check=True)
|
||||
async def feedback_stats(
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
request: Request,
|
||||
) -> dict[str, Any]:
|
||||
"""Get aggregated feedback stats (positive/negative counts) for a run."""
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
return await feedback_repo.aggregate_by_run(thread_id, run_id)
|
||||
|
||||
|
||||
@router.delete("/{thread_id}/runs/{run_id}/feedback/{feedback_id}")
|
||||
@require_permission("threads", "delete", owner_check=True, require_existing=True)
|
||||
async def delete_feedback(
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
feedback_id: str,
|
||||
request: Request,
|
||||
) -> dict[str, bool]:
|
||||
"""Delete a feedback record."""
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
# Verify feedback belongs to the specified thread/run before deleting
|
||||
existing = await feedback_repo.get(feedback_id)
|
||||
if existing is None:
|
||||
raise HTTPException(status_code=404, detail=f"Feedback {feedback_id} not found")
|
||||
if existing.get("thread_id") != thread_id or existing.get("run_id") != run_id:
|
||||
raise HTTPException(status_code=404, detail=f"Feedback {feedback_id} not found in run {run_id}")
|
||||
deleted = await feedback_repo.delete(feedback_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=404, detail=f"Feedback {feedback_id} not found")
|
||||
return {"success": True}
|
||||
|
||||
@ -14,7 +14,7 @@ from fastapi import APIRouter, HTTPException, Query, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.gateway.authz import require_permission
|
||||
from app.gateway.deps import get_feedback_repo, get_run_event_store, get_run_manager, get_run_store, get_stream_bridge
|
||||
from app.gateway.deps import get_run_event_store, get_run_manager, get_run_store, get_stream_bridge
|
||||
from app.gateway.pagination import trim_run_message_page
|
||||
from app.gateway.routers.thread_runs import RunCreateRequest
|
||||
from app.gateway.services import build_checkpoint_state_accessor, sse_consumer, start_run, wait_for_run_completion
|
||||
@ -134,12 +134,3 @@ async def run_messages(
|
||||
)
|
||||
data, has_more = trim_run_message_page(rows, limit=limit, after_seq=after_seq)
|
||||
return {"data": data, "has_more": has_more}
|
||||
|
||||
|
||||
@router.get("/{run_id}/feedback")
|
||||
@require_permission("runs", "read")
|
||||
async def run_feedback(run_id: str, request: Request) -> list[dict]:
|
||||
"""Return all feedback for a run."""
|
||||
run = await _resolve_run(run_id, request)
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
return await feedback_repo.list_by_run(run["thread_id"], run_id)
|
||||
|
||||
@ -23,7 +23,7 @@ from langchain_core.messages import BaseMessage
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.gateway.authz import require_permission
|
||||
from app.gateway.deps import get_current_user, get_feedback_repo, get_run_event_store, get_run_manager, get_run_store, get_stream_bridge
|
||||
from app.gateway.deps import get_current_user, get_feedback_service, get_run_event_store, get_run_manager, get_run_store, get_stream_bridge
|
||||
from app.gateway.pagination import trim_run_message_page
|
||||
from app.gateway.services import build_checkpoint_state_accessor, build_thread_checkpoint_state_accessor, sse_consumer, start_run, wait_for_run_completion
|
||||
from deerflow.runtime import CancelOutcome, RunRecord, RunStatus, serialize_channel_values_for_api
|
||||
@ -731,10 +731,10 @@ async def list_thread_messages(
|
||||
# Attach feedback to the last AI message of each run. Only query when there
|
||||
# is an AI message to attach it to — threads with no completed AI turn yet
|
||||
# would otherwise pay for a grouped feedback lookup whose result is unused.
|
||||
feedback_map: dict[str, dict] = {}
|
||||
feedback_map = {}
|
||||
if last_ai_per_run:
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
feedback_map = await feedback_repo.list_by_thread_grouped(thread_id, user_id=user_id)
|
||||
feedback_service = get_feedback_service(request)
|
||||
feedback_map = await feedback_service.latest_per_run_in_thread(thread_id, user_id=user_id)
|
||||
|
||||
last_ai_indices = set(last_ai_per_run.values())
|
||||
for i, msg in enumerate(messages):
|
||||
@ -743,9 +743,10 @@ async def list_thread_messages(
|
||||
fb = feedback_map.get(run_id)
|
||||
msg["feedback"] = (
|
||||
{
|
||||
"feedback_id": fb["feedback_id"],
|
||||
"rating": fb["rating"],
|
||||
"comment": fb.get("comment"),
|
||||
"feedback_id": fb.feedback_id,
|
||||
"rating": fb.rating,
|
||||
"comment": fb.comment,
|
||||
"tags": list(fb.tags),
|
||||
}
|
||||
if fb
|
||||
else None
|
||||
@ -851,11 +852,11 @@ async def _enrich_thread_message_page(
|
||||
|
||||
event_store = get_run_event_store(request)
|
||||
last_ai_seq_by_run = await event_store.get_last_visible_ai_seq_by_run(thread_id, run_ids, user_id=user_id)
|
||||
feedback_map: dict[str, dict] = {}
|
||||
feedback_map = {}
|
||||
feedback_run_ids = {run_id for row in data if isinstance((run_id := row.get("run_id")), str) and row.get("seq") == last_ai_seq_by_run.get(run_id)}
|
||||
if feedback_run_ids:
|
||||
feedback_repo = get_feedback_repo(request)
|
||||
feedback_map = await feedback_repo.list_by_run_ids(thread_id, feedback_run_ids, user_id=user_id)
|
||||
feedback_service = get_feedback_service(request)
|
||||
feedback_map = await feedback_service.latest_for_runs(thread_id, feedback_run_ids, user_id=user_id)
|
||||
|
||||
for row in data:
|
||||
run_id = row.get("run_id")
|
||||
@ -864,9 +865,10 @@ async def _enrich_thread_message_page(
|
||||
feedback = feedback_map.get(run_id)
|
||||
if feedback:
|
||||
row["feedback"] = {
|
||||
"feedback_id": feedback["feedback_id"],
|
||||
"rating": feedback["rating"],
|
||||
"comment": feedback.get("comment"),
|
||||
"feedback_id": feedback.feedback_id,
|
||||
"rating": feedback.rating,
|
||||
"comment": feedback.comment,
|
||||
"tags": list(feedback.tags),
|
||||
}
|
||||
|
||||
content = row.get("content")
|
||||
|
||||
@ -245,67 +245,6 @@ def _make_feedback_repo(rows: list[dict]):
|
||||
return repo
|
||||
|
||||
|
||||
def _make_feedback(run_id: str, idx: int) -> dict:
|
||||
return {"id": f"fb-{idx}", "run_id": run_id, "thread_id": "thread-x", "value": "up"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TestRunFeedback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRunFeedback:
|
||||
def test_returns_list_of_feedback_dicts(self):
|
||||
"""GET /api/runs/{run_id}/feedback returns a list of feedback dicts."""
|
||||
run_record = {"run_id": "run-fb-1", "thread_id": "thread-fb-1"}
|
||||
rows = [_make_feedback("run-fb-1", i) for i in range(3)]
|
||||
app = _make_app(
|
||||
run_store=_make_run_store(run_record),
|
||||
feedback_repo=_make_feedback_repo(rows),
|
||||
)
|
||||
with TestClient(app) as client:
|
||||
response = client.get("/api/runs/run-fb-1/feedback")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert isinstance(body, list)
|
||||
assert len(body) == 3
|
||||
|
||||
def test_404_when_run_not_found(self):
|
||||
"""Returns 404 when run store returns None."""
|
||||
app = _make_app(
|
||||
run_store=_make_run_store(None),
|
||||
feedback_repo=_make_feedback_repo([]),
|
||||
)
|
||||
with TestClient(app) as client:
|
||||
response = client.get("/api/runs/missing-run/feedback")
|
||||
assert response.status_code == 404
|
||||
assert "missing-run" in response.json()["detail"]
|
||||
|
||||
def test_empty_list_when_no_feedback(self):
|
||||
"""Returns empty list when no feedback exists for the run."""
|
||||
run_record = {"run_id": "run-fb-2", "thread_id": "thread-fb-2"}
|
||||
app = _make_app(
|
||||
run_store=_make_run_store(run_record),
|
||||
feedback_repo=_make_feedback_repo([]),
|
||||
)
|
||||
with TestClient(app) as client:
|
||||
response = client.get("/api/runs/run-fb-2/feedback")
|
||||
assert response.status_code == 200
|
||||
assert response.json() == []
|
||||
|
||||
def test_503_when_feedback_repo_not_configured(self):
|
||||
"""Returns 503 when feedback_repo is None (no DB configured)."""
|
||||
run_record = {"run_id": "run-fb-3", "thread_id": "thread-fb-3"}
|
||||
app = _make_app(
|
||||
run_store=_make_run_store(run_record),
|
||||
)
|
||||
# Explicitly set feedback_repo to None to simulate missing DB
|
||||
app.state.feedback_repo = None
|
||||
with TestClient(app) as client:
|
||||
response = client.get("/api/runs/run-fb-3/feedback")
|
||||
assert response.status_code == 503
|
||||
|
||||
|
||||
def test_resolve_thread_id_handles_null_configurable():
|
||||
"""A client may send ``config.configurable`` as JSON ``null``.
|
||||
|
||||
|
||||
@ -15,6 +15,7 @@ from _router_auth_helpers import make_authed_test_app
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.gateway.routers import thread_runs
|
||||
from deerflow.domain.feedback import Feedback
|
||||
|
||||
|
||||
def _make_app(messages, feedback_grouped):
|
||||
@ -25,9 +26,9 @@ def _make_app(messages, feedback_grouped):
|
||||
event_store.list_messages = AsyncMock(return_value=messages)
|
||||
app.state.run_event_store = event_store
|
||||
|
||||
feedback_repo = MagicMock()
|
||||
feedback_repo.list_by_thread_grouped = AsyncMock(return_value=feedback_grouped)
|
||||
app.state.feedback_repo = feedback_repo
|
||||
feedback_service = MagicMock()
|
||||
feedback_service.latest_per_run_in_thread = AsyncMock(return_value=feedback_grouped)
|
||||
app.state.feedback_service = feedback_service
|
||||
|
||||
# list_thread_messages also calls run_manager.list_by_thread to inject
|
||||
# turn durations; stub it to return no runs so that path stays inert.
|
||||
@ -35,7 +36,7 @@ def _make_app(messages, feedback_grouped):
|
||||
run_manager.list_by_thread = AsyncMock(return_value=[])
|
||||
app.state.run_manager = run_manager
|
||||
|
||||
return app, feedback_repo
|
||||
return app, feedback_service
|
||||
|
||||
|
||||
def _ai(run_id: str, seq: int, content: str) -> dict:
|
||||
@ -53,8 +54,8 @@ def test_feedback_attached_to_last_ai_message_per_run():
|
||||
_ai("r1", 3, "final answer"), # last AI of r1 -> should get feedback
|
||||
_ai("r2", 4, "other run"), # last AI of r2 -> no feedback row
|
||||
]
|
||||
grouped = {"r1": {"feedback_id": "fb-1", "rating": "up", "comment": "nice"}}
|
||||
app, feedback_repo = _make_app(messages, grouped)
|
||||
grouped = {"r1": Feedback(feedback_id="fb-1", run_id="r1", thread_id="t1", rating=1, comment="nice")}
|
||||
app, feedback_service = _make_app(messages, grouped)
|
||||
|
||||
resp = TestClient(app).get("/api/threads/t1/messages")
|
||||
assert resp.status_code == 200
|
||||
@ -62,21 +63,21 @@ def test_feedback_attached_to_last_ai_message_per_run():
|
||||
|
||||
by_seq = {m["seq"]: m for m in data}
|
||||
# The bug: this used to be None for every message.
|
||||
assert by_seq[3]["feedback"] == {"feedback_id": "fb-1", "rating": "up", "comment": "nice"}
|
||||
assert by_seq[3]["feedback"] == {"feedback_id": "fb-1", "rating": 1, "comment": "nice", "tags": []}
|
||||
# Earlier AI message of the same run and the human message get no feedback.
|
||||
assert by_seq[2]["feedback"] is None
|
||||
assert by_seq[1]["feedback"] is None
|
||||
# r2's last AI message has no feedback row.
|
||||
assert by_seq[4]["feedback"] is None
|
||||
feedback_repo.list_by_thread_grouped.assert_awaited_once()
|
||||
feedback_service.latest_per_run_in_thread.assert_awaited_once()
|
||||
|
||||
|
||||
def test_no_feedback_query_when_thread_has_no_ai_message():
|
||||
messages = [_human("r1", 1)]
|
||||
app, feedback_repo = _make_app(messages, {})
|
||||
app, feedback_service = _make_app(messages, {})
|
||||
|
||||
resp = TestClient(app).get("/api/threads/t1/messages")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()[0]["feedback"] is None
|
||||
# No AI message -> the grouped feedback query must not run.
|
||||
feedback_repo.list_by_thread_grouped.assert_not_awaited()
|
||||
feedback_service.latest_per_run_in_thread.assert_not_awaited()
|
||||
|
||||
@ -11,6 +11,7 @@ from _router_auth_helpers import make_authed_test_app
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.gateway.routers import thread_runs
|
||||
from deerflow.domain.feedback import Feedback
|
||||
from deerflow.runtime import RunRecord
|
||||
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||
|
||||
@ -23,9 +24,9 @@ def _make_app(event_store: MemoryRunEventStore, *, superseded: set[str] | None =
|
||||
run_manager.list_successful_regenerate_sources.return_value = superseded or set()
|
||||
run_manager.get_many_by_thread.return_value = records or {}
|
||||
app.state.run_manager = run_manager
|
||||
feedback_repo = AsyncMock()
|
||||
feedback_repo.list_by_run_ids.return_value = feedback or {}
|
||||
app.state.feedback_repo = feedback_repo
|
||||
feedback_service = AsyncMock()
|
||||
feedback_service.latest_for_runs.return_value = feedback or {}
|
||||
app.state.feedback_service = feedback_service
|
||||
return app
|
||||
|
||||
|
||||
@ -172,7 +173,7 @@ def test_thread_page_feedback_only_attaches_to_global_last_ai_row():
|
||||
original_get_last_visible_ai_seq_by_run = store.get_last_visible_ai_seq_by_run
|
||||
store.list_messages = AsyncMock(wraps=original_list_messages)
|
||||
store.get_last_visible_ai_seq_by_run = AsyncMock(wraps=original_get_last_visible_ai_seq_by_run)
|
||||
feedback = {"run-1": {"feedback_id": "fb-1", "rating": 1, "comment": "good"}}
|
||||
feedback = {"run-1": Feedback(feedback_id="fb-1", run_id="run-1", thread_id="thread-1", rating=1, comment="good")}
|
||||
app = _make_app(store, feedback=feedback)
|
||||
with TestClient(app) as client:
|
||||
response = client.get("/api/threads/thread-1/messages/page?limit=3")
|
||||
@ -180,13 +181,13 @@ def test_thread_page_feedback_only_attaches_to_global_last_ai_row():
|
||||
data = response.json()["data"]
|
||||
assert data[0]["feedback"] is None
|
||||
assert data[1]["feedback"] is None
|
||||
assert data[2]["feedback"] == {"feedback_id": "fb-1", "rating": 1, "comment": "good"}
|
||||
assert data[2]["feedback"] == {"feedback_id": "fb-1", "rating": 1, "comment": "good", "tags": []}
|
||||
scan_user_id = store.list_messages.await_args.kwargs["user_id"]
|
||||
enrichment_user_id = store.get_last_visible_ai_seq_by_run.await_args.kwargs["user_id"]
|
||||
assert enrichment_user_id == scan_user_id
|
||||
feedback_repo = app.state.feedback_repo
|
||||
feedback_repo.list_by_run_ids.assert_awaited_once_with("thread-1", {"run-1"}, user_id=scan_user_id)
|
||||
feedback_repo.list_by_thread_grouped.assert_not_awaited()
|
||||
feedback_service = app.state.feedback_service
|
||||
feedback_service.latest_for_runs.assert_awaited_once_with("thread-1", {"run-1"}, user_id=scan_user_id)
|
||||
feedback_service.latest_per_run_in_thread.assert_not_awaited()
|
||||
|
||||
|
||||
def test_thread_page_helpers_forward_explicit_user_without_request_context():
|
||||
@ -199,7 +200,7 @@ def test_thread_page_helpers_forward_explicit_user_without_request_context():
|
||||
request = MagicMock()
|
||||
request.app.state.run_event_store = event_store
|
||||
request.app.state.run_manager = run_manager
|
||||
request.app.state.feedback_repo = AsyncMock()
|
||||
request.app.state.feedback_service = AsyncMock()
|
||||
|
||||
async def exercise_helpers():
|
||||
await thread_runs._scan_thread_message_page(
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user