711 lines
32 KiB
Python
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="研究任务不存在")
|