100 lines
3.5 KiB
Python
100 lines
3.5 KiB
Python
"""Failure-ordering regression coverage for final Deep Research report writes."""
|
|
|
|
import pytest
|
|
|
|
from app.gateway import deep_research_job_executor as executor_module
|
|
from app.gateway.deep_research_job_executor import DeepResearchJobExecutor
|
|
from deerflow.agents.deep_research.events import DeepResearchEvent
|
|
from deerflow.agents.deep_research.types import DeepResearchResult
|
|
|
|
|
|
class _JobStore:
|
|
def __init__(self, session_store) -> None:
|
|
self._session_store = session_store
|
|
self.progress_calls: list[dict] = []
|
|
|
|
async def update_progress(self, job_id: str, **fields): # noqa: ARG002
|
|
assert self._session_store.report_updates, "report must be persisted before job completion"
|
|
self.progress_calls.append(fields)
|
|
return {"id": job_id, **fields}
|
|
|
|
|
|
class _SessionStore:
|
|
def __init__(self, *, available: bool = True) -> None:
|
|
self.available = available
|
|
self.report_updates: list[dict] = []
|
|
|
|
async def update(self, session_id: str, **fields): # noqa: ARG002
|
|
if not self.available:
|
|
raise RuntimeError("temporary session database failure")
|
|
self.report_updates.append(fields)
|
|
return {"id": session_id, **fields}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_report_is_persisted_before_job_is_marked_completed():
|
|
sessions = _SessionStore()
|
|
jobs = _JobStore(sessions)
|
|
runner = DeepResearchJobExecutor(jobs, sessions, event_store=None)
|
|
|
|
await runner._mark_completed(
|
|
"drj_1",
|
|
"drs_1",
|
|
"worker_1",
|
|
DeepResearchResult(report_markdown="# 报告", source_ids=["src_1"]),
|
|
user_id="user_1",
|
|
query="课题",
|
|
)
|
|
|
|
assert sessions.report_updates[0]["report_markdown"] == "# 报告"
|
|
assert jobs.progress_calls[0]["status"] == "completed"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_report_projection_failure_never_marks_job_completed(monkeypatch):
|
|
monkeypatch.setattr(executor_module, "_REPORT_PROJECTION_RETRY_ATTEMPTS", 2)
|
|
monkeypatch.setattr(executor_module, "_REPORT_PROJECTION_RETRY_DELAY_SECONDS", 0)
|
|
sessions = _SessionStore(available=False)
|
|
jobs = _JobStore(sessions)
|
|
runner = DeepResearchJobExecutor(jobs, sessions, event_store=None)
|
|
|
|
with pytest.raises(RuntimeError, match="研究报告未能保存"):
|
|
await runner._mark_completed(
|
|
"drj_1",
|
|
"drs_1",
|
|
"worker_1",
|
|
DeepResearchResult(report_markdown="# 报告"),
|
|
user_id="user_1",
|
|
query="课题",
|
|
)
|
|
|
|
assert jobs.progress_calls == []
|
|
|
|
|
|
def test_terminal_timeline_events_are_valid_durable_event_types():
|
|
assert DeepResearchEvent(type="job_completed", phase="done").type == "job_completed"
|
|
assert DeepResearchEvent(type="job_failed", phase="done").type == "job_failed"
|
|
|
|
|
|
def test_zero_source_fallback_is_a_valid_completed_report():
|
|
from app.gateway.deep_research_job_executor import _terminal_report_validation_error
|
|
from deerflow.agents.deep_research.runners.basic import _build_report_fallback
|
|
|
|
report = _build_report_fallback("没有命中的课题", "", 0)
|
|
assert _terminal_report_validation_error(report) is None
|
|
assert "没有命中的课题" in report
|
|
|
|
|
|
def test_source_table_reload_does_not_become_session_error():
|
|
from app.gateway.deep_research_job_executor import _session_diagnostic_error
|
|
|
|
assert _session_diagnostic_error(
|
|
[
|
|
{
|
|
"code": "client_materials_reloaded_before_write",
|
|
"stage": "collecting",
|
|
"detail": "写作前再次从素材表读到 15 条",
|
|
}
|
|
]
|
|
) is None
|