"""Regression tests for owner-scoped enterprise-research task persistence.""" from __future__ import annotations import asyncio from types import SimpleNamespace import pytest import pytest_asyncio from fastapi import FastAPI from httpx import ASGITransport, AsyncClient from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine from deerflow.persistence.base import Base from deerflow.persistence.enterprise_research import ( EnterpriseResearchReportJobRepository, EnterpriseResearchTaskRepository, MemoryEnterpriseResearchReportJobStore, MemoryEnterpriseResearchTaskStore, make_enterprise_research_task_store, ) @pytest_asyncio.fixture async def store(tmp_path): engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'enterprise-research.db'}") async with engine.begin() as conn: import deerflow.persistence.enterprise_research.model # noqa: F401 await conn.run_sync(Base.metadata.create_all) return EnterpriseResearchTaskRepository(async_sessionmaker(engine, expire_on_commit=False)) async def _create(store, *, task_id: str, user_id: str = "u1"): return await store.create_task( { "id": task_id, "user_id": user_id, "subject": "新能源汽车供应链风险", "focus": "关键原材料", "template_id": "topic", "knowledge_base_ids": ["kb-1"], "queries": ["新能源汽车供应链风险", "关键原材料"], "status": "draft", } ) @pytest.mark.asyncio async def test_task_evidence_roundtrip_and_owner_isolation(store: EnterpriseResearchTaskRepository): created = await _create(store, task_id="ert-1") assert created["sources"] == [] assert created["status"] == "draft" updated = await store.update_task( "ert-1", "u1", { "status": "ready", "selected_source_ids": ["chunk-1"], "plan_markdown": "# 供应链风险研究计划", "sources": [ { "chunk_id": "chunk-1", "title": "供应链报告", "content": "证据内容", "knowledge_base_id": "kb-1", } ], }, ) assert updated is not None assert updated["status"] == "ready" assert updated["sources"][0]["chunk_id"] == "chunk-1" assert updated["selected_source_ids"] == ["chunk-1"] assert updated["plan_markdown"] == "# 供应链风险研究计划" assert await store.get_task("ert-1", "u2") is None assert await store.update_task("ert-1", "u2", {"status": "failed"}) is None @pytest.mark.asyncio async def test_list_is_owner_scoped_and_delete_requires_owner(store: EnterpriseResearchTaskRepository): await _create(store, task_id="ert-u1", user_id="u1") await _create(store, task_id="ert-u2", user_id="u2") assert [task["id"] for task in await store.list_tasks("u1")] == ["ert-u1"] assert await store.delete_task("ert-u1", "u2") is False assert await store.delete_task("ert-u1", "u1") is True assert await store.get_task("ert-u1", "u1") is None @pytest.mark.asyncio async def test_memory_store_matches_sql_contract(): store = make_enterprise_research_task_store(None) assert isinstance(store, MemoryEnterpriseResearchTaskStore) await _create(store, task_id="memory-1") assert (await store.get_task("memory-1", "u1"))["subject"] == "新能源汽车供应链风险" assert await store.delete_task("memory-1", "u1") is True @pytest.mark.asyncio async def test_report_job_freezes_sources_and_is_owner_scoped(tmp_path): engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'enterprise-report-job.db'}") async with engine.begin() as conn: import deerflow.persistence.enterprise_research.model # noqa: F401 await conn.run_sync(Base.metadata.create_all) jobs = EnterpriseResearchReportJobRepository(async_sessionmaker(engine, expire_on_commit=False)) created = await jobs.create_job( { "id": "erj-1", "task_id": "ert-1", "user_id": "u1", "plan_snapshot": "# 计划", "sources": [{"chunk_id": "source-1", "content": "冻结证据"}], "action": "rewrite", "instruction": "优化结构", "original_report": "# 原报告", } ) assert created["status"] == "queued" assert created["sources"][0]["content"] == "冻结证据" assert created["action"] == "rewrite" assert created["original_report"] == "# 原报告" deduplicated = await jobs.create_job( { "id": "erj-duplicate", "task_id": "ert-1", "user_id": "u1", "plan_snapshot": "# 其他计划", "sources": [], "action": "initial", } ) assert deduplicated["id"] == "erj-1" assert await jobs.get_job("erj-1", "u2") is None assert (await jobs.get_active_for_task("ert-1", "u1"))["id"] == "erj-1" assert await jobs.claim_job("erj-1", "u1") is True assert await jobs.claim_job("erj-1", "u1") is False running = (await jobs.list_recoverable())[0] assert running["status"] == "running" assert await jobs.requeue_job("erj-1", "u1", expected_updated_at="stale-snapshot") is False assert await jobs.requeue_job("erj-1", "u1", expected_updated_at=running["updated_at"]) is True assert await jobs.requeue_job("erj-1", "u1", expected_updated_at=running["updated_at"]) is False assert await jobs.claim_job("erj-1", "u1") is True completed = await jobs.transition_job( "erj-1", "u1", expected_statuses={"running"}, data={"status": "completed", "phase": "completed", "report_markdown": "# 报告"}, ) assert completed is not None and completed["report_markdown"] == "# 报告" assert await jobs.transition_job( "erj-1", "u1", expected_statuses={"running"}, data={"status": "cancelled"} ) is None assert await jobs.get_active_for_task("ert-1", "u1") is None next_job = await jobs.create_job( {"id": "erj-2", "task_id": "ert-1", "user_id": "u1", "plan_snapshot": "# 新计划", "sources": []} ) assert next_job["id"] == "erj-2" @pytest.mark.asyncio async def test_memory_report_job_store_has_matching_owner_contract(): jobs = MemoryEnterpriseResearchReportJobStore() await jobs.create_job({"id": "memory-job", "task_id": "task", "user_id": "u1", "plan_snapshot": "x"}) assert await jobs.get_job("memory-job", "u2") is None assert (await jobs.list_jobs("task", "u1"))[0]["id"] == "memory-job" assert await jobs.claim_job("memory-job", "u1") is True assert await jobs.claim_job("memory-job", "u1") is False running = await jobs.get_job("memory-job", "u1") assert running is not None assert await jobs.requeue_job("memory-job", "u1", expected_updated_at=running["updated_at"]) is True assert await jobs.claim_job("memory-job", "u1") is True completed = await jobs.transition_job( "memory-job", "u1", expected_statuses={"running"}, data={"status": "completed", "phase": "completed"} ) assert completed is not None and completed["status"] == "completed" def test_evidence_snapshot_deduplicates_and_never_exposes_unmapped_remote_bases(): from app.gateway.routers.enterprise_research import _dedupe_evidence mappings = {"remote-kb-1": {"id": "local-kb-1", "name": "用户可见知识库", "weknora_id": "remote-kb-1"}} evidence = _dedupe_evidence( [ {"chunk_id": "chunk-1", "content": "第一段", "knowledge_base_id": "remote-kb-1", "title": "文档 A"}, {"chunk_id": "chunk-1", "content": "重复段", "knowledge_base_id": "remote-kb-1", "title": "文档 A"}, {"chunk_id": "private-chunk", "content": "不应显示", "knowledge_base_id": "other-user-kb", "title": "私有文档"}, ], mappings, ) assert len(evidence) == 1 assert evidence[0]["knowledge_base_id"] == "local-kb-1" assert "other-user-kb" not in str(evidence) def test_report_executor_uses_distinct_initial_rewrite_and_followup_prompts(): from app.gateway.enterprise_research_report_executor import EnterpriseResearchReportExecutor base = { "plan_snapshot": "# 计划", "sources": [{"chunk_id": "source-1", "title": "内部资料", "content": "证据正文"}], "original_report": "# 原报告\n原始正文", "instruction": "压缩重复段落", "target_length": 3_000, "style": "analytical", } initial = EnterpriseResearchReportExecutor._messages({**base, "action": "initial"}) rewrite = EnterpriseResearchReportExecutor._messages({**base, "action": "rewrite"}) followup = EnterpriseResearchReportExecutor._messages({**base, "action": "followup"}) assert "完整 Markdown 报告" in initial[0]["content"] assert "原报告" not in initial[1]["content"] assert "全文改写" in rewrite[1]["content"] and "# 原报告" in rewrite[1]["content"] assert "不要重写整篇报告" in followup[0]["content"] assert "压缩重复段落" in followup[1]["content"] summary = EnterpriseResearchReportExecutor._summary( "# 新报告\n\n新正文", 2, "rewrite", original_report="# 旧报告\n\n旧正文", instruction="调整结构" ) assert "原文" in summary and "标题节点 1→1" in summary and "调整结构" in summary @pytest.mark.asyncio async def test_executor_recovery_requeues_exact_running_snapshot(monkeypatch): from app.gateway.enterprise_research_report_executor import EnterpriseResearchReportExecutor tasks = MemoryEnterpriseResearchTaskStore() jobs = MemoryEnterpriseResearchReportJobStore() await jobs.create_job({"id": "recover-job", "task_id": "task", "user_id": "u1", "plan_snapshot": "# 计划"}) assert await jobs.claim_job("recover-job", "u1") is True executor = EnterpriseResearchReportExecutor(tasks, jobs) started: list[tuple[str, str]] = [] async def remember(job_id: str, user_id: str) -> None: started.append((job_id, user_id)) monkeypatch.setattr(executor, "ensure_started", remember) await executor.recover_pending() recovered = await jobs.get_job("recover-job", "u1") assert recovered is not None and recovered["status"] == "queued" assert started == [("recover-job", "u1")] @pytest.mark.asyncio async def test_enterprise_research_http_plan_job_dedupe_and_cancel(monkeypatch): import app.gateway.routers.enterprise_research as enterprise_router tasks = MemoryEnterpriseResearchTaskStore() jobs = MemoryEnterpriseResearchReportJobStore() class NoopExecutor: def __init__(self) -> None: self.started: list[tuple[str, str]] = [] async def ensure_started(self, job_id: str, user_id: str) -> None: self.started.append((job_id, user_id)) async def cancel(self, job_id: str, user_id: str): job = await jobs.get_job(job_id, user_id) if job is None: return None cancelled = await jobs.transition_job( job_id, user_id, expected_statuses={"queued", "running"}, data={"status": "cancelled", "phase": "cancelled"}, ) await tasks.update_task(job["task_id"], user_id, {"active_report_job_id": None}) return cancelled async def anonymous_actor(_request): return None monkeypatch.setattr(enterprise_router, "get_optional_user_from_request", anonymous_actor) app = FastAPI() app.include_router(enterprise_router.router) app.state.enterprise_research_task_store = tasks app.state.enterprise_research_report_job_store = jobs app.state.enterprise_research_report_executor = NoopExecutor() async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: created = await client.post( "/api/enterprise-research/tasks", json={"subject": "企业风险", "queries": ["企业风险"], "knowledge_base_ids": []}, ) assert created.status_code == 201 task_id = created.json()["id"] await tasks.update_task( task_id, "default", { "status": "ready", "selected_source_ids": ["source-1"], "sources": [{"chunk_id": "source-1", "title": "内部材料", "content": "已核验事实"}], }, ) plan = await client.post(f"/api/enterprise-research/tasks/{task_id}/plan/draft") assert plan.status_code == 200 and plan.json()["plan_markdown"] saved = await client.put( f"/api/enterprise-research/tasks/{task_id}/plan", json={"plan_markdown": "# 企业风险报告", "selected_source_ids": ["source-1"]}, ) assert saved.status_code == 200 first = await client.post( f"/api/enterprise-research/tasks/{task_id}/report-jobs", json={"action": "initial", "style": "formal", "target_length": 2_000}, ) assert first.status_code == 201 duplicate = await client.post( f"/api/enterprise-research/tasks/{task_id}/report-jobs", json={"action": "initial"}, ) assert duplicate.status_code == 201 and duplicate.json()["id"] == first.json()["id"] cancelled = await client.post( f"/api/enterprise-research/tasks/{task_id}/report-jobs/{first.json()['id']}/cancel" ) assert cancelled.status_code == 200 and cancelled.json()["status"] == "cancelled" @pytest.mark.asyncio async def test_executor_stream_commit_separates_reasoning_and_followup(monkeypatch): import app.gateway.enterprise_research_report_executor as executor_module tasks = MemoryEnterpriseResearchTaskStore() jobs = MemoryEnterpriseResearchReportJobStore() await _create(tasks, task_id="stream-task") await tasks.update_task( "stream-task", "u1", {"plan_markdown": "# 研究计划", "report_markdown": "# 旧报告", "report_summary": "旧摘要"}, ) class FakeBackend: def __init__(self, _config) -> None: pass async def stream_complete(self, *, messages, on_delta, on_reasoning, **_kwargs): is_followup = "不要重写整篇报告" in messages[0]["content"] text = "追问回答" if is_followup else "# 新报告\n正文 [资料:source-1]" await on_reasoning("模型内部思考") for delta in (text[:4], text[4:]): await on_delta(delta) return SimpleNamespace(text=text, model="fake-model") monkeypatch.setattr(executor_module, "DeerFlowCompletionBackend", FakeBackend) executor = executor_module.EnterpriseResearchReportExecutor(tasks, jobs) await jobs.create_job( { "id": "initial-job", "task_id": "stream-task", "user_id": "u1", "plan_snapshot": "# 研究计划", "sources": [{"chunk_id": "source-1", "content": "证据"}], "action": "initial", } ) await executor._run("initial-job", "u1") completed = await jobs.get_job("initial-job", "u1") task = await tasks.get_task("stream-task", "u1") assert completed is not None and completed["status"] == "completed" assert completed["report_markdown"] == "# 新报告\n正文 [资料:source-1]" assert "模型内部思考" not in completed["report_markdown"] assert task is not None and task["report_markdown"] == completed["report_markdown"] await jobs.create_job( { "id": "followup-job", "task_id": "stream-task", "user_id": "u1", "plan_snapshot": "# 研究计划", "sources": [{"chunk_id": "source-1", "content": "证据"}], "action": "followup", "instruction": "结论是什么", "original_report": completed["report_markdown"], } ) await executor._run("followup-job", "u1") followup = await jobs.get_job("followup-job", "u1") unchanged = await tasks.get_task("stream-task", "u1") assert followup is not None and followup["report_markdown"] == "追问回答" assert unchanged is not None and unchanged["report_markdown"] == completed["report_markdown"] @pytest.mark.asyncio async def test_cross_process_cancel_cannot_be_overwritten_by_late_model_result(monkeypatch): import app.gateway.enterprise_research_report_executor as executor_module tasks = MemoryEnterpriseResearchTaskStore() jobs = MemoryEnterpriseResearchReportJobStore() await _create(tasks, task_id="cancel-task") await tasks.update_task("cancel-task", "u1", {"report_markdown": "# 已保存报告"}) await jobs.create_job( {"id": "cancel-job", "task_id": "cancel-task", "user_id": "u1", "plan_snapshot": "# 计划"} ) started = asyncio.Event() release = asyncio.Event() class SlowBackend: def __init__(self, _config) -> None: pass async def stream_complete(self, *, on_delta, **_kwargs): started.set() await release.wait() await on_delta("# 不应提交的迟到结果") return SimpleNamespace(text="# 不应提交的迟到结果", model="fake-model") monkeypatch.setattr(executor_module, "DeerFlowCompletionBackend", SlowBackend) writer = executor_module.EnterpriseResearchReportExecutor(tasks, jobs) remote_gateway = executor_module.EnterpriseResearchReportExecutor(tasks, jobs) running = asyncio.create_task(writer._run("cancel-job", "u1")) await started.wait() cancelled = await remote_gateway.cancel("cancel-job", "u1") assert cancelled is not None and cancelled["status"] == "cancelled" release.set() with pytest.raises(asyncio.CancelledError): await running final_job = await jobs.get_job("cancel-job", "u1") final_task = await tasks.get_task("cancel-task", "u1") assert final_job is not None and final_job["status"] == "cancelled" assert final_task is not None and final_task["report_markdown"] == "# 已保存报告"