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

397 lines
18 KiB
Python

"""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"] == "# 已保存报告"