"""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" "" ), 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", 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"