mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-16 17:46:20 +00:00
The ORM row stays behind (single alembic chain until the infrastructure move); everything else is served by the domain slice. Repository tests become a contract suite that runs the same cases against the SQL adapter and an in-memory fake; ownership and timezone tests use explicit user_id instead of the AUTO-sentinel context.
97 lines
3.7 KiB
Python
97 lines
3.7 KiB
Python
"""FeedbackService tests with injected fakes — zero IO, no engine.
|
|
|
|
The seam the ports create is what makes these tests possible: the service
|
|
is exercised end-to-end against dict-backed doubles.
|
|
"""
|
|
|
|
import pytest
|
|
from test_feedback import InMemoryFeedbackRepository
|
|
|
|
from deerflow.domain.feedback import FeedbackService, InvalidRatingError, RunNotFoundError
|
|
|
|
|
|
class FakeRunLookup:
|
|
"""RunLookup double backed by a run_id -> thread_id mapping."""
|
|
|
|
def __init__(self, runs: dict[str, str]):
|
|
self._runs = runs
|
|
|
|
async def thread_of(self, run_id: str) -> str | None:
|
|
return self._runs.get(run_id)
|
|
|
|
|
|
def _service(runs: dict[str, str] | None = None) -> FeedbackService:
|
|
return FeedbackService(
|
|
repository=InMemoryFeedbackRepository(),
|
|
runs=FakeRunLookup(runs if runs is not None else {"r1": "t1"}),
|
|
)
|
|
|
|
|
|
class TestRateRun:
|
|
@pytest.mark.anyio
|
|
async def test_stores_rating(self):
|
|
svc = _service()
|
|
fb = await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
assert fb.rating == 1
|
|
assert fb.user_id == "u1"
|
|
|
|
@pytest.mark.anyio
|
|
async def test_idempotent_upsert_keeps_identity(self):
|
|
svc = _service()
|
|
first = await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
second = await svc.rate_run("t1", "r1", rating=-1, comment="meh", user_id="u1")
|
|
assert second.feedback_id == first.feedback_id
|
|
assert second.rating == -1
|
|
|
|
@pytest.mark.anyio
|
|
async def test_enrich_with_tags(self):
|
|
svc = _service()
|
|
await svc.rate_run("t1", "r1", rating=-1, comment=None, user_id="u1")
|
|
fb = await svc.rate_run("t1", "r1", rating=-1, comment="wrong", user_id="u1", tags=["incorrect"])
|
|
assert fb.tags == ("incorrect",)
|
|
|
|
@pytest.mark.anyio
|
|
async def test_unknown_run_rejected(self):
|
|
svc = _service({})
|
|
with pytest.raises(RunNotFoundError):
|
|
await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
|
|
@pytest.mark.anyio
|
|
async def test_cross_thread_run_rejected(self):
|
|
"""A run id belonging to another thread must not be ratable through
|
|
this thread's URL (ownership check)."""
|
|
svc = _service({"r1": "other-thread"})
|
|
with pytest.raises(RunNotFoundError):
|
|
await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
|
|
@pytest.mark.anyio
|
|
async def test_invalid_rating_propagates_before_persistence(self):
|
|
svc = _service()
|
|
with pytest.raises(InvalidRatingError):
|
|
await svc.rate_run("t1", "r1", rating=0, comment=None, user_id="u1")
|
|
|
|
|
|
class TestRetractAndReads:
|
|
@pytest.mark.anyio
|
|
async def test_retract(self):
|
|
svc = _service()
|
|
await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
assert await svc.retract_run_rating("t1", "r1", user_id="u1") is True
|
|
assert await svc.retract_run_rating("t1", "r1", user_id="u1") is False
|
|
|
|
@pytest.mark.anyio
|
|
async def test_latest_per_run_in_thread(self):
|
|
svc = _service({"r1": "t1", "r2": "t1"})
|
|
await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
await svc.rate_run("t1", "r2", rating=-1, comment=None, user_id="u1")
|
|
grouped = await svc.latest_per_run_in_thread("t1", user_id="u1")
|
|
assert {run_id: fb.rating for run_id, fb in grouped.items()} == {"r1": 1, "r2": -1}
|
|
|
|
@pytest.mark.anyio
|
|
async def test_latest_for_runs(self):
|
|
svc = _service({"r1": "t1", "r2": "t1"})
|
|
await svc.rate_run("t1", "r1", rating=1, comment=None, user_id="u1")
|
|
await svc.rate_run("t1", "r2", rating=-1, comment=None, user_id="u1")
|
|
grouped = await svc.latest_for_runs("t1", {"r2"}, user_id="u1")
|
|
assert set(grouped) == {"r2"}
|