1083 lines
39 KiB
Python
1083 lines
39 KiB
Python
"""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"
|