1108 lines
44 KiB
Python
1108 lines
44 KiB
Python
"""AI 写作路由 —— 业务表 CRUD + transcript 持久化 + admin cleanup。
|
||
|
||
写作主流程(``/start /stream /resume /state``)已迁到 LangGraph Server 的
|
||
``ai_writing`` assistant(前端走 ``useAIWritingStream`` hook + SDK),路由层
|
||
不再承载 graph 推进,只保留以下「不依赖 graph 状态」的辅助接口:
|
||
|
||
GET /api/ai-writing/sessions 列出用户的会话历史
|
||
GET /api/ai-writing/sessions/{id} 获取单个历史会话
|
||
PATCH /api/ai-writing/sessions/{id} 重命名会话标题
|
||
DELETE /api/ai-writing/sessions/{id} 删除会话
|
||
PUT /api/ai-writing/sessions/{id}/transcript 保存完整对话时间线
|
||
GET /api/ai-writing/article-types 列出所有文章类型
|
||
POST /api/ai-writing/article-types 新增文章类型
|
||
PATCH /api/ai-writing/article-types/{id} 更新文章类型
|
||
DELETE /api/ai-writing/article-types/{id} 删除文章类型
|
||
POST /api/ai-writing/sample/extract 提取样文仿写上传文件文本
|
||
POST /api/ai-writing/admin/cleanup (admin)手动清理会话
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
import tempfile
|
||
from pathlib import Path
|
||
from typing import Optional
|
||
|
||
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile
|
||
from fastapi.responses import StreamingResponse
|
||
from pydantic import BaseModel, Field
|
||
|
||
from app.gateway.authz import require_auth
|
||
from app.gateway.deps import get_config
|
||
from deerflow.utils.file_conversion import CONVERTIBLE_EXTENSIONS, convert_file_to_markdown
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
router = APIRouter(prefix="/api/ai-writing", tags=["ai-writing"])
|
||
|
||
_SAMPLE_TEXT_EXTENSIONS = {".txt", ".md", ".markdown"}
|
||
_MAX_SAMPLE_BYTES = 20 * 1024 * 1024
|
||
_MAX_SAMPLE_TEXT_CHARS = 120_000
|
||
|
||
# ── 请求/响应模型 ─────────────────────────────────────────────────────────────
|
||
|
||
class RenameSessionRequest(BaseModel):
|
||
title: str
|
||
|
||
|
||
class UpdateDraftRequest(BaseModel):
|
||
# 历史回看时用户在右侧编辑器里修改后的正文 Markdown
|
||
draft_markdown: str
|
||
# 可选:标题一并更新(用户改了文章标题时)
|
||
draft_title: Optional[str] = None
|
||
|
||
|
||
class AdminCleanupRequest(BaseModel):
|
||
"""手动清理接口的可选覆盖参数。
|
||
|
||
全部留空 → 用 config.yaml 的 ai_writing.cleanup 段默认值;显式传值
|
||
→ 覆盖配置(一般用于运维想验证一个临时阈值,例如清掉「最近 1 小时
|
||
以前的」会话来模拟整轮调度)。
|
||
"""
|
||
retention_days: Optional[int] = Field(default=None, ge=0, description="保留多少天以内的会话;0 = 不保留")
|
||
only_finished: Optional[bool] = Field(default=None, description="True 时只清 done/error 状态")
|
||
batch_size: Optional[int] = Field(default=None, ge=1, le=10000, description="单批最多删多少行")
|
||
|
||
|
||
class AdminCleanupResponse(BaseModel):
|
||
deleted_db: int
|
||
deleted_checkpoints: int
|
||
# 因「仍在跑 / 等用户干预」被保护掉的会话数(threads_meta.status in
|
||
# running/interrupted)。即使 retention_days=0 也不会秒杀用户正在写的稿。
|
||
skipped_active: int = 0
|
||
elapsed_seconds: float
|
||
params: dict
|
||
|
||
|
||
class SaveTranscriptRequest(BaseModel):
|
||
# 前端组装的完整对话时间线(progressEvents + completedInterventions)
|
||
transcript: dict
|
||
|
||
|
||
class SampleExtractResponse(BaseModel):
|
||
filename: str
|
||
text: str
|
||
char_count: int
|
||
truncated: bool = False
|
||
|
||
|
||
class ArticleTypeRequest(BaseModel):
|
||
key: str
|
||
label: str
|
||
description: Optional[str] = None
|
||
prompt_hint: Optional[str] = None
|
||
default_word_count: int = 800
|
||
sort_order: int = 0
|
||
|
||
|
||
class UpdateArticleTypeRequest(BaseModel):
|
||
key: Optional[str] = None
|
||
label: Optional[str] = None
|
||
description: Optional[str] = None
|
||
prompt_hint: Optional[str] = None
|
||
default_word_count: Optional[int] = None
|
||
sort_order: Optional[int] = None
|
||
|
||
|
||
class ArticleTypeResponse(BaseModel):
|
||
id: str
|
||
key: str
|
||
label: str
|
||
description: Optional[str] = None
|
||
prompt_hint: Optional[str] = None
|
||
default_word_count: int
|
||
sort_order: int
|
||
created_at: str
|
||
updated_at: str
|
||
|
||
|
||
class ResearchTemplateRequest(BaseModel):
|
||
name: str
|
||
description: Optional[str] = None
|
||
article_type_key: Optional[str] = None
|
||
outline: Optional[str] = None
|
||
material_source: str = "general"
|
||
word_count: int = 800
|
||
audience: Optional[str] = None
|
||
strict_mode: bool = False
|
||
topic: Optional[str] = None
|
||
sort_order: int = 0
|
||
|
||
|
||
class UpdateResearchTemplateRequest(BaseModel):
|
||
name: Optional[str] = None
|
||
description: Optional[str] = None
|
||
article_type_key: Optional[str] = None
|
||
outline: Optional[str] = None
|
||
material_source: Optional[str] = None
|
||
word_count: Optional[int] = None
|
||
audience: Optional[str] = None
|
||
strict_mode: Optional[bool] = None
|
||
topic: Optional[str] = None
|
||
sort_order: Optional[int] = None
|
||
|
||
|
||
class ResearchTemplateResponse(BaseModel):
|
||
id: str
|
||
name: str
|
||
description: Optional[str] = None
|
||
article_type_key: Optional[str] = None
|
||
outline: Optional[str] = None
|
||
material_source: str
|
||
word_count: int
|
||
audience: Optional[str] = None
|
||
strict_mode: bool
|
||
topic: Optional[str] = None
|
||
sort_order: int
|
||
created_at: str
|
||
updated_at: str
|
||
|
||
|
||
class SessionResponse(BaseModel):
|
||
id: str
|
||
user_id: Optional[str] = None
|
||
title: str
|
||
user_intent: Optional[str] = None
|
||
status: str
|
||
draft_title: Optional[str] = None
|
||
draft_markdown: Optional[str] = None
|
||
completed_interventions: Optional[list] = None
|
||
review_result: Optional[dict] = None
|
||
transcript: Optional[dict] = None
|
||
created_at: str
|
||
updated_at: str
|
||
|
||
|
||
async def _extract_sample_text(file: UploadFile) -> tuple[str, str]:
|
||
name = file.filename or "sample"
|
||
ext = Path(name).suffix.lower()
|
||
raw = await file.read()
|
||
if not raw:
|
||
raise HTTPException(status_code=400, detail="文件为空")
|
||
if len(raw) > _MAX_SAMPLE_BYTES:
|
||
raise HTTPException(status_code=413, detail=f"文件过大,最大支持 {_MAX_SAMPLE_BYTES // (1024 * 1024)}MB")
|
||
|
||
if ext in _SAMPLE_TEXT_EXTENSIONS:
|
||
text = raw.decode("utf-8", errors="ignore")
|
||
elif ext in CONVERTIBLE_EXTENSIONS:
|
||
with tempfile.TemporaryDirectory() as tmp:
|
||
tmp_path = Path(tmp) / f"sample{ext}"
|
||
tmp_path.write_bytes(raw)
|
||
md_path = await convert_file_to_markdown(tmp_path)
|
||
if md_path is None:
|
||
raise HTTPException(status_code=422, detail=f"无法解析文件:{name}")
|
||
text = md_path.read_text(encoding="utf-8", errors="ignore")
|
||
else:
|
||
raise HTTPException(status_code=415, detail=f"不支持的文件类型:{ext or '未知'}(支持 doc/docx/pdf/md/txt)")
|
||
|
||
text = text.strip()
|
||
if not text:
|
||
raise HTTPException(status_code=422, detail="文件内容为空或无法提取文本")
|
||
return name, text
|
||
|
||
|
||
@router.post("/sample/extract", response_model=SampleExtractResponse)
|
||
@require_auth
|
||
async def extract_sample_article(request: Request, file: UploadFile = File(...)) -> SampleExtractResponse:
|
||
"""Extract text from a sample article for AI-writing imitation mode."""
|
||
_ = request
|
||
name, text = await _extract_sample_text(file)
|
||
char_count = len(text)
|
||
truncated = char_count > _MAX_SAMPLE_TEXT_CHARS
|
||
if truncated:
|
||
text = text[:_MAX_SAMPLE_TEXT_CHARS]
|
||
return SampleExtractResponse(
|
||
filename=name,
|
||
text=text,
|
||
char_count=char_count,
|
||
truncated=truncated,
|
||
)
|
||
|
||
|
||
# ── 路由 ──────────────────────────────────────────────────────────────────────
|
||
@router.get("/sessions", response_model=list[SessionResponse])
|
||
async def list_sessions(request: Request):
|
||
"""列出当前用户的 AI 写作历史会话。"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
from app.gateway.routers._ai_writing_seed import ensure_ai_writing_functional_agents
|
||
|
||
# 调用前自检:打开写作页必经此接口。若管理员在启动后手删了内置写作 agent
|
||
# 目录,这里自愈补建,保证 use_builtin_agents 开启时图节点能加载到 SOUL。
|
||
ensure_ai_writing_functional_agents()
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
return []
|
||
user_id = await get_current_user(request)
|
||
return await repo.list(user_id=user_id)
|
||
|
||
|
||
@router.get("/sessions/{session_id}", response_model=SessionResponse)
|
||
async def get_session(session_id: str, request: Request):
|
||
"""获取单个历史会话。"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
record = await repo.get(session_id, user_id=user_id)
|
||
if record is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
return record
|
||
|
||
|
||
@router.patch("/sessions/{session_id}", response_model=SessionResponse)
|
||
async def rename_session(session_id: str, body: RenameSessionRequest, request: Request):
|
||
"""重命名会话标题。"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
updated = await repo.update(session_id, user_id=user_id, title=body.title.strip())
|
||
if updated is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
return updated
|
||
|
||
|
||
@router.put("/sessions/{session_id}/draft", response_model=SessionResponse)
|
||
async def update_session_draft(session_id: str, body: UpdateDraftRequest, request: Request):
|
||
"""更新会话的成稿正文(历史回看时用户在右侧编辑器里直接修改并保存)。"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
# draft_title 缺省(None)时不动原标题,传了才更新
|
||
extra = {} if body.draft_title is None else {"draft_title": body.draft_title}
|
||
updated = await repo.update(
|
||
session_id,
|
||
user_id=user_id,
|
||
draft_markdown=body.draft_markdown,
|
||
**extra,
|
||
)
|
||
if updated is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
return updated
|
||
|
||
|
||
@router.delete("/sessions/{session_id}")
|
||
async def delete_session(session_id: str, request: Request):
|
||
"""删除历史会话。"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
deleted = await repo.delete(session_id, user_id=user_id)
|
||
if not deleted:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
return {"success": True}
|
||
|
||
|
||
@router.put("/sessions/{session_id}/transcript")
|
||
async def save_transcript(session_id: str, body: SaveTranscriptRequest, request: Request):
|
||
"""保存会话的完整对话时间线,供历史会话回显出与实时对话一致的富时间线。
|
||
|
||
时间线持久化属于「锦上添花」功能:即便保存失败(如内网部署缺少 transcript
|
||
列、或 MySQL TEXT 列超长),也绝不能阻断写作流程。任何异常都吞掉并返回
|
||
success=False,由前端忽略即可。
|
||
"""
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
return {"success": False, "reason": "store_unavailable"}
|
||
try:
|
||
user_id = await get_current_user(request)
|
||
updated = await repo.update(session_id, user_id=user_id, transcript=body.transcript)
|
||
except Exception:
|
||
logger.exception("ai_writing: failed to save transcript: session=%s", session_id)
|
||
return {"success": False, "reason": "save_failed"}
|
||
if updated is None:
|
||
# 会话不存在或不属于当前用户:同样不抛 500,返回 success=False 即可
|
||
return {"success": False, "reason": "session_not_found"}
|
||
return {"success": True}
|
||
|
||
|
||
# ── 意图解析(对话驱动写作台) ────────────────────────────────────────────────
|
||
|
||
|
||
class IntentParseRequest(BaseModel):
|
||
# 用户在底部智能输入框发的自由文本
|
||
text: str
|
||
# 前端当前 WritingStatus(仅参考;暂停点以 checkpoint 里真实 interrupt 为准)
|
||
status: Optional[str] = None
|
||
# 前端认为的当前暂停点(checkpoint 读取失败时的兜底)
|
||
pause_point: Optional[str] = None
|
||
|
||
|
||
class IntentParseResponse(BaseModel):
|
||
intent: str
|
||
pause_point: Optional[str] = None
|
||
action: Optional[str] = None
|
||
# 驼峰键(approvedMaterialIds / userQuery / outlineFeedback / userRevisionNotes /
|
||
# overrideVerdict / selectedIssueIndices / sectionDecisions),前端可直接喂 submitIntervention
|
||
payload: dict = Field(default_factory=dict)
|
||
need_research: bool = False
|
||
research_query: str = ""
|
||
answer: str = ""
|
||
confidence: float = 0.0
|
||
clarify: str = ""
|
||
|
||
|
||
class WritingSuggestionMessage(BaseModel):
|
||
role: str
|
||
content: str
|
||
|
||
|
||
class WritingSuggestionsRequest(BaseModel):
|
||
messages: list[WritingSuggestionMessage] = Field(default_factory=list)
|
||
n: int = Field(default=3, ge=1, le=5)
|
||
model_name: Optional[str] = None
|
||
|
||
|
||
class WritingSuggestionsResponse(BaseModel):
|
||
suggestions: list[str] = Field(default_factory=list)
|
||
|
||
|
||
class AnswerStreamRequest(BaseModel):
|
||
text: str
|
||
status: Optional[str] = None
|
||
pause_point: Optional[str] = None
|
||
model_name: Optional[str] = None
|
||
|
||
|
||
async def _prepare_intent_parse_context(
|
||
session_id: str,
|
||
body: IntentParseRequest,
|
||
request: Request,
|
||
) -> tuple[str, dict, Optional[dict], Optional[str]]:
|
||
"""鉴权并读取会话 checkpoint,为一次性/流式意图解析提供同一份上下文。"""
|
||
text = (body.text or "").strip()
|
||
if not text:
|
||
raise HTTPException(status_code=400, detail="text 不能为空")
|
||
|
||
await _ensure_session_owner(session_id, request)
|
||
values, interrupt_value, pause_point = await _read_ai_writing_checkpoint_context(
|
||
session_id,
|
||
body.pause_point,
|
||
request,
|
||
)
|
||
return text, values, interrupt_value, pause_point
|
||
|
||
|
||
async def _read_ai_writing_checkpoint_context(
|
||
session_id: str,
|
||
fallback_pause_point: Optional[str],
|
||
request: Request,
|
||
) -> tuple[dict, Optional[dict], Optional[str]]:
|
||
from app.gateway.ai_writing_job_executor import _pending_interrupt
|
||
|
||
# 读 checkpoint:暂停点以 pending interrupt 为权威(frontend 声明只作读取失败兜底)。
|
||
values: dict = {}
|
||
interrupt_value: Optional[dict] = None
|
||
pause_point: Optional[str] = None
|
||
checkpoint_ok = False
|
||
_app_state = getattr(getattr(request, "app", None), "state", None)
|
||
checkpointer = getattr(_app_state, "checkpointer", None)
|
||
if checkpointer is not None:
|
||
from deerflow.agents.ai_writing.graph import make_ai_writing_graph
|
||
|
||
graph = make_ai_writing_graph()
|
||
graph.checkpointer = checkpointer
|
||
try:
|
||
snapshot = await graph.aget_state({"configurable": {"thread_id": session_id}})
|
||
values = getattr(snapshot, "values", None) or {}
|
||
interrupt_value = _pending_interrupt(snapshot)
|
||
checkpoint_ok = bool(values)
|
||
if interrupt_value:
|
||
pause_point = interrupt_value.get("pause_point")
|
||
except Exception:
|
||
logger.warning("ai_writing intent: read checkpoint failed for %s", session_id, exc_info=True)
|
||
if pause_point is None and not checkpoint_ok:
|
||
pause_point = (fallback_pause_point or "").strip() or None
|
||
|
||
return values, interrupt_value, pause_point
|
||
|
||
|
||
async def _get_session_owner_record(session_id: str, request: Request) -> dict:
|
||
from app.gateway.deps import get_ai_writing_session_store, get_current_user
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
record = await repo.get(session_id, user_id=user_id)
|
||
if record is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
return record
|
||
|
||
|
||
async def _ensure_session_owner(session_id: str, request: Request) -> None:
|
||
await _get_session_owner_record(session_id, request)
|
||
|
||
|
||
def _clip_answer_context(text: str, limit: int) -> str:
|
||
text = (text or "").strip()
|
||
if len(text) <= limit:
|
||
return text
|
||
return text[:limit] + "\n…(内容过长已截断)"
|
||
|
||
|
||
def _build_session_answer_context(record: dict) -> str:
|
||
"""Build a fast context from the persisted session row for plain Q&A."""
|
||
parts: list[str] = []
|
||
title = str(record.get("title") or record.get("draft_title") or "").strip()
|
||
user_intent = str(record.get("user_intent") or "").strip()
|
||
status = str(record.get("status") or "").strip()
|
||
if title:
|
||
parts.append(f"会话标题:{title}")
|
||
if user_intent:
|
||
parts.append(f"写作主题:{_clip_answer_context(user_intent, 300)}")
|
||
if status:
|
||
parts.append(f"流程状态:{status}")
|
||
|
||
draft_markdown = str(record.get("draft_markdown") or "").strip()
|
||
has_substantive_context = bool(draft_markdown)
|
||
if draft_markdown:
|
||
draft_title = str(record.get("draft_title") or title or "").strip()
|
||
header = f"草稿标题:《{draft_title}》" if draft_title else "当前草稿"
|
||
parts.append(f"{header}\n--- 草稿正文 ---\n{_clip_answer_context(draft_markdown, 6000)}")
|
||
|
||
review = record.get("review_result") if isinstance(record.get("review_result"), dict) else {}
|
||
if review:
|
||
has_substantive_context = True
|
||
lines = [
|
||
f"编辑结论:{review.get('verdict', '')}(总分 {review.get('overall_score', review.get('overallScore', ''))})"
|
||
]
|
||
notes = review.get("revision_notes") or review.get("revisionNotes") or []
|
||
if isinstance(notes, list):
|
||
for i, note in enumerate(notes[:8]):
|
||
if not isinstance(note, dict):
|
||
continue
|
||
issue = _clip_answer_context(str(note.get("issue") or ""), 120)
|
||
suggestion = _clip_answer_context(str(note.get("suggestion") or ""), 120)
|
||
if issue or suggestion:
|
||
lines.append(f"意见 {i + 1}:{issue} → {suggestion}")
|
||
parts.append("\n".join(lines))
|
||
|
||
if not has_substantive_context:
|
||
return ""
|
||
return "\n\n".join(part for part in parts if part.strip())
|
||
|
||
|
||
@router.post("/sessions/{session_id}/intent", response_model=IntentParseResponse)
|
||
async def parse_session_intent(session_id: str, body: IntentParseRequest, request: Request):
|
||
"""对话驱动写作台:把用户自由文本翻译成结构化写作指令(或问答/澄清)。
|
||
|
||
读该会话 graph checkpoint 的素材/大纲/草稿/审核摘要 + 当前 pending interrupt
|
||
(权威暂停点),交给 ``deerflow.agents.ai_writing.intent_router`` 的 LLM 意图
|
||
路由。动作白名单 / 高代价置信度门槛 / payload 清洗都在 intent_router 内完成,
|
||
这里只做鉴权与 state 读取。模型失败不抛 5xx —— intent_router 会降级成 clarify。
|
||
"""
|
||
from deerflow.agents.ai_writing.intent_router import parse_user_intent
|
||
|
||
text, values, interrupt_value, pause_point = await _prepare_intent_parse_context(session_id, body, request)
|
||
|
||
result = await parse_user_intent(
|
||
text=text,
|
||
pause_point=pause_point,
|
||
values=values,
|
||
interrupt_value=interrupt_value,
|
||
model_name=str(values.get("model_name") or "") or None,
|
||
)
|
||
return IntentParseResponse(**result)
|
||
|
||
|
||
@router.post("/sessions/{session_id}/intent/stream")
|
||
async def stream_session_intent(session_id: str, body: IntentParseRequest, request: Request):
|
||
"""对话驱动写作台:流式意图解析。
|
||
|
||
SSE 事件:
|
||
- ``thinking``: ``{"chunk": "..."}``,模型推理 / ``<think>`` 内容
|
||
- ``text``: ``{"chunk": "..."}``,模型正文输出(通常是结构化 JSON)
|
||
- ``result``: 最终规范化后的 ``IntentParseResponse`` 字段
|
||
- ``done`` / ``error``
|
||
"""
|
||
from app.gateway.services import format_sse
|
||
from deerflow.agents.ai_writing.intent_router import stream_user_intent
|
||
|
||
text, values, interrupt_value, pause_point = await _prepare_intent_parse_context(session_id, body, request)
|
||
|
||
async def event_stream():
|
||
try:
|
||
async for event in stream_user_intent(
|
||
text=text,
|
||
pause_point=pause_point,
|
||
values=values,
|
||
interrupt_value=interrupt_value,
|
||
model_name=str(values.get("model_name") or "") or None,
|
||
):
|
||
if event.get("type") == "chunk":
|
||
chunk = str(event.get("chunk") or "")
|
||
if not chunk:
|
||
continue
|
||
yield format_sse(
|
||
"thinking" if event.get("is_thinking") else "text",
|
||
{"chunk": chunk},
|
||
)
|
||
elif event.get("type") == "result":
|
||
yield format_sse("result", event.get("result") or {})
|
||
yield format_sse("done", {})
|
||
except Exception:
|
||
logger.exception("ai_writing intent stream failed: session=%s", session_id)
|
||
yield format_sse("error", {"message": "意图解析流式接口异常,请稍后重试"})
|
||
|
||
return StreamingResponse(
|
||
event_stream(),
|
||
media_type="text/event-stream",
|
||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||
)
|
||
|
||
|
||
@router.post("/sessions/{session_id}/suggestions", response_model=WritingSuggestionsResponse)
|
||
async def generate_session_suggestions(
|
||
session_id: str,
|
||
body: WritingSuggestionsRequest,
|
||
request: Request,
|
||
config=Depends(get_config),
|
||
):
|
||
"""为 AI 写作对话版生成窄栏追问建议。
|
||
|
||
与主对话页的 ``/api/threads/{id}/suggestions`` 语义一致,但鉴权对象换成
|
||
AI 写作 session,避免把写作 sessionId 冒充 workspace threadId。
|
||
"""
|
||
from langchain_core.messages import HumanMessage, SystemMessage
|
||
|
||
from app.gateway.routers.suggestions import (
|
||
_extract_response_text,
|
||
_format_conversation,
|
||
_parse_json_string_list,
|
||
SuggestionMessage,
|
||
)
|
||
from deerflow.models import create_chat_model
|
||
|
||
await _ensure_session_owner(session_id, request)
|
||
|
||
if not body.messages:
|
||
return WritingSuggestionsResponse()
|
||
|
||
messages = [
|
||
SuggestionMessage(role=m.role, content=m.content)
|
||
for m in body.messages
|
||
if (m.content or "").strip()
|
||
]
|
||
conversation = _format_conversation(messages)
|
||
if not conversation:
|
||
return WritingSuggestionsResponse()
|
||
|
||
n = body.n
|
||
system_instruction = (
|
||
"You are generating compact follow-up questions for a narrow AI writing sidebar.\n"
|
||
f"Based on the conversation below, produce EXACTLY {n} short questions the user might ask next.\n"
|
||
"Requirements:\n"
|
||
"- Questions must be relevant to the article or writing workflow.\n"
|
||
"- Prefer questions about summary, structure, evidence, style, revisions, or next actions.\n"
|
||
"- Questions must be written in the same language as the user.\n"
|
||
"- Keep each question very concise (<= 24 Chinese characters when Chinese).\n"
|
||
"- Do NOT include numbering, markdown, or any extra text.\n"
|
||
"- Output MUST be a JSON array of strings only.\n"
|
||
)
|
||
user_content = f"Conversation Context:\n{conversation}\n\nGenerate {n} follow-up questions"
|
||
|
||
try:
|
||
model = create_chat_model(name=body.model_name, thinking_enabled=False, app_config=config)
|
||
response = await model.ainvoke(
|
||
[SystemMessage(content=system_instruction), HumanMessage(content=user_content)],
|
||
config={"run_name": "ai_writing_suggestions"},
|
||
)
|
||
raw = _extract_response_text(response.content)
|
||
suggestions = _parse_json_string_list(raw) or []
|
||
cleaned = [s.replace("\n", " ").strip() for s in suggestions if s.strip()]
|
||
return WritingSuggestionsResponse(suggestions=cleaned[:n])
|
||
except Exception:
|
||
logger.exception("ai_writing: failed to generate suggestions: session=%s", session_id)
|
||
return WritingSuggestionsResponse()
|
||
|
||
|
||
@router.post("/sessions/{session_id}/answer/stream")
|
||
async def stream_session_answer(session_id: str, body: AnswerStreamRequest, request: Request):
|
||
"""AI 写作对话版问答旁路:纯文本 SSE 回答。
|
||
|
||
意图路由仍负责操作类消息;明显的提问走这里,避免把答案塞进 JSON 的
|
||
``answer`` 字段导致“模型完整生成后才一次性显示”。
|
||
"""
|
||
from app.gateway.services import format_sse
|
||
from deerflow.agents.ai_writing.intent_router import build_context_snapshot
|
||
from deerflow.agents.ai_writing.nodes._utils import ThinkTagParser, _content_to_text, reasoning_delta
|
||
from deerflow.models import create_chat_model
|
||
from langchain_core.messages import HumanMessage, SystemMessage
|
||
|
||
text = (body.text or "").strip()
|
||
if not text:
|
||
raise HTTPException(status_code=400, detail="text 不能为空")
|
||
record = await _get_session_owner_record(session_id, request)
|
||
session_context = _build_session_answer_context(record)
|
||
|
||
async def event_stream():
|
||
import time
|
||
|
||
started = time.monotonic()
|
||
parser = ThinkTagParser()
|
||
chunk_count = 0
|
||
first_chunk_logged = False
|
||
yield format_sse("start", {})
|
||
try:
|
||
values: dict = {}
|
||
if session_context:
|
||
context = session_context
|
||
model_name = (body.model_name or "").strip() or None
|
||
logger.info(
|
||
"ai_writing answer context ready: session=%s elapsed=%.3fs source=session_row pause_point=%s",
|
||
session_id,
|
||
time.monotonic() - started,
|
||
body.pause_point,
|
||
)
|
||
else:
|
||
values, interrupt_value, pause_point = await _read_ai_writing_checkpoint_context(
|
||
session_id,
|
||
body.pause_point,
|
||
request,
|
||
)
|
||
context = build_context_snapshot(values, pause_point, interrupt_value)
|
||
model_name = (body.model_name or "").strip() or str(values.get("model_name") or "") or None
|
||
logger.info(
|
||
"ai_writing answer context ready: session=%s elapsed=%.3fs source=checkpoint has_values=%s pause_point=%s",
|
||
session_id,
|
||
time.monotonic() - started,
|
||
bool(values),
|
||
pause_point,
|
||
)
|
||
|
||
messages = [
|
||
SystemMessage(content=(
|
||
"你是 AI 智能写作工作台里的问答助手。用户正在写一篇文章,"
|
||
"你只能根据给定的写作上下文回答问题,不执行确认、检索、修改、定稿等流程动作。"
|
||
"如果上下文没有足够信息,请明确说明,不要编造。回答用中文,简洁但有信息量。"
|
||
)),
|
||
HumanMessage(content=f"【写作上下文】\n{context}\n\n【用户问题】\n{text}\n\n请直接回答:"),
|
||
]
|
||
|
||
model = create_chat_model(
|
||
name=model_name,
|
||
thinking_enabled=False,
|
||
force_disable_thinking=True,
|
||
)
|
||
async for chunk in model.astream(messages, config={"run_name": "ai_writing_answer"}):
|
||
reasoning = reasoning_delta(chunk)
|
||
if reasoning:
|
||
if not first_chunk_logged:
|
||
first_chunk_logged = True
|
||
logger.info(
|
||
"ai_writing answer first chunk: session=%s elapsed=%.3fs kind=thinking",
|
||
session_id,
|
||
time.monotonic() - started,
|
||
)
|
||
yield format_sse("thinking", {"chunk": reasoning})
|
||
content = _content_to_text(getattr(chunk, "content", "") or "")
|
||
if not content:
|
||
continue
|
||
for is_thinking, part in parser.feed(content):
|
||
if not part:
|
||
continue
|
||
if not first_chunk_logged:
|
||
first_chunk_logged = True
|
||
logger.info(
|
||
"ai_writing answer first chunk: session=%s elapsed=%.3fs kind=%s",
|
||
session_id,
|
||
time.monotonic() - started,
|
||
"thinking" if is_thinking else "text",
|
||
)
|
||
chunk_count += 1
|
||
yield format_sse("thinking" if is_thinking else "text", {"chunk": part})
|
||
logger.info(
|
||
"ai_writing answer stream done: session=%s chunks=%d elapsed=%.3fs first_chunk=%s",
|
||
session_id,
|
||
chunk_count,
|
||
time.monotonic() - started,
|
||
first_chunk_logged,
|
||
)
|
||
yield format_sse("done", {})
|
||
except Exception:
|
||
logger.exception("ai_writing answer stream failed: session=%s", session_id)
|
||
yield format_sse("error", {"message": "回答生成失败,请稍后重试"})
|
||
|
||
return StreamingResponse(
|
||
event_stream(),
|
||
media_type="text/event-stream",
|
||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||
)
|
||
|
||
|
||
class BackgroundResponse(BaseModel):
|
||
# 是否新拉起了后台任务(False = 已在后台跑,幂等返回)
|
||
started: bool
|
||
# 会话当前状态(background = 挂起中)
|
||
status: str
|
||
|
||
|
||
@router.post("/sessions/{session_id}/background", response_model=BackgroundResponse)
|
||
async def suspend_to_background(session_id: str, request: Request):
|
||
"""把这篇写作转入「后台挂起」:后端自动驱动剩余流程到定稿。
|
||
|
||
前端在用户点「后台挂起」或离开页面时调用。幂等:同一会话已在后台跑则直接返回。
|
||
驱动逻辑见 ``app.gateway.ai_writing_job_executor``(素材/大纲自动确认、缺素材
|
||
通用知识续写、草稿送审、审核不过自动重写直到通过或用完次数再定稿)。
|
||
"""
|
||
from app.gateway.deps import (
|
||
get_ai_writing_job_executor,
|
||
get_ai_writing_session_store,
|
||
get_current_user,
|
||
)
|
||
|
||
executor = get_ai_writing_job_executor(request)
|
||
if executor is None:
|
||
raise HTTPException(status_code=503, detail="Background executor not available")
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
|
||
user_id = await get_current_user(request)
|
||
# 鉴权 + 存在性:会话必须属于当前用户(get 已带 user_id 过滤)。
|
||
record = await repo.get(session_id, user_id=user_id)
|
||
if record is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
if record.get("status") == "done":
|
||
# 已定稿,无需后台驱动。
|
||
return BackgroundResponse(started=False, status="done")
|
||
|
||
started = executor.start(thread_id=session_id, user_id=user_id)
|
||
return BackgroundResponse(started=started, status="background")
|
||
|
||
|
||
class ProgressResponse(BaseModel):
|
||
# 业务表 status(background = 挂起中 / done / error / in_progress)
|
||
sessionStatus: str
|
||
# 规范化阶段 key(intent/research/material_confirm/outline/outline_confirm/draft/draft_confirm/review/done),
|
||
# 从 graph checkpoint 推导;None = 还没建 checkpoint。
|
||
stage: Optional[str] = None
|
||
# 阶段细化说明(素材不足求助 / 编辑打回)
|
||
note: Optional[str] = None
|
||
revisionCount: int = 0
|
||
maxRevisions: int = 3
|
||
reviewVerdict: Optional[str] = None
|
||
# 后台驱动任务是否在本进程在跑(重启后可能为 False 但 status 仍 background,由 reconcile 续起)
|
||
isRunning: bool = False
|
||
|
||
|
||
@router.get("/sessions/{session_id}/progress", response_model=ProgressResponse)
|
||
async def get_session_progress(session_id: str, request: Request):
|
||
"""读取一篇写作的实际进度(当前阶段),供「查看进度」流程图弹窗轮询。
|
||
|
||
阶段从 graph checkpoint 推导,反映后端真实推进位置;鉴权要求会话属于当前用户。
|
||
"""
|
||
from app.gateway.deps import (
|
||
get_ai_writing_job_executor,
|
||
get_ai_writing_session_store,
|
||
get_current_user,
|
||
)
|
||
|
||
repo = get_ai_writing_session_store(request)
|
||
if repo is None:
|
||
raise HTTPException(status_code=503, detail="Session store not available")
|
||
user_id = await get_current_user(request)
|
||
record = await repo.get(session_id, user_id=user_id)
|
||
if record is None:
|
||
raise HTTPException(status_code=404, detail="Session not found")
|
||
|
||
executor = get_ai_writing_job_executor(request)
|
||
progress = await executor.get_progress(session_id) if executor is not None else {}
|
||
return ProgressResponse(
|
||
sessionStatus=str(record.get("status") or "in_progress"),
|
||
stage=progress.get("stage") or None,
|
||
note=progress.get("note"),
|
||
revisionCount=int(progress.get("revisionCount") or 0),
|
||
maxRevisions=int(progress.get("maxRevisions") or 3),
|
||
reviewVerdict=progress.get("reviewVerdict"),
|
||
isRunning=bool(executor.is_running(session_id)) if executor is not None else False,
|
||
)
|
||
|
||
|
||
@router.get("/article-types", response_model=list[ArticleTypeResponse])
|
||
async def list_article_types(request: Request):
|
||
"""列出所有文章类型,按 sort_order 升序排列。"""
|
||
from app.gateway.deps import get_article_type_store
|
||
from app.gateway.routers._ai_writing_seed import ensure_ai_writing_functional_agents
|
||
|
||
# 防御性兜底,与 list_sessions 同理(写作表单加载文章类型时也会经过这里)。
|
||
ensure_ai_writing_functional_agents()
|
||
|
||
store = get_article_type_store(request)
|
||
if store is None:
|
||
return []
|
||
return await store.list()
|
||
|
||
|
||
@router.post("/article-types", response_model=ArticleTypeResponse, status_code=201)
|
||
async def create_article_type(body: ArticleTypeRequest, request: Request):
|
||
"""新增文章类型。key 必须全局唯一。"""
|
||
from app.gateway.deps import get_article_type_store
|
||
|
||
store = get_article_type_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Article type store not available")
|
||
existing = await store.get_by_key(body.key)
|
||
if existing is not None:
|
||
raise HTTPException(status_code=409, detail=f"Article type key '{body.key}' already exists")
|
||
return await store.create(
|
||
key=body.key,
|
||
label=body.label,
|
||
description=body.description,
|
||
prompt_hint=body.prompt_hint,
|
||
default_word_count=body.default_word_count,
|
||
sort_order=body.sort_order,
|
||
)
|
||
|
||
|
||
@router.patch("/article-types/{article_type_id}", response_model=ArticleTypeResponse)
|
||
async def update_article_type(article_type_id: str, body: UpdateArticleTypeRequest, request: Request):
|
||
"""更新文章类型字段(只传需要修改的字段)。"""
|
||
from app.gateway.deps import get_article_type_store
|
||
|
||
store = get_article_type_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Article type store not available")
|
||
# 如果修改了 key,检查唯一性
|
||
if body.key is not None:
|
||
existing = await store.get_by_key(body.key)
|
||
if existing is not None and existing["id"] != article_type_id:
|
||
raise HTTPException(status_code=409, detail=f"Article type key '{body.key}' already exists")
|
||
|
||
_UNSET = object.__new__(object)
|
||
kwargs = {}
|
||
for field in ("key", "label", "description", "prompt_hint", "default_word_count", "sort_order"):
|
||
val = getattr(body, field)
|
||
if val is not None:
|
||
kwargs[field] = val
|
||
|
||
updated = await store.update(article_type_id, **kwargs)
|
||
if updated is None:
|
||
raise HTTPException(status_code=404, detail="Article type not found")
|
||
return updated
|
||
|
||
|
||
@router.delete("/article-types/{article_type_id}")
|
||
async def delete_article_type(article_type_id: str, request: Request):
|
||
"""删除文章类型。"""
|
||
from app.gateway.deps import get_article_type_store
|
||
|
||
store = get_article_type_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Article type store not available")
|
||
deleted = await store.delete(article_type_id)
|
||
if not deleted:
|
||
raise HTTPException(status_code=404, detail="Article type not found")
|
||
return {"success": True}
|
||
|
||
|
||
# ── 研究模板(课题研究左栏一键预填表单,需求 0616 #5)────────────────────────────
|
||
|
||
|
||
@router.get("/research-templates", response_model=list[ResearchTemplateResponse])
|
||
async def list_research_templates(request: Request):
|
||
"""列出所有研究模板,按 sort_order 升序。
|
||
|
||
新机器/空表时自动写入内置范例模板(网摘/周报/要讯),幂等仅在表为空时执行。
|
||
"""
|
||
from app.gateway.deps import get_research_template_store
|
||
|
||
store = get_research_template_store(request)
|
||
if store is None:
|
||
return []
|
||
try:
|
||
await store.seed_defaults()
|
||
except Exception:
|
||
logger.exception("ai_writing: failed to seed default research templates")
|
||
return await store.list()
|
||
|
||
|
||
@router.post("/research-templates", response_model=ResearchTemplateResponse, status_code=201)
|
||
async def create_research_template(body: ResearchTemplateRequest, request: Request):
|
||
"""新增研究模板。"""
|
||
from app.gateway.deps import get_research_template_store
|
||
|
||
store = get_research_template_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Research template store not available")
|
||
return await store.create(
|
||
name=body.name,
|
||
description=body.description,
|
||
article_type_key=body.article_type_key,
|
||
outline=body.outline,
|
||
material_source=body.material_source,
|
||
word_count=body.word_count,
|
||
audience=body.audience,
|
||
strict_mode=body.strict_mode,
|
||
topic=body.topic,
|
||
sort_order=body.sort_order,
|
||
)
|
||
|
||
|
||
@router.patch("/research-templates/{template_id}", response_model=ResearchTemplateResponse)
|
||
async def update_research_template(template_id: str, body: UpdateResearchTemplateRequest, request: Request):
|
||
"""更新研究模板字段(只传需要修改的字段)。"""
|
||
from app.gateway.deps import get_research_template_store
|
||
|
||
store = get_research_template_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Research template store not available")
|
||
|
||
kwargs = {}
|
||
for field in (
|
||
"name",
|
||
"description",
|
||
"article_type_key",
|
||
"outline",
|
||
"material_source",
|
||
"word_count",
|
||
"audience",
|
||
"strict_mode",
|
||
"topic",
|
||
"sort_order",
|
||
):
|
||
val = getattr(body, field)
|
||
if val is not None:
|
||
kwargs[field] = val
|
||
|
||
updated = await store.update(template_id, **kwargs)
|
||
if updated is None:
|
||
raise HTTPException(status_code=404, detail="Research template not found")
|
||
return updated
|
||
|
||
|
||
@router.delete("/research-templates/{template_id}")
|
||
async def delete_research_template(template_id: str, request: Request):
|
||
"""删除研究模板。"""
|
||
from app.gateway.deps import get_research_template_store
|
||
|
||
store = get_research_template_store(request)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Research template store not available")
|
||
deleted = await store.delete(template_id)
|
||
if not deleted:
|
||
raise HTTPException(status_code=404, detail="Research template not found")
|
||
return {"success": True}
|
||
|
||
|
||
# ── 管理接口 ──────────────────────────────────────────────────────────────────
|
||
|
||
async def _require_admin(request: Request) -> None:
|
||
"""放行 admin / auth-disabled,其余 403。
|
||
|
||
与 scheduled_tasks._require_admin 同套语义:auth_disabled 模式下
|
||
get_optional_user_from_request 返回 None,此时放行(开发环境一致体验)。
|
||
"""
|
||
from app.gateway.deps import get_optional_user_from_request
|
||
|
||
user = await get_optional_user_from_request(request)
|
||
if user is None:
|
||
return
|
||
if getattr(user, "system_role", None) != "admin":
|
||
raise HTTPException(status_code=403, detail="AI 写作清理管理接口仅限管理员")
|
||
|
||
|
||
def _resolve_cleanup_params(override: AdminCleanupRequest) -> dict:
|
||
"""合并 config.yaml 默认值 + 请求显式覆盖。
|
||
|
||
显式传 0 / False 也算覆盖(用 ``is not None`` 判定),方便测「保留 0 天 =
|
||
全删」这种极端场景。
|
||
"""
|
||
from deerflow.config import get_app_config
|
||
|
||
raw = getattr(get_app_config(), "ai_writing", None) or {}
|
||
cfg = (raw.get("cleanup") if isinstance(raw, dict) else None) or {}
|
||
|
||
return {
|
||
"retention_days": (
|
||
override.retention_days if override.retention_days is not None
|
||
else int(cfg.get("retention_days", 7))
|
||
),
|
||
"only_finished": (
|
||
override.only_finished if override.only_finished is not None
|
||
else bool(cfg.get("delete_only_finished", False))
|
||
),
|
||
"batch_size": (
|
||
override.batch_size if override.batch_size is not None
|
||
else int(cfg.get("batch_size", 500))
|
||
),
|
||
}
|
||
|
||
|
||
@router.post("/admin/cleanup", response_model=AdminCleanupResponse)
|
||
async def admin_cleanup(
|
||
body: AdminCleanupRequest,
|
||
request: Request,
|
||
) -> AdminCleanupResponse:
|
||
"""**仅管理员** 手动触发一次会话清理(不等凌晨 3 点定时调度)。
|
||
|
||
用途:
|
||
- 验证 cleanup 配置改动是否生效
|
||
- 内网紧急清理(磁盘满 / MySQL 行数过多)
|
||
- 调试:传 ``retention_days: 0`` 一刷拉清空,方便重新跑回归
|
||
|
||
所有参数留空 → 沿用 config.yaml 的 ``ai_writing.cleanup`` 段默认。显式
|
||
传值会覆盖配置但**不修改 config.yaml**(一次性参数,下次定时调度还是用
|
||
配置文件的值)。
|
||
"""
|
||
import time
|
||
|
||
from app.gateway.ai_writing_cleanup import run_cleanup_once
|
||
|
||
await _require_admin(request)
|
||
params = _resolve_cleanup_params(body)
|
||
|
||
# 透传主 LangGraph checkpointer / thread_store —— 新路径下 ai_writing 的
|
||
# thread 状态在主 checkpointer 里,thread_store 用来按 status 过滤活跃 thread。
|
||
# getattr 嵌套防御:测试用 stub Request 时 ``request.app`` 可能不存在。
|
||
_app_state = getattr(getattr(request, "app", None), "state", None)
|
||
main_checkpointer = getattr(_app_state, "checkpointer", None)
|
||
main_thread_store = getattr(_app_state, "thread_store", None)
|
||
|
||
started = time.monotonic()
|
||
try:
|
||
result = await run_cleanup_once(
|
||
**params,
|
||
checkpointer=main_checkpointer,
|
||
thread_store=main_thread_store,
|
||
)
|
||
except Exception as e:
|
||
logger.exception("ai_writing admin cleanup failed: params=%s", params)
|
||
# 抛 500 是合理的 —— 管理员触发清理失败需要看到失败原因,跟「业务
|
||
# 接口被失败拖累」是两回事
|
||
raise HTTPException(
|
||
status_code=500,
|
||
detail=f"清理执行失败: {type(e).__name__}: {e}",
|
||
) from e
|
||
|
||
elapsed = time.monotonic() - started
|
||
logger.info(
|
||
"ai_writing admin cleanup done: params=%s deleted_db=%d deleted_cp=%d elapsed=%.2fs",
|
||
params,
|
||
result["deleted_db"],
|
||
result["deleted_checkpoints"],
|
||
elapsed,
|
||
)
|
||
return AdminCleanupResponse(
|
||
deleted_db=result["deleted_db"],
|
||
deleted_checkpoints=result["deleted_checkpoints"],
|
||
skipped_active=result.get("skipped_active", 0),
|
||
elapsed_seconds=round(elapsed, 3),
|
||
params=params,
|
||
)
|