deerflow-code/offline-backend-20260512/backend/app/gateway/routers/enterprise_research.py
2026-09-07 18:24:55 +08:00

711 lines
32 KiB
Python

"""Independent enterprise-research workbench API.
This router intentionally does not alter the established ``/api/deep-research``
workflow. It stores a small owner-scoped research task and calls WeKnora through
the same DeerFlow-owned authorization mapping used by the knowledge-base page.
"""
from __future__ import annotations
import asyncio
import json
import logging
from typing import Any, Literal
from urllib.parse import urlencode
from uuid import uuid4
from fastapi import APIRouter, HTTPException, Query, Request, status
from pydantic import BaseModel, Field
from starlette.responses import StreamingResponse
from app.gateway.deps import get_optional_user_from_request
from deerflow.agents.deep_research.adapters.llm import DeerFlowCompletionBackend
from deerflow.agents.deep_research.config import DeepResearchConfig
from deerflow.integrations.weknora.client import WeKnoraError
from deerflow.integrations.weknora.runtime import build_weknora_client, get_resolved_llmwiki_runtime
from deerflow.persistence.enterprise_research import EnterpriseResearchReportJobStore, EnterpriseResearchTaskStore
from deerflow.persistence.llmwiki import LlmWikiStore
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/enterprise-research", tags=["enterprise-research"])
TaskStatus = Literal["draft", "collecting", "ready", "partial", "failed"]
class EnterpriseResearchEvidence(BaseModel):
chunk_id: str = Field(default="", max_length=256)
content: str = Field(default="", max_length=4_000)
knowledge_id: str = Field(default="", max_length=256)
knowledge_base_id: str = Field(default="", max_length=256)
knowledge_base_name: str = Field(default="", max_length=512)
title: str = Field(default="", max_length=1_024)
filename: str = Field(default="", max_length=1_024)
score: float = 0.0
chunk_index: int = 0
url: str = Field(default="", max_length=4_096)
class EnterpriseResearchTaskResponse(BaseModel):
id: str
subject: str
focus: str
template_id: str
knowledge_base_ids: list[str]
queries: list[str]
selected_source_ids: list[str] = Field(default_factory=list)
sources: list[EnterpriseResearchEvidence]
plan_markdown: str | None = None
report_markdown: str | None = None
report_summary: str | None = None
report_model_name: str | None = None
active_report_job_id: str | None = None
status: TaskStatus
warning: str | None = None
created_at: str | None = None
updated_at: str | None = None
class EnterpriseResearchTaskListResponse(BaseModel):
tasks: list[EnterpriseResearchTaskResponse]
class EnterpriseResearchTaskCreate(BaseModel):
subject: str = Field(..., min_length=1, max_length=180)
focus: str = Field(default="", max_length=800)
template_id: str = Field(default="topic", min_length=1, max_length=64)
knowledge_base_ids: list[str] = Field(default_factory=list, max_length=100)
queries: list[str] = Field(..., min_length=1, max_length=8)
class EnterpriseResearchPlanUpdate(BaseModel):
plan_markdown: str = Field(..., min_length=1, max_length=16_000)
selected_source_ids: list[str] = Field(default_factory=list, max_length=48)
class EnterpriseResearchAiPlanCreate(BaseModel):
model_name: str | None = Field(default=None, max_length=256)
instruction: str = Field(default="", max_length=2_000)
class EnterpriseResearchReportJobCreate(BaseModel):
model_name: str | None = Field(default=None, max_length=256)
action: Literal["initial", "rewrite", "followup"] = "initial"
instruction: str = Field(default="", max_length=4_000)
style: Literal["formal", "analytical", "concise"] = "formal"
target_length: int | None = Field(default=None, ge=800, le=30_000)
class EnterpriseResearchReportJobResponse(BaseModel):
id: str
task_id: str
status: Literal["queued", "running", "completed", "failed", "cancelled"]
phase: str
model_name: str | None = None
action: Literal["initial", "rewrite", "followup"] = "initial"
instruction: str | None = None
style: str | None = None
target_length: int | None = None
original_report: str | None = None
report_markdown: str = ""
report_summary: str | None = None
error_message: str | None = None
created_at: str | None = None
updated_at: str | None = None
class EnterpriseResearchReportJobListResponse(BaseModel):
jobs: list[EnterpriseResearchReportJobResponse]
def _task_store(request: Request) -> EnterpriseResearchTaskStore:
store = getattr(request.app.state, "enterprise_research_task_store", None)
if store is None:
raise HTTPException(status_code=503, detail="企业深度研究任务存储暂不可用")
return store
def _report_job_store(request: Request) -> EnterpriseResearchReportJobStore:
store = getattr(request.app.state, "enterprise_research_report_job_store", None)
if store is None:
raise HTTPException(status_code=503, detail="企业研究报告任务存储暂不可用")
return store
def _report_executor(request: Request):
executor = getattr(request.app.state, "enterprise_research_report_executor", None)
if executor is None:
raise HTTPException(status_code=503, detail="企业研究报告执行器暂不可用")
return executor
def _report_start_lock(request: Request) -> asyncio.Lock:
lock = getattr(request.app.state, "enterprise_research_report_start_lock", None)
if lock is None:
lock = asyncio.Lock()
request.app.state.enterprise_research_report_start_lock = lock
return lock
def _llmwiki_store(request: Request) -> LlmWikiStore:
store = getattr(request.app.state, "llmwiki_store", None)
if store is None:
raise HTTPException(status_code=503, detail="知识库权限存储暂不可用")
return store
async def _actor(request: Request) -> tuple[str, bool]:
user = await get_optional_user_from_request(request)
if user is None:
return "default", True
return str(user.id), getattr(user, "system_role", None) == "admin"
def _response(row: dict[str, Any]) -> EnterpriseResearchTaskResponse:
return EnterpriseResearchTaskResponse(
id=str(row["id"]),
subject=str(row.get("subject") or ""),
focus=str(row.get("focus") or ""),
template_id=str(row.get("template_id") or "topic"),
knowledge_base_ids=[str(item) for item in row.get("knowledge_base_ids") or []],
queries=[str(item) for item in row.get("queries") or []],
selected_source_ids=[str(item) for item in row.get("selected_source_ids") or []],
sources=[EnterpriseResearchEvidence.model_validate(item) for item in row.get("sources") or [] if isinstance(item, dict)],
plan_markdown=str(row["plan_markdown"]) if row.get("plan_markdown") else None,
report_markdown=str(row["report_markdown"]) if row.get("report_markdown") else None,
report_summary=str(row["report_summary"]) if row.get("report_summary") else None,
report_model_name=str(row["report_model_name"]) if row.get("report_model_name") else None,
active_report_job_id=str(row["active_report_job_id"]) if row.get("active_report_job_id") else None,
status=row.get("status") if row.get("status") in {"draft", "collecting", "ready", "partial", "failed"} else "failed",
warning=str(row["warning"]) if row.get("warning") else None,
created_at=row.get("created_at"),
updated_at=row.get("updated_at"),
)
def _report_job_response(row: dict[str, Any]) -> EnterpriseResearchReportJobResponse:
job_status = str(row.get("status") or "failed")
if job_status not in {"queued", "running", "completed", "failed", "cancelled"}:
job_status = "failed"
action = str(row.get("action") or "initial")
if action not in {"initial", "rewrite", "followup"}:
action = "initial"
return EnterpriseResearchReportJobResponse(
id=str(row["id"]), task_id=str(row["task_id"]), status=job_status, phase=str(row.get("phase") or job_status),
model_name=str(row["model_name"]) if row.get("model_name") else None,
action=action,
instruction=str(row["instruction"]) if row.get("instruction") else None,
style=str(row["style"]) if row.get("style") else None,
target_length=int(row["target_length"]) if row.get("target_length") else None,
original_report=str(row["original_report"]) if row.get("original_report") else None,
report_markdown=str(row.get("report_markdown") or ""),
report_summary=str(row["report_summary"]) if row.get("report_summary") else None,
error_message=str(row["error_message"]) if row.get("error_message") else None,
created_at=row.get("created_at"), updated_at=row.get("updated_at"),
)
def _normalise_ids(values: list[str]) -> list[str]:
seen: set[str] = set()
result: list[str] = []
for value in values:
item = str(value).strip()
if item and item not in seen:
seen.add(item)
result.append(item)
return result[:100]
def _normalise_queries(values: list[str]) -> list[str]:
seen: set[str] = set()
result: list[str] = []
for value in values:
item = str(value).strip()[:512]
if item and item not in seen:
seen.add(item)
result.append(item)
return result[:8]
def _default_plan(task: dict[str, Any], selected_source_ids: list[str]) -> str:
"""Build a transparent editable plan; report prose is a later, separate job."""
subject = str(task.get("subject") or "研究主题")
focus = str(task.get("focus") or "")
template_id = str(task.get("template_id") or "topic")
if template_id == "event":
chapters = [
"一、事件边界与时间线",
"二、相关主体与立场",
"三、影响传导与风险研判",
"四、后续观察与应对建议",
]
elif template_id == "policy":
chapters = [
"一、政策与行业背景",
"二、现状、数据与变化趋势",
"三、传导机制与关键影响",
"四、风险、机会与行动建议",
]
else:
chapters = [
"一、研究对象与事实背景",
"二、关键问题、观点与分歧",
"三、影响、风险与机会",
"四、结论、待验证问题与建议",
]
sources_by_id = {str(item.get("chunk_id") or ""): item for item in task.get("sources") or [] if isinstance(item, dict)}
selected_titles = [
str(sources_by_id[item].get("title") or sources_by_id[item].get("filename") or "未命名资料")
for item in selected_source_ids
if item in sources_by_id
]
source_lines = "\n".join(f"- {title}" for title in selected_titles[:12]) or "- 暂无已选资料"
scope = focus or "围绕主题的事实、影响与可验证判断"
return (
f"# {subject}研究计划\n\n"
"## 研究目标\n"
f"围绕“{subject}”形成可追溯、可核验的研究结论;重点范围:{scope}。\n\n"
"## 核心问题\n"
"1. 当前已知事实、背景和时间范围是什么?\n"
f"2. 与“{subject}”相关的主要主体、观点或变化是什么?\n"
"3. 证据能够支持哪些影响、风险、机会或待验证判断?\n\n"
"## 报告结构\n"
+ "\n".join(f"{chapter}\n- 说明本章需回答的核心问题,并在正文中标记资料依据。" for chapter in chapters)
+ "\n\n## 已选证据\n"
+ source_lines
+ "\n\n## 写作约束\n"
"- 仅使用已选资料作为事实依据;不补充未经验证的外部事实。\n"
"- 对资料不足或观点相互矛盾之处明确标记,避免将推测写成结论。\n"
"- 生成报告前可继续编辑本计划与已选证据。\n"
)
def _validate_selected_source_ids(task: dict[str, Any], values: list[str]) -> list[str]:
available_ids = {str(item.get("chunk_id") or "") for item in task.get("sources") or [] if isinstance(item, dict)}
selected_ids = _normalise_ids(values)[:48]
if any(item not in available_ids for item in selected_ids):
raise HTTPException(status_code=422, detail="选中的资料已变化,请刷新后重新选择")
return selected_ids
async def _visible_knowledge_bases(request: Request, task: dict[str, Any], user_id: str, is_admin: bool) -> list[dict[str, Any]]:
store = _llmwiki_store(request)
selected_ids = _normalise_ids(task.get("knowledge_base_ids") or [])
if not selected_ids:
return await store.list_visible(user_id, scope="all", is_admin=False)
rows: list[dict[str, Any]] = []
for mapping_id in selected_ids:
row = await store.get_authorized(mapping_id, user_id, write=False, is_admin=is_admin)
if row is None:
# Never reveal whether this belongs to a different account.
raise HTTPException(status_code=404, detail="知识库不存在或无访问权限")
rows.append(row)
return rows
def _client(request: Request):
runtime = get_resolved_llmwiki_runtime(getattr(request.app.state, "config", None))
if not runtime.weknora_enabled:
raise HTTPException(status_code=409, detail="WeKnora 知识服务尚未启用")
try:
return build_weknora_client(runtime)
except RuntimeError as exc:
raise HTTPException(status_code=503, detail=str(exc)) from None
def _evidence_view(result: dict[str, Any], mappings: dict[str, dict[str, Any]]) -> dict[str, Any] | None:
remote_base_id = str(result.get("knowledge_base_id") or "")
mapping = mappings.get(remote_base_id)
if mapping is None:
return None
params = urlencode({"knowledge_base_id": mapping["id"], "chunk_id": result.get("chunk_id") or ""})
return {
"chunk_id": str(result.get("chunk_id") or ""),
"content": str(result.get("content") or "")[:4_000],
"knowledge_id": str(result.get("knowledge_id") or ""),
"knowledge_base_id": str(mapping["id"]),
"knowledge_base_name": str(mapping.get("name") or ""),
"title": str(result.get("title") or ""),
"filename": str(result.get("filename") or ""),
"score": float(result.get("score") or 0),
"chunk_index": int(result.get("chunk_index") or 0),
"url": f"/weknora-source-preview?{params}",
}
def _dedupe_evidence(results: list[dict[str, Any]], mappings: dict[str, dict[str, Any]]) -> list[dict[str, Any]]:
seen: set[str] = set()
evidence: list[dict[str, Any]] = []
for result in results:
item = _evidence_view(result, mappings)
if item is None:
continue
key = item["chunk_id"] or f"{item['knowledge_id']}:{item['chunk_index']}:{item['title']}"
if key in seen:
continue
seen.add(key)
evidence.append(item)
if len(evidence) >= 48:
break
return evidence
@router.get("/tasks", response_model=EnterpriseResearchTaskListResponse)
async def list_tasks(request: Request, limit: int = Query(default=20, ge=1, le=50)) -> EnterpriseResearchTaskListResponse:
user_id, _ = await _actor(request)
rows = await _task_store(request).list_tasks(user_id, limit=limit)
return EnterpriseResearchTaskListResponse(tasks=[_response(row) for row in rows])
@router.post("/tasks", response_model=EnterpriseResearchTaskResponse, status_code=status.HTTP_201_CREATED)
async def create_task(request: Request, body: EnterpriseResearchTaskCreate) -> EnterpriseResearchTaskResponse:
user_id, _ = await _actor(request)
queries = _normalise_queries(body.queries)
if not queries:
raise HTTPException(status_code=422, detail="至少需要一个有效检索词")
row = await _task_store(request).create_task(
{
"id": f"ert_{uuid4().hex}",
"user_id": user_id,
"subject": body.subject.strip(),
"focus": body.focus.strip(),
"template_id": body.template_id.strip(),
"knowledge_base_ids": _normalise_ids(body.knowledge_base_ids),
"queries": queries,
"status": "draft",
}
)
return _response(row)
@router.post("/tasks/{task_id}/collect", response_model=EnterpriseResearchTaskResponse)
async def collect_task(request: Request, task_id: str) -> EnterpriseResearchTaskResponse:
"""Collect a bounded, user-authorized WeKnora evidence snapshot for a task."""
user_id, is_admin = await _actor(request)
task_store = _task_store(request)
task = await task_store.get_task(task_id, user_id)
if task is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
task = await task_store.update_task(task_id, user_id, {"status": "collecting", "warning": None})
assert task is not None
try:
bases = await _visible_knowledge_bases(request, task, user_id, is_admin)
if not bases:
row = await task_store.update_task(task_id, user_id, {"status": "failed", "warning": "当前账号没有可用知识库。"})
return _response(row or task)
client = _client(request)
remote_ids = [str(row["weknora_id"]) for row in bases]
outcomes = await asyncio.gather(
*(client.search(query, remote_ids) for query in task["queries"]),
return_exceptions=True,
)
raw_results: list[dict[str, Any]] = []
failed_queries = 0
for outcome in outcomes:
if isinstance(outcome, Exception):
failed_queries += 1
logger.warning("Enterprise research query failed: %s", outcome)
else:
raw_results.extend(outcome)
sources = _dedupe_evidence(raw_results, {str(row["weknora_id"]): row for row in bases})
if not sources:
task_status: TaskStatus = "failed"
warning = "未检索到可用证据,请调整研究主题或检查知识库解析状态。"
elif failed_queries:
task_status = "partial"
warning = f"{failed_queries} 个检索方向暂时不可用,已保留其余结果。"
else:
task_status = "ready"
warning = None
row = await task_store.update_task(task_id, user_id, {"status": task_status, "sources": sources, "warning": warning})
return _response(row or task)
except HTTPException:
raise
except WeKnoraError as exc:
await task_store.update_task(task_id, user_id, {"status": "failed", "warning": str(exc)})
raise HTTPException(status_code=exc.status_code, detail=str(exc)) from None
except Exception:
logger.exception("Enterprise research collection failed for %s", task_id)
await task_store.update_task(task_id, user_id, {"status": "failed", "warning": "证据收集出现异常,请稍后重试。"})
raise HTTPException(status_code=502, detail="证据收集失败,请稍后重试") from None
@router.post("/tasks/{task_id}/plan/draft", response_model=EnterpriseResearchTaskResponse)
async def create_plan_draft(request: Request, task_id: str) -> EnterpriseResearchTaskResponse:
"""Create a deterministic, editable research-plan draft from persisted evidence."""
user_id, _ = await _actor(request)
task_store = _task_store(request)
task = await task_store.get_task(task_id, user_id)
if task is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
if not task.get("sources"):
raise HTTPException(status_code=409, detail="请先完成证据收集,再生成研究计划")
selected_source_ids = _validate_selected_source_ids(task, [
str(item.get("chunk_id") or "") for item in task["sources"][:24] if isinstance(item, dict)
])
row = await task_store.update_task(
task_id,
user_id,
{
"selected_source_ids": selected_source_ids,
"plan_markdown": _default_plan(task, selected_source_ids),
},
)
return _response(row or task)
@router.put("/tasks/{task_id}/plan", response_model=EnterpriseResearchTaskResponse)
async def update_plan(request: Request, task_id: str, body: EnterpriseResearchPlanUpdate) -> EnterpriseResearchTaskResponse:
"""Persist the user's selected evidence and editable plan before report generation."""
user_id, _ = await _actor(request)
task_store = _task_store(request)
task = await task_store.get_task(task_id, user_id)
if task is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
if not task.get("sources"):
raise HTTPException(status_code=409, detail="当前任务没有可用于编排的证据")
selected_source_ids = _validate_selected_source_ids(task, body.selected_source_ids)
plan_markdown = body.plan_markdown.strip()
if not plan_markdown:
raise HTTPException(status_code=422, detail="研究计划不能为空")
row = await task_store.update_task(
task_id,
user_id,
{"selected_source_ids": selected_source_ids, "plan_markdown": plan_markdown},
)
return _response(row or task)
@router.post("/tasks/{task_id}/plan/ai", response_model=EnterpriseResearchTaskResponse)
async def create_ai_plan(
request: Request, task_id: str, body: EnterpriseResearchAiPlanCreate
) -> EnterpriseResearchTaskResponse:
"""Use the configured model to improve a plan, without collecting new evidence."""
user_id, _ = await _actor(request)
task_store = _task_store(request)
task = await task_store.get_task(task_id, user_id)
if task is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
if not task.get("sources"):
raise HTTPException(status_code=409, detail="请先完成证据收集")
selected_ids = list(task.get("selected_source_ids") or [])
if not selected_ids:
selected_ids = [
str(item.get("chunk_id") or "")
for item in task["sources"][:24]
if isinstance(item, dict) and item.get("chunk_id")
]
selected_ids = _validate_selected_source_ids(task, selected_ids)
selected = {
str(item.get("chunk_id") or ""): item
for item in task.get("sources") or []
if isinstance(item, dict)
}
evidence_titles = "\n".join(
f"- {selected[item].get('title') or selected[item].get('filename') or item}"
for item in selected_ids
if item in selected
)
current = str(task.get("plan_markdown") or _default_plan(task, selected_ids))
backend = DeerFlowCompletionBackend(
DeepResearchConfig(mode="basic", language="zh-CN", smart_model=(body.model_name or "").strip() or None).clamp()
)
try:
result = await backend.stream_complete(
model_role="smart",
operation="enterprise_research_plan",
messages=[
{
"role": "system",
"content": (
"你是企业研究规划师。仅优化研究目标、核心问题、章节结构和证据使用计划;"
"不要撰写报告正文,不要联网,不要虚构资料。仅输出 Markdown 研究计划。"
),
},
{
"role": "user",
"content": (
f"研究主题:{task['subject']}\n重点范围:{task.get('focus') or '未限定'}\n"
f"用户补充要求:{body.instruction or '提高结构完整性和章节关联性'}\n\n"
f"当前计划:\n{current}\n\n已选资料标题:\n{evidence_titles}"
),
},
],
)
except Exception as exc:
logger.exception("enterprise research AI plan generation failed task=%s", task_id)
raise HTTPException(status_code=502, detail="研究计划生成失败,请检查模型配置后重试") from exc
plan = result.text.strip()
if plan.startswith("```"):
lines = plan.splitlines()[1:]
if lines and lines[-1].strip().startswith("```"):
lines.pop()
plan = "\n".join(lines).strip()
if not plan:
raise HTTPException(status_code=502, detail="模型没有返回可用研究计划")
row = await task_store.update_task(
task_id, user_id, {"plan_markdown": plan[:16_000], "selected_source_ids": selected_ids}
)
return _response(row or task)
@router.get("/tasks/{task_id}/report-jobs", response_model=EnterpriseResearchReportJobListResponse)
async def list_report_jobs(
request: Request, task_id: str, limit: int = Query(default=10, ge=1, le=20)
) -> EnterpriseResearchReportJobListResponse:
user_id, _ = await _actor(request)
if await _task_store(request).get_task(task_id, user_id) is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
rows = await _report_job_store(request).list_jobs(task_id, user_id, limit=limit)
return EnterpriseResearchReportJobListResponse(jobs=[_report_job_response(row) for row in rows])
@router.post("/tasks/{task_id}/report-jobs", response_model=EnterpriseResearchReportJobResponse, status_code=status.HTTP_201_CREATED)
async def create_report_job(
request: Request, task_id: str, body: EnterpriseResearchReportJobCreate
) -> EnterpriseResearchReportJobResponse:
"""Freeze the approved plan/evidence, then launch an independent report job."""
user_id, _ = await _actor(request)
task_store = _task_store(request)
job_store = _report_job_store(request)
executor = _report_executor(request)
async with _report_start_lock(request):
task = await task_store.get_task(task_id, user_id)
if task is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
active = await job_store.get_active_for_task(task_id, user_id)
if active is not None:
await executor.ensure_started(active["id"], user_id)
return _report_job_response(active)
plan = str(task.get("plan_markdown") or "").strip()
selected_ids = _validate_selected_source_ids(task, list(task.get("selected_source_ids") or []))
if not plan:
raise HTTPException(status_code=409, detail="请先生成并保存研究计划")
if not selected_ids:
raise HTTPException(status_code=409, detail="请至少选择一条资料后再撰写报告")
sources_by_id = {
str(item.get("chunk_id") or ""): item
for item in task.get("sources") or []
if isinstance(item, dict)
}
sources = [sources_by_id[source_id] for source_id in selected_ids if source_id in sources_by_id]
if not sources:
raise HTTPException(status_code=409, detail="已选资料不可用,请重新保存研究计划")
if body.action in {"rewrite", "followup"} and not task.get("report_markdown"):
raise HTTPException(status_code=409, detail="当前任务还没有可供改写或追问的报告")
if body.action == "followup" and not body.instruction.strip():
raise HTTPException(status_code=422, detail="请填写报告追问")
row = await job_store.create_job(
{
"id": f"erj_{uuid4().hex}", "task_id": task_id, "user_id": user_id,
"model_name": (body.model_name or "").strip() or None,
"action": body.action, "instruction": body.instruction.strip(), "style": body.style,
"target_length": body.target_length, "plan_snapshot": plan, "sources": sources,
"original_report": task.get("report_markdown") if body.action in {"rewrite", "followup"} else None,
}
)
await task_store.update_task(task_id, user_id, {"active_report_job_id": row["id"]})
await executor.ensure_started(row["id"], user_id)
return _report_job_response(row)
@router.get("/tasks/{task_id}/report-jobs/{job_id}/stream")
async def stream_report_job(request: Request, task_id: str, job_id: str) -> StreamingResponse:
"""SSE stream: durable draft snapshot first, then live Markdown/reasoning deltas."""
user_id, _ = await _actor(request)
job = await _report_job_store(request).get_job(job_id, user_id)
if job is None or job.get("task_id") != task_id:
raise HTTPException(status_code=404, detail="报告任务不存在")
if job["status"] == "queued":
await _report_executor(request).ensure_started(job_id, user_id)
async def body():
async for event in _report_executor(request).stream_events(job_id, user_id):
yield "event: enterprise_research\n"
yield "data: " + json.dumps(event, ensure_ascii=False, separators=(",", ":")) + "\n\n"
return StreamingResponse(
body(), media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
@router.post(
"/tasks/{task_id}/report-jobs/{job_id}/cancel",
response_model=EnterpriseResearchReportJobResponse,
)
async def cancel_report_job(request: Request, task_id: str, job_id: str) -> EnterpriseResearchReportJobResponse:
user_id, _ = await _actor(request)
job = await _report_job_store(request).get_job(job_id, user_id)
if job is None or job.get("task_id") != task_id:
raise HTTPException(status_code=404, detail="报告任务不存在")
cancelled = await _report_executor(request).cancel(job_id, user_id)
return _report_job_response(cancelled or job)
@router.post(
"/tasks/{task_id}/report-jobs/{job_id}/restore",
response_model=EnterpriseResearchTaskResponse,
)
async def restore_report_version(request: Request, task_id: str, job_id: str) -> EnterpriseResearchTaskResponse:
"""Restore the before-image frozen by a completed full-report rewrite."""
user_id, _ = await _actor(request)
job = await _report_job_store(request).get_job(job_id, user_id)
if job is None or job.get("task_id") != task_id:
raise HTTPException(status_code=404, detail="报告任务不存在")
original = str(job.get("original_report") or "")
if job.get("status") != "completed" or job.get("action") != "rewrite" or not original.strip():
raise HTTPException(status_code=409, detail="该任务没有可恢复的改写前版本")
if await _report_job_store(request).get_active_for_task(task_id, user_id) is not None:
raise HTTPException(status_code=409, detail="报告任务正在执行,请完成或停止后再恢复版本")
task_store = _task_store(request)
current = await task_store.get_task(task_id, user_id)
if current is None:
raise HTTPException(status_code=404, detail="研究任务不存在")
current_report = str(current.get("report_markdown") or "")
if not current_report.strip():
raise HTTPException(status_code=409, detail="当前没有可替换的报告")
# The restore itself is a versioned operation: its before-image allows the
# user to reverse a mistaken restore through the same version UI.
await _report_job_store(request).create_job(
{
"id": f"erj_{uuid4().hex}", "task_id": task_id, "user_id": user_id,
"status": "completed", "phase": "completed", "action": "rewrite",
"instruction": f"恢复报告任务 {job_id} 的改写前版本", "style": "formal",
"model_name": current.get("report_model_name"),
"plan_snapshot": str(current.get("plan_markdown") or ""),
"sources": [
item for item in current.get("sources") or []
if isinstance(item, dict) and item.get("chunk_id") in set(current.get("selected_source_ids") or [])
],
"original_report": current_report, "report_markdown": original,
"report_summary": f"已恢复到报告任务 {job_id} 的改写前版本。",
}
)
task = await task_store.update_task(
task_id,
user_id,
{
"report_markdown": original,
"report_summary": f"已恢复到报告任务 {job_id} 的改写前版本。",
"active_report_job_id": None,
},
)
assert task is not None
return _response(task)
@router.delete("/tasks/{task_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_task(request: Request, task_id: str) -> None:
user_id, _ = await _actor(request)
if not await _task_store(request).delete_task(task_id, user_id):
raise HTTPException(status_code=404, detail="研究任务不存在")