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

1108 lines
44 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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,
)