deerflow-code/offline-backend-20260512/backend/tests/test_deep_research_start_job.py
2026-09-07 18:24:55 +08:00

1083 lines
39 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Regression coverage for idempotent Deep Research job submission."""
import json
import os
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
import pytest
from app.gateway import deep_research_report_rewrite as report_rewrite
from app.gateway.routers import deep_research
class _SessionStore:
def __init__(self) -> None:
self.updates: list[dict] = []
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return {
"id": session_id,
"query": "测试主题",
"title": "测试主题",
"config_snapshot": {"mode": "basic"},
}
async def count_active_by_user(self, *, user_id: str) -> int: # noqa: ARG002
return 0
async def list_by_user(self, *, user_id: str, status: str | None = None, limit: int = 20, **_: object): # noqa: ARG002
return []
async def update(self, session_id: str, *, user_id: str, **fields): # noqa: ARG002
self.updates.append(fields)
return {"id": session_id, **fields}
class _JobStore:
async def get_active_for_session(self, session_id: str, *, user_id: str): # noqa: ARG002
return None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
# Simulate a retry after an earlier request committed the job but did
# not finish responding to the browser.
return {"id": "drj_existing", "status": "queued"}, False
class _Dispatcher:
def __init__(self) -> None:
self.nudges = 0
def nudge(self) -> None:
self.nudges += 1
def cancel_running(self, job_id: str) -> bool: # noqa: ARG002
return False
def _request(*, sessions, jobs, dispatcher, sources=None):
return SimpleNamespace(
app=SimpleNamespace(
state=SimpleNamespace(
deep_research_session_store=sessions,
deep_research_job_store=jobs,
deep_research_dispatcher=dispatcher,
deep_research_source_store=sources,
)
)
)
@pytest.mark.asyncio
async def test_retry_of_existing_queued_job_nudges_dispatcher(monkeypatch):
"""An idempotent retry must recover a queued-but-not-yet-dispatched job."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sessions = _SessionStore()
dispatcher = _Dispatcher()
request = _request(sessions=sessions, jobs=_JobStore(), dispatcher=dispatcher)
response = await deep_research.start_job(
"drs_1",
deep_research.StartJobRequest(entry="legacy"),
request,
)
assert response == {
"jobId": "drj_existing",
"sessionId": "drs_1",
"status": "queued",
"reused": True,
"collectorMaterialCount": 0,
"persistedSourceCount": 0,
"diagnostics": [],
}
assert dispatcher.nudges == 1
# Restarting clears any diagnostic left by the previous attempt, so the
# frontend does not append a stale error to a healthy new run.
assert sessions.updates == [
{"status": "running", "active_job_id": "drj_existing", "error": None}
]
@pytest.mark.asyncio
async def test_queued_job_still_starts_when_session_projection_temporarily_fails(monkeypatch):
"""A committed job is more important than its non-authoritative UI projection."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _ProjectionFailureSessionStore(_SessionStore):
async def update(self, session_id: str, *, user_id: str, **fields): # noqa: ARG002
raise RuntimeError("temporary database readback failure")
class _NewJobStore(_JobStore):
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
return {"id": "drj_new", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
dispatcher = _Dispatcher()
response = await deep_research.start_job(
"drs_1",
deep_research.StartJobRequest(entry="legacy"),
_request(
sessions=_ProjectionFailureSessionStore(),
jobs=_NewJobStore(),
dispatcher=dispatcher,
),
)
assert response["jobId"] == "drj_new"
assert response["reused"] is False
assert dispatcher.nudges == 1
assert any(item["code"] == "session_projection_failed" for item in response["diagnostics"])
@pytest.mark.asyncio
async def test_chat_report_card_can_enable_images_after_collection(monkeypatch):
"""The write-time card override must win over the collection-time default."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _ChatSessionStore(_SessionStore):
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return {
"id": session_id,
"query": "测试主题",
"title": "测试主题",
"runtime_thread_id": "collector-thread",
"config_snapshot": {
"mode": "basic",
"collection_mode": "chat",
"generate_images": False,
"max_generated_images": 2,
},
}
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_images", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sessions = _ChatSessionStore()
jobs = _CapturingJobStore()
response = await deep_research.start_job(
"drs_1",
deep_research.StartJobRequest(
entry="chat",
config={"generate_images": True, "max_generated_images": 3},
),
_request(sessions=sessions, jobs=jobs, dispatcher=_Dispatcher()),
)
assert response["jobId"] == "drj_images"
assert jobs.created is not None
snapshot = jobs.created["input_snapshot"]["config"]
assert snapshot["generate_images"] is True
assert snapshot["max_generated_images"] == 3
# The session snapshot is persisted too, so a reload/replay keeps the
# write-time choice instead of reverting to the creation-time false.
assert any(
update.get("config_snapshot", {}).get("generate_images") is True
for update in sessions.updates
)
@pytest.mark.asyncio
async def test_report_variant_uses_normal_job_with_selected_sources(monkeypatch):
"""Regeneration must queue the original report engine, never a rewrite job."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _VariantSessionStore(_SessionStore):
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return {
"id": session_id,
"query": "新能源汽车市场",
"title": "新能源汽车市场 · 新结构报告",
"status": "draft",
"runtime_thread_id": "dr_variant_thread",
"config_snapshot": {
"mode": "basic",
"collection_mode": "chat",
"report_variant": True,
"parent_session_id": "drs_parent",
},
}
class _SelectedSources:
async def count_by_session(self, *args, **kwargs): # noqa: ANN002, ANN003
return 12
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_regenerate", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sessions = _VariantSessionStore()
jobs = _CapturingJobStore()
response = await deep_research.start_job(
"drs_variant",
deep_research.StartJobRequest(
entry="regenerate",
config={
"custom_outline": "# 新结构\n\n## 第四章",
"report_instruction": "统一表达并保留可核验事实。",
},
),
_request(
sessions=sessions,
jobs=jobs,
dispatcher=_Dispatcher(),
sources=_SelectedSources(),
),
)
assert response["jobId"] == "drj_regenerate"
snapshot = jobs.created["input_snapshot"]
assert snapshot["entry"] == "regenerate"
assert snapshot["config"]["custom_outline"].endswith("## 第四章")
assert snapshot["config"]["report_instruction"] == "统一表达并保留可核验事实。"
assert "report_variant" not in snapshot["config"]
assert "parent_session_id" not in snapshot["config"]
assert any(
update.get("config_snapshot", {}).get("report_variant") is True
for update in sessions.updates
)
assert any(
update.get("config_snapshot", {}).get("parent_session_id") == "drs_parent"
for update in sessions.updates
)
@pytest.mark.asyncio
async def test_stream_stays_open_when_durable_replay_has_a_transient_failure(monkeypatch):
"""A stream DB blip must yield a recoverable warning, not an ASGI error group."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _StreamJobStore:
async def get(self, job_id: str, *, user_id: str): # noqa: ARG002
return {"id": job_id, "session_id": "drs_1"}
async def get_unscoped(self, job_id: str): # noqa: ARG002
return None
class _UnavailableEventStore:
async def list_after(self, *args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("temporary replica failure")
async def is_connected() -> bool:
return False
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
request = SimpleNamespace(
app=SimpleNamespace(
state=SimpleNamespace(
deep_research_job_store=_StreamJobStore(),
deep_research_event_store=_UnavailableEventStore(),
)
),
is_disconnected=is_connected,
)
response = await deep_research.stream_job("drj_1", request)
chunk = await anext(response.body_iterator)
await response.body_iterator.aclose()
assert "durable_replay_temporarily_unavailable" in chunk
@pytest.mark.asyncio
async def test_full_report_rewrite_persists_selected_report_structure(monkeypatch):
"""A selected structure must survive durable replay and constrain the worker."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _CompletedSessionStore(_SessionStore):
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return {
"id": session_id,
"query": "测试主题",
"title": "测试主题",
"status": "completed",
"report_markdown": "# 原报告\n\n原有内容。",
"config_snapshot": {"mode": "basic"},
}
class _SourceStore:
async def list_by_session(self, *args, **kwargs): # noqa: ANN002, ANN003
# SQL-backed historical sources return native datetimes, whereas
# the job snapshot is persisted as a JSON column.
return [{"id": "src_1", "created_at": datetime(2026, 8, 21, tzinfo=UTC)}]
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
# Match the repository's JSON persistence boundary. This used to
# raise TypeError for historical source rows with datetime fields.
json.dumps(kwargs["input_snapshot"], ensure_ascii=False)
self.created = kwargs
return {"id": "drj_rewrite", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
jobs = _CapturingJobStore()
outline = "# 执行摘要\n\n# 一、核心发现\n\n# 二、行动建议"
response = await deep_research.start_full_report_rewrite_job(
"drs_1",
deep_research.FullReportRewriteRequest(
instruction="压缩重复段落,保持事实准确。",
reportOutline=outline,
operation="generate",
),
_request(
sessions=_CompletedSessionStore(),
jobs=jobs,
dispatcher=_Dispatcher(),
sources=_SourceStore(),
),
)
assert response["jobId"] == "drj_rewrite"
assert jobs.created is not None
snapshot = jobs.created["input_snapshot"]
assert snapshot["operation"] == "generate"
assert snapshot["report_outline"] == outline
assert "【报告结构要求】" in snapshot["instruction"]
assert "# 一、核心发现" in snapshot["instruction"]
assert snapshot["sources"][0]["created_at"] == "2026-08-21T00:00:00+00:00"
@pytest.mark.asyncio
async def test_regeneration_forks_report_and_selected_sources(monkeypatch):
"""Choosing another structure must never mutate the completed parent report."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _VariantSessionStore:
def __init__(self) -> None:
self.parent = {
"id": "drs_parent",
"user_id": "u1",
"title": "新能源研究",
"query": "新能源市场",
"mode": "basic",
"status": "completed",
"report_markdown": (
"# 原报告\n\n"
"已保存引用 [来源:src_parent];"
"流式引用 [[source:src_parent]];"
"已取消的资料 [[source:src_unselected]]。"
),
"config_snapshot": {"mode": "basic", "custom_outline": "# 原结构"},
"source_count": 1,
}
self.created: dict | None = None
self.updated: dict | None = None
self.deleted: list[str] = []
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return self.parent if session_id == "drs_parent" else None
async def create(self, **fields): # noqa: ANN003
self.created = fields
return {**fields, "report_markdown": None, "source_count": 0}
async def update(self, session_id: str, *, user_id: str, **fields): # noqa: ARG002
self.updated = {"session_id": session_id, **fields}
return {
"id": session_id,
"user_id": "u1",
"title": self.created["title"] if self.created else "",
"query": "新能源市场",
"mode": "basic",
"status": self.created["status"] if self.created else "draft",
"config_snapshot": self.created["config_snapshot"] if self.created else {},
**fields,
}
async def delete(self, session_id: str, *, user_id: str): # noqa: ARG002
self.deleted.append(session_id)
return True
class _VariantSourceStore:
async def clone_selected_to_session(self, source_session_id: str, target_session_id: str, *, user_id: str): # noqa: ARG002
assert source_session_id == "drs_parent"
assert target_session_id.startswith("drs_")
return {"src_parent": "src_variant"}
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
async def fake_ensure_runtime_thread(request, session, user_id): # noqa: ANN001, ARG001
await sessions.update(session["id"], user_id=user_id, runtime_thread_id="dr_variant_thread")
sessions = _VariantSessionStore()
monkeypatch.setattr(deep_research, "_ensure_chat_runtime_thread", fake_ensure_runtime_thread)
response = await deep_research.create_report_variant(
"drs_parent",
deep_research.CreateReportVariantRequest(reportOutline="# 新结构\n\n## 结论"),
_request(
sessions=sessions,
jobs=_JobStore(),
dispatcher=_Dispatcher(),
sources=_VariantSourceStore(),
),
)
assert response.id != "drs_parent"
assert response.title == "新能源研究 · 新结构报告"
assert response.report_markdown is None
assert response.status == "draft"
assert response.config == {
"mode": "basic",
"custom_outline": "# 新结构\n\n## 结论",
"report_variant": True,
"parent_session_id": "drs_parent",
}
assert sessions.parent["report_markdown"] == (
"# 原报告\n\n"
"已保存引用 [来源:src_parent];"
"流式引用 [[source:src_parent]];"
"已取消的资料 [[source:src_unselected]]。"
)
assert sessions.updated is not None
assert sessions.updated["source_count"] == 1
assert sessions.deleted == []
def test_normalise_citations_accepts_display_ready_and_legacy_markers() -> None:
source_ids = {"drs_src_1234567890abcdef1234567890abcdef"}
result, cited = report_rewrite._normalise_citations( # noqa: SLF001
"甲[来源:drs_src_1234567890abcdef1234567890abcdef];"
"乙[[source:drs_src_1234567890abcdef1234567890abcdef]];"
"丙[来源: drs_src_1234567890abcdef1234567890abcdef]。",
source_ids,
)
assert result == (
"甲[来源:drs_src_1234567890abcdef1234567890abcdef];"
"乙[来源:drs_src_1234567890abcdef1234567890abcdef];"
"丙[来源:drs_src_1234567890abcdef1234567890abcdef]。"
)
assert cited == ["drs_src_1234567890abcdef1234567890abcdef"]
def test_normalise_citations_repairs_unique_truncated_selected_id() -> None:
full_id = "drs_src_ac653c4a6f364e59a30b044e98d32e2f"
result, cited = report_rewrite._normalise_citations( # noqa: SLF001
"事实[来源:drs_src_ac653c4a6f364e59a30b044e98d32e]。",
{full_id},
)
assert result == f"事实[来源:{full_id}]。"
assert cited == [full_id]
def test_normalise_citations_still_rejects_unknown_source() -> None:
with pytest.raises(ValueError, match="未选资料来源"):
report_rewrite._normalise_citations( # noqa: SLF001
"事实[来源:drs_src_ffffffffffffffffffffffffffffffff]。",
{"drs_src_ac653c4a6f364e59a30b044e98d32e2f"},
)
def test_normalise_citations_removes_provider_truncated_marker_without_discarding_report() -> None:
result, cited = report_rewrite._normalise_citations( # noqa: SLF001
"# 参考资料\n\n已完成的正文。\n[来源:drs_src_ac653c4a6f364e59a30b044e98d32",
{"drs_src_ac653c4a6f364e59a30b044e98d32e2f"},
)
assert result == "# 参考资料\n\n已完成的正文。"
assert cited == []
@pytest.mark.asyncio
async def test_generate_operation_uses_selected_sources_not_inherited_report(monkeypatch) -> None:
captured: dict = {}
class _CompletionBackend:
def __init__(self, config) -> None: # noqa: ANN001
captured["config"] = config
async def stream_complete(self, **kwargs): # noqa: ANN003
captured.update(kwargs)
return SimpleNamespace(
text=(
"# 新报告\n\n事实[来源:drs_src_1234567890abcdef1234567890abcdef]。\n\n"
"<!-- DEERFLOW_REPORT_COMPLETE -->"
),
usage={},
)
monkeypatch.setattr(report_rewrite, "DeerFlowCompletionBackend", _CompletionBackend)
async def _ignore(_delta: str) -> None:
return None
result, cited, _usage = await report_rewrite.stream_full_report_rewrite(
session={"config_snapshot": {"mode": "basic"}, "query": "新能源汽车"},
report="# 不应进入模型的旧报告\n\n旧结构和旧措辞。",
sources=[
{
"id": "drs_src_1234567890abcdef1234567890abcdef",
"title": "已选资料",
"raw_content": "可核验事实",
}
],
instruction="按新结构生成",
style="formal_analysis",
model_name=None,
on_delta=_ignore,
operation="generate",
)
prompt = "\n".join(str(message["content"]) for message in captured["messages"])
assert "不应进入模型的旧报告" not in prompt
assert "# 已选来源" in prompt
assert captured["operation"] == "full_report_generate"
assert captured["max_tokens"] == 8_192
assert result.endswith("[来源:drs_src_1234567890abcdef1234567890abcdef]。")
assert cited == ["drs_src_1234567890abcdef1234567890abcdef"]
@pytest.mark.asyncio
async def test_report_generation_continues_when_first_segment_has_no_completion_marker(monkeypatch) -> None:
calls: list[str] = []
class _CompletionBackend:
def __init__(self, _config) -> None: # noqa: ANN001
pass
async def stream_complete(self, **kwargs): # noqa: ANN003
calls.append(str(kwargs["operation"]))
if len(calls) == 1:
return SimpleNamespace(text="# 新报告\n\n第一部分。", usage={"completion_tokens": 100})
return SimpleNamespace(
text="## 第二部分\n\n续写完成。\n\n<!-- DEERFLOW_REPORT_COMPLETE -->",
usage={"completion_tokens": 50},
)
monkeypatch.setattr(report_rewrite, "DeerFlowCompletionBackend", _CompletionBackend)
async def _ignore(_delta: str) -> None:
return None
result, cited, usage = await report_rewrite.stream_full_report_rewrite(
session={"config_snapshot": {"mode": "basic"}, "query": "新能源汽车"},
report="# 旧报告",
sources=[],
instruction="按新结构生成",
style="formal_analysis",
model_name=None,
on_delta=_ignore,
operation="generate",
)
assert calls == ["full_report_generate", "full_report_generate_continue"]
assert result == "# 新报告\n\n第一部分。\n\n## 第二部分\n\n续写完成。"
assert cited == []
assert usage["completion_tokens"] == 150
@pytest.mark.asyncio
async def test_write_report_treats_browser_materials_as_authoritative(monkeypatch):
"""POST /report must freeze the request body as the writing input."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_write", "status": "queued"}, True
class _SourceStore:
def __init__(self) -> None:
self.rows: list[dict] = []
async def upsert(self, **kwargs): # noqa: ANN003
self.rows.append(kwargs)
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
jobs = _CapturingJobStore()
sources = _SourceStore()
materials = [
{
"title": "来源甲",
"url": "https://a.test/1",
"raw_content": "正文甲",
"snippet": "摘要甲",
}
]
response = await deep_research.write_report(
"drs_1",
deep_research.WriteReportRequest(
query="前端课题",
title="前端标题",
config={"mode": "basic"},
collector_materials=materials,
),
_request(sessions=_SessionStore(), jobs=jobs, dispatcher=_Dispatcher(), sources=sources),
)
assert response["jobId"] == "drj_write"
assert response["collectorMaterialCount"] == 1
assert response["persistedSourceCount"] == 1
snapshot = jobs.created["input_snapshot"]
assert snapshot["entry"] == "chat"
assert snapshot["query"] == "前端课题"
assert snapshot["title"] == "前端标题"
assert snapshot["client_materials_authoritative"] is True
assert snapshot["collector_materials"] == materials
assert snapshot["collector_material_count"] == 1
assert sources.rows[0]["title"] == "来源甲"
assert sources.rows[0]["selected"] is True
assert response["diagnostics"] == []
@pytest.mark.asyncio
async def test_write_report_caps_job_snapshot_to_fallback_limit(monkeypatch):
"""The job snapshot only keeps 20 fallback rows even if the client sent more."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_cap", "status": "queued"}, True
class _SourceStore:
def __init__(self) -> None:
self.rows: list[dict] = []
async def upsert(self, **kwargs): # noqa: ANN003
self.rows.append(kwargs)
return kwargs
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
jobs = _CapturingJobStore()
sources = _SourceStore()
materials = [
{"title": f"来源{index}", "url": f"https://a.test/{index}", "raw_content": f"正文{index}"}
for index in range(25)
]
response = await deep_research.write_report(
"drs_1",
deep_research.WriteReportRequest(query="大量素材", collector_materials=materials),
_request(sessions=_SessionStore(), jobs=jobs, dispatcher=_Dispatcher(), sources=sources),
)
snapshot = jobs.created["input_snapshot"]
assert response["collectorMaterialCount"] == 25
assert response["persistedSourceCount"] == 25
assert len(snapshot["collector_materials"]) == 20
assert snapshot["collector_materials"][0]["title"] == "来源0"
assert snapshot["collector_materials"][-1]["title"] == "来源19"
@pytest.mark.asyncio
async def test_write_report_empty_materials_stays_authoritative(monkeypatch):
"""An empty browser list must freeze as the writing input, not trigger harvest."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_empty", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
jobs = _CapturingJobStore()
response = await deep_research.write_report(
"drs_1",
deep_research.WriteReportRequest(query="空素材课题", collector_materials=[]),
_request(sessions=_SessionStore(), jobs=jobs, dispatcher=_Dispatcher()),
)
assert response["collectorMaterialCount"] == 0
assert response["persistedSourceCount"] == 0
snapshot = jobs.created["input_snapshot"]
assert snapshot["client_materials_authoritative"] is True
assert snapshot["collector_materials"] == []
assert snapshot["collector_material_count"] == 0
assert snapshot["query"] == "空素材课题"
assert any(item["code"] == "empty_client_materials" for item in response["diagnostics"])
@pytest.mark.asyncio
async def test_write_report_force_no_materials_ignores_browser_snapshot(monkeypatch):
"""The card checkbox must freeze an empty pool even if the browser still has hits."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _CapturingJobStore(_JobStore):
def __init__(self) -> None:
self.created: dict | None = None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_force_empty", "status": "queued"}, True
class _SourceStore:
def __init__(self) -> None:
self.rows: list[dict] = []
async def upsert(self, **kwargs): # noqa: ANN003
self.rows.append(kwargs)
return kwargs
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
jobs = _CapturingJobStore()
sources = _SourceStore()
response = await deep_research.write_report(
"drs_1",
deep_research.WriteReportRequest(
query="强制无素材",
collector_materials=[{"id": "src_keep", "title": "来源甲", "raw_content": "正文"}],
force_no_materials=True,
),
_request(sessions=_SessionStore(), jobs=jobs, dispatcher=_Dispatcher(), sources=sources),
)
assert response["collectorMaterialCount"] == 0
assert response["persistedSourceCount"] == 0
assert sources.rows == []
snapshot = jobs.created["input_snapshot"]
assert snapshot["force_no_materials"] is True
assert snapshot["client_materials_authoritative"] is True
assert snapshot["collector_materials"] == []
assert snapshot["collector_material_count"] == 0
assert any(item["code"] == "force_no_materials" for item in response["diagnostics"])
@pytest.mark.asyncio
async def test_update_session_persists_collector_materials(monkeypatch):
"""Collection-complete PATCH must copy browser materials into the source table."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _SourceStore:
def __init__(self) -> None:
self.rows: list[dict] = []
async def upsert(self, **kwargs): # noqa: ANN003
self.rows.append(kwargs)
return kwargs
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sessions = _SessionStore()
sources = _SourceStore()
response = await deep_research.update_session(
"drs_1",
deep_research.UpdateSessionRequest(
status="awaiting_report",
collector_materials=[
{"title": "来源甲", "url": "https://a.test/1", "raw_content": "正文甲"},
],
),
_request(sessions=sessions, jobs=_JobStore(), dispatcher=_Dispatcher(), sources=sources),
)
assert response.status == "awaiting_report"
assert response.source_count == 1
assert sources.rows[0]["title"] == "来源甲"
assert sources.rows[0]["selected"] is True
@pytest.mark.asyncio
async def test_update_session_persists_incremental_search_batches(monkeypatch):
"""Each search batch is upserted on its own PATCH, not held until the turn ends."""
async def fake_require_user(request): # noqa: ANN001
return "u1"
class _SourceStore:
def __init__(self) -> None:
self.rows: list[dict] = []
async def upsert(self, **kwargs): # noqa: ANN003
self.rows.append(kwargs)
return kwargs
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sources = _SourceStore()
request = _request(sessions=_SessionStore(), jobs=_JobStore(), dispatcher=_Dispatcher(), sources=sources)
await deep_research.update_session(
"drs_1",
deep_research.UpdateSessionRequest(
collector_materials=[
{"title": "检索1-甲", "url": "https://a.test/1", "raw_content": "正文1"},
{"title": "检索1-乙", "url": "https://a.test/2", "raw_content": "正文2"},
],
),
request,
)
await deep_research.update_session(
"drs_1",
deep_research.UpdateSessionRequest(
collector_materials=[
{"title": "检索2-甲", "url": "https://b.test/1", "raw_content": "正文3"},
],
),
request,
)
assert [row["title"] for row in sources.rows] == ["检索1-甲", "检索1-乙", "检索2-甲"]
class _RepairSessionStore:
def __init__(self, row: dict) -> None:
self.row = row
self.updates: list[dict] = []
async def update(self, session_id: str, **fields): # noqa: ARG002
self.updates.append(fields)
self.row = {**self.row, **fields}
return self.row
class _RepairSourceStore:
def __init__(self, rows: list[dict] | None = None) -> None:
self.rows = rows or []
async def list_by_session(self, session_id, **kwargs): # noqa: ANN001, ARG002
return self.rows
def _repair_request(session_row: dict, sources: list[dict] | None = None):
sessions = _RepairSessionStore(session_row)
return _request(
sessions=sessions,
jobs=_JobStore(),
dispatcher=_Dispatcher(),
sources=_RepairSourceStore(sources),
), sessions
@pytest.mark.asyncio
async def test_session_read_repair_skips_running_job():
request, sessions = _repair_request(
{
"id": "drs_live",
"status": "running",
"active_job_id": "drj_live",
"query": "课题",
"report_markdown": "无法生成研究报告",
"error": None,
}
)
result = await deep_research._repair_invalid_persisted_report(
request,
sessions.row,
user_id="u1",
)
assert sessions.updates == []
assert result["status"] == "running"
assert result["active_job_id"] == "drj_live"
assert result["report_markdown"] == "无法生成研究报告"
assert result.get("error") is None
@pytest.mark.asyncio
async def test_session_read_repair_skips_completed_session_that_still_owns_a_job():
request, sessions = _repair_request(
{
"id": "drs_owned",
"status": "completed",
"active_job_id": "drj_still_writing",
"query": "课题",
"report_markdown": "无法生成研究报告",
"error": None,
}
)
result = await deep_research._repair_invalid_persisted_report(
request,
sessions.row,
user_id="u1",
)
assert sessions.updates == []
assert result["active_job_id"] == "drj_still_writing"
assert result["report_markdown"] == "无法生成研究报告"
@pytest.mark.asyncio
async def test_session_read_repair_replaces_refusal_without_marking_session_failed():
request, sessions = _repair_request(
{
"id": "drs_old",
"status": "completed",
"active_job_id": None,
"query": "课题",
"source_count": 0,
"report_markdown": "无法生成研究报告",
"error": "old",
}
)
result = await deep_research._repair_invalid_persisted_report(
request,
sessions.row,
user_id="u1",
)
assert sessions.updates
assert sessions.updates[0]["status"] == "completed"
assert sessions.updates[0]["error"] is None
assert "active_job_id" not in sessions.updates[0]
assert result["error"] is None
assert result["report_markdown"] != "无法生成研究报告"
assert "课题" in result["report_markdown"]
def test_job_is_stale_when_lease_owner_pid_is_dead() -> None:
now = datetime.now(UTC)
dead = {
"status": "running",
"cancel_requested": False,
"lease_owner": "w-999999-deadbeef",
"lease_until": now + timedelta(seconds=180),
}
live = {
"status": "running",
"cancel_requested": False,
"lease_owner": f"w-{os.getpid()}-alive000",
"lease_until": now + timedelta(seconds=180),
}
assert deep_research._job_is_stale(dead, now=now) is True
assert deep_research._job_is_stale(live, now=now) is False
assert deep_research._job_is_stale(
{**live, "cancel_requested": True},
now=now,
) is True
@pytest.mark.asyncio
async def test_start_job_releases_dead_worker_running_sessions(monkeypatch) -> None:
async def fake_require_user(request): # noqa: ANN001
return "u1"
now = datetime.now(UTC)
zombie = {
"id": "drs_dead",
"user_id": "u1",
"query": "旧课题",
"title": "旧课题",
"status": "running",
"active_job_id": "drj_dead",
"runtime_thread_id": "dr_t",
"config_snapshot": {"mode": "basic", "collection_mode": "chat"},
}
target = {
"id": "drs_1",
"user_id": "u1",
"query": "新课题",
"title": "新课题",
"status": "draft",
"config_snapshot": {"mode": "basic"},
}
class _Sessions:
def __init__(self) -> None:
self.rows = {"drs_dead": dict(zombie), "drs_1": dict(target)}
async def get(self, session_id: str, *, user_id: str): # noqa: ARG002
return self.rows.get(session_id)
async def list_by_user(self, *, user_id: str, status: str | None = None, limit: int = 20, **_: object): # noqa: ARG002
return [row for row in self.rows.values() if status is None or row.get("status") == status]
async def count_active_by_user(self, *, user_id: str) -> int: # noqa: ARG002
return sum(1 for row in self.rows.values() if row.get("status") in ("running", "awaiting_input"))
async def update(self, session_id: str, *, user_id: str, **fields): # noqa: ARG002
self.rows[session_id] = {**self.rows[session_id], **fields}
return self.rows[session_id]
class _Jobs:
def __init__(self) -> None:
self.jobs = {
"drj_dead": {
"id": "drj_dead",
"session_id": "drs_dead",
"status": "running",
"lease_owner": "w-999999-deadbeef",
"lease_until": now + timedelta(seconds=180),
"cancel_requested": False,
}
}
self.created = None
async def get(self, job_id: str, *, user_id: str | None = None): # noqa: ARG002
return self.jobs.get(job_id)
async def request_cancel(self, job_id: str, *, user_id: str | None = None): # noqa: ARG002
self.jobs[job_id] = {**self.jobs[job_id], "cancel_requested": True}
return self.jobs[job_id]
async def finalize_cancel(self, job_id: str, *, lease_owner: str | None = None): # noqa: ARG002
self.jobs[job_id] = {
**self.jobs[job_id],
"status": "cancelled",
"lease_owner": None,
"lease_until": None,
}
return self.jobs[job_id]
async def get_active_for_session(self, session_id: str, *, user_id: str): # noqa: ARG002
return None
async def try_create_or_get_active(self, **kwargs): # noqa: ANN003
self.created = kwargs
return {"id": "drj_new", "status": "queued"}, True
monkeypatch.setattr(deep_research, "_require_user", fake_require_user)
sessions = _Sessions()
jobs = _Jobs()
response = await deep_research.start_job(
"drs_1",
deep_research.StartJobRequest(entry="legacy"),
_request(sessions=sessions, jobs=jobs, dispatcher=_Dispatcher()),
)
assert response["jobId"] == "drj_new"
assert sessions.rows["drs_dead"]["status"] == "awaiting_report"
assert sessions.rows["drs_dead"]["active_job_id"] is None
assert jobs.jobs["drj_dead"]["status"] == "cancelled"