397 lines
18 KiB
Python
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"] == "# 已保存报告"
|