464 lines
22 KiB
Python
464 lines
22 KiB
Python
"""AI 写作 LangGraph Graph,含 4 个用户干预点。"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
|
||
from langchain_core.runnables import RunnableConfig
|
||
from langgraph.graph import END, StateGraph
|
||
from langgraph.types import interrupt
|
||
|
||
from .nodes.editor import editor_node
|
||
from .nodes.intent_parser import intent_parser_node
|
||
from .nodes.researcher import researcher_node
|
||
from .nodes.sample_analyzer import sample_analyzer_node
|
||
from .nodes.writer_draft import writer_draft_node
|
||
from .nodes.writer_outline import writer_outline_node
|
||
from .state import AIWritingState
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
async def _mark_session_done(config: RunnableConfig | None) -> None:
|
||
"""把业务表行 status 写成 ``done``。
|
||
|
||
pause_draft_confirm action=finalize / pause_review_confirm force_finalize 时调,
|
||
让 LangGraph Server 路径下新会话有正确终态(老路径靠 ``_stream_graph_segment``
|
||
finally 块写终态,新路径完全不走它)。失败不影响 graph 流程。
|
||
"""
|
||
if not config:
|
||
return
|
||
thread_id = (config.get("configurable") or {}).get("thread_id")
|
||
if not thread_id:
|
||
return
|
||
from deerflow.persistence.ai_writing_sessions import get_default_store
|
||
from deerflow.runtime.user_context import AUTO
|
||
|
||
repo = get_default_store()
|
||
if repo is None:
|
||
return
|
||
try:
|
||
await repo.update(thread_id, user_id=AUTO, status="done")
|
||
except Exception:
|
||
logger.exception("ai_writing: mark_session_done failed (non-fatal) thread_id=%s", thread_id)
|
||
|
||
# ── 干预节点 ──────────────────────────────────────────────────────────────────
|
||
|
||
async def pause_material_confirm(state: AIWritingState) -> dict:
|
||
"""PAUSE-1:素材确认"""
|
||
pkg = state["material_package"]
|
||
user_input: dict = interrupt({
|
||
"pause_point": "material_confirm",
|
||
"message": "素材收集专家已完成素材搜集,请确认素材后继续。",
|
||
"materials": pkg["materials"],
|
||
"keywords": pkg["keywords"],
|
||
"summary": pkg["summary"],
|
||
})
|
||
|
||
action = user_input.get("action", "confirm")
|
||
|
||
# 单输入框文本:可能是技能名 / 来源 / 检索意图,re_search 和 enable_skill_search 都用得到
|
||
user_query = (user_input.get("userQuery") or user_input.get("user_query") or "").strip()
|
||
requested_source = user_input.get("materialSource") or user_input.get("material_source")
|
||
fallback_source = state.get("material_source") or "general"
|
||
material_source = requested_source if requested_source in ("general", "knowledge_base", "notebook") else fallback_source
|
||
if material_source not in ("general", "knowledge_base", "notebook"):
|
||
material_source = "general"
|
||
source_label = {
|
||
"general": "通用检索",
|
||
"knowledge_base": "知识库检索",
|
||
"notebook": "我的空间",
|
||
}.get(material_source, "通用检索")
|
||
|
||
if action == "confirm":
|
||
approved_ids = set(user_input.get("approvedMaterialIds") or [m["id"] for m in pkg["materials"]])
|
||
filtered = [m for m in pkg["materials"] if m["id"] in approved_ids]
|
||
updated_pkg = {**pkg, "materials": filtered}
|
||
# 用户自带大纲:researcher 已在「搜集素材」阶段解析大纲并归类素材 → 跳过大纲生成步,
|
||
# 仅把被用户取消勾选的素材从各章节 material_ids 里剔除,直接进入大纲确认。
|
||
prebuilt = state.get("current_outline")
|
||
if prebuilt and (state.get("user_outline_raw") or "").strip():
|
||
return {
|
||
"material_package": updated_pkg,
|
||
"current_outline": _prune_outline_material_ids(prebuilt, approved_ids),
|
||
"last_user_intervention": user_input,
|
||
"status": "awaiting_outline_confirm",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "已采用你提供的大纲并归类素材,请确认大纲"}],
|
||
}
|
||
return {
|
||
"material_package": updated_pkg,
|
||
"last_user_intervention": user_input,
|
||
"status": "writing_outline",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "用户已确认素材,开始规划大纲"}],
|
||
}
|
||
|
||
# enable_skill_search:旧版「启用技能检索」按钮的兼容入口——通用检索现在
|
||
# 本来就按 agent 配置技能编排,researcher 把它按 user_query 追加检索处理
|
||
if action == "enable_skill_search":
|
||
msg = f"按输入「{user_query[:30]}」使用{source_label}追加素材…" if user_query else f"使用{source_label}重新检索中…"
|
||
return {
|
||
"last_user_intervention": {**user_input, "action": "enable_skill_search", "user_query": user_query, "material_source": material_source},
|
||
"material_source": material_source,
|
||
"status": "researching",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": msg}],
|
||
}
|
||
|
||
# re_search:带 user_query = 在配置技能中匹配检索并**追加**素材;
|
||
# 不带输入 = 重新生成检索词、覆盖式重检(兼容旧版 extraKeywords 数组追加)
|
||
extra_kws = user_input.get("extraKeywords") or []
|
||
msg = f"按输入「{user_query[:30]}」使用{source_label}追加素材…" if user_query else f"使用{source_label}重新搜索中..."
|
||
return {
|
||
"last_user_intervention": {**user_input, "_extra_keywords": extra_kws, "user_query": user_query, "material_source": material_source},
|
||
"material_source": material_source,
|
||
"status": "researching",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": msg}],
|
||
}
|
||
|
||
|
||
def _prune_outline_material_ids(outline: dict, approved_ids: set) -> dict:
|
||
"""把大纲各章节 material_ids 里、用户在素材确认卡取消勾选的素材编号剔除掉。
|
||
|
||
用户自带大纲时素材归类发生在「搜集素材」阶段(researcher),早于素材确认;
|
||
用户在确认卡里去掉某些素材后,用这个函数把对应编号从大纲剔除,保持一致。
|
||
"""
|
||
if not outline:
|
||
return outline
|
||
sections = []
|
||
for s in outline.get("sections", []):
|
||
kept = [mid for mid in (s.get("material_ids") or []) if mid in approved_ids]
|
||
sections.append({**s, "material_ids": kept})
|
||
return {**outline, "sections": sections}
|
||
|
||
|
||
def _normalize_outline(outline: dict) -> dict:
|
||
"""将前端回传的 camelCase 大纲统一为后端 snake_case 格式。"""
|
||
if not outline:
|
||
return outline
|
||
sections = []
|
||
for s in outline.get("sections", []):
|
||
sections.append({
|
||
"section_title": s.get("section_title") or s.get("sectionTitle", ""),
|
||
"key_points": s.get("key_points") or s.get("keyPoints", []),
|
||
"material_ids": s.get("material_ids") or s.get("materialIds", []),
|
||
})
|
||
return {"title": outline.get("title", ""), "sections": sections}
|
||
|
||
|
||
def _check_outline_material_support(outline: dict) -> str:
|
||
"""轻量启发式:统计未关联任何素材的章节,返回预警文案(空串=无需预警)。
|
||
|
||
仅做提示、不拦截流程;真正的素材不足会在写作节点被检出并触发求助暂停。
|
||
"""
|
||
sections = (outline or {}).get("sections", [])
|
||
weak = [
|
||
(s.get("section_title") or s.get("sectionTitle") or "")
|
||
for s in sections
|
||
if not (s.get("material_ids") or s.get("materialIds"))
|
||
]
|
||
weak = [t for t in weak if t]
|
||
if not weak:
|
||
return ""
|
||
return (
|
||
f"有 {len(weak)} 个章节未关联检索素材({('、'.join(weak))})。"
|
||
"严格依据素材模式下,这些章节在写作时可能因素材不足而暂停向你求助。"
|
||
)
|
||
|
||
|
||
async def pause_outline_confirm(state: AIWritingState) -> dict:
|
||
"""PAUSE-2:大纲确认"""
|
||
user_input: dict = interrupt({
|
||
"pause_point": "outline_confirm",
|
||
"message": "作家已完成大纲规划,请确认或编辑大纲后继续。",
|
||
"outline": state["current_outline"],
|
||
})
|
||
|
||
action = user_input.get("action", "confirm")
|
||
|
||
if action == "confirm":
|
||
raw_outline = user_input.get("editedOutline") or state["current_outline"]
|
||
edited = _normalize_outline(raw_outline) if raw_outline else state["current_outline"]
|
||
events: list[dict] = [{"type": "resumed", "agent_name": "系统", "message": "用户已确认大纲,开始写作"}]
|
||
# 阶段四:严格模式下对大纲做轻量素材充分性预警(仅提示,不拦截流程)
|
||
if state.get("writing_mode") == "strict":
|
||
warning = _check_outline_material_support(edited)
|
||
if warning:
|
||
events.append({
|
||
"type": "notice", "agent_name": "系统",
|
||
"title": "素材充分性提示", "detail": warning,
|
||
})
|
||
return {
|
||
"current_outline": edited,
|
||
"last_user_intervention": user_input,
|
||
"status": "writing_draft",
|
||
"progress_events": events,
|
||
}
|
||
|
||
# re_outline
|
||
feedback = user_input.get("outlineFeedback", "")
|
||
updates: dict = {
|
||
"last_user_intervention": {**user_input, "_feedback": feedback},
|
||
"status": "writing_outline",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": f"重新规划大纲({feedback[:30]})"}],
|
||
}
|
||
# 用户在卡片里手动改过大纲 → 以手动版作为本次重规划的编辑基准
|
||
raw_edited = user_input.get("editedOutline")
|
||
if raw_edited:
|
||
updates["current_outline"] = _normalize_outline(raw_edited)
|
||
return updates
|
||
|
||
|
||
async def pause_draft_confirm(state: AIWritingState, config: RunnableConfig) -> dict:
|
||
"""PAUSE-3:草稿确认"""
|
||
draft = state.get("current_draft")
|
||
# 防御:草稿缺失(写作节点异常未产出)时不进入确认暂停,转 error 优雅结束,
|
||
# 避免 draft["title"] 在 None 上崩溃。正常路径下 after_writer_draft 已拦截
|
||
# status:error,这里是兜底第二道防线(如 editor 异常等边角情形)。
|
||
if not draft:
|
||
return {
|
||
"status": "error",
|
||
"progress_events": [{
|
||
"type": "error", "agent_name": "作家",
|
||
"title": "草稿缺失", "detail": "草稿尚未生成,无法进入确认环节,请重试。",
|
||
"message": "草稿生成失败",
|
||
}],
|
||
}
|
||
user_input: dict = interrupt({
|
||
"pause_point": "draft_confirm",
|
||
"message": "草稿已完成,请确认或给出修改意见。",
|
||
"draft_title": draft["title"],
|
||
"revision_count": draft["revision_count"],
|
||
"review_result": state.get("review_result"),
|
||
})
|
||
|
||
action = user_input.get("action", "to_editor")
|
||
|
||
if action == "finalize":
|
||
await _mark_session_done(config)
|
||
return {
|
||
"last_user_intervention": user_input,
|
||
"status": "done",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "用户确认定稿"}],
|
||
}
|
||
|
||
if action == "user_revise":
|
||
notes = user_input.get("userRevisionNotes", "")
|
||
return {
|
||
"last_user_intervention": {**user_input, "_user_notes": notes},
|
||
"revision_count": state.get("revision_count", 0) + 1,
|
||
"status": "revising",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": f"按用户意见重写({notes[:30]})"}],
|
||
}
|
||
|
||
# to_editor:用户可能在「您的修改意见」里填了关注点再点「让编辑继续审核」,
|
||
# 把它带给编辑节点,让三维度审查重点核查用户在意的方面(不再被丢弃)。
|
||
editor_notes = (user_input.get("userRevisionNotes") or "").strip()
|
||
msg = (
|
||
f"提交编辑审核(已附带你的修改意见:{editor_notes[:30]})"
|
||
if editor_notes else "提交编辑审核"
|
||
)
|
||
return {
|
||
"last_user_intervention": {**user_input, "_editor_notes": editor_notes},
|
||
"status": "reviewing",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": msg}],
|
||
}
|
||
|
||
|
||
async def pause_section_help(state: AIWritingState) -> dict:
|
||
"""素材不足求助点:严格模式下作家发现章节缺素材,让用户逐章决定如何处理。
|
||
|
||
用户对每个被阻塞章节选择:
|
||
- supplement:补充检索素材(整体回到检索流程重走)
|
||
- delete:删除该章节
|
||
- loose:按通用知识续写该章节
|
||
"""
|
||
blocked = state.get("blocked_sections") or []
|
||
user_input: dict = interrupt({
|
||
"pause_point": "section_help",
|
||
"message": "部分章节缺少足够素材支撑写作,请决定如何处理。",
|
||
"blocked_sections": blocked,
|
||
})
|
||
|
||
# sectionDecisions: [{"sectionTitle": "...", "decision": "supplement|delete|loose"}]
|
||
by_title: dict = {}
|
||
for d in user_input.get("sectionDecisions") or []:
|
||
title = d.get("sectionTitle") or d.get("section_title")
|
||
if title:
|
||
by_title[title] = d.get("decision", "loose")
|
||
|
||
# 任一章节选「补充检索」→ 整体回到检索,重走「检索 → 素材确认 → 大纲 → 写作」
|
||
supplement_reasons = [
|
||
b["reason"] for b in blocked if by_title.get(b["section_title"]) == "supplement"
|
||
]
|
||
if supplement_reasons:
|
||
return {
|
||
# supplement_search:补充检索(保留已有素材、追加新素材),完成后直接回写作、不重排大纲
|
||
"last_user_intervention": {**user_input, "action": "supplement_search", "_extra_keywords": supplement_reasons},
|
||
"status": "researching",
|
||
"blocked_sections": None,
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "补充检索素材中…"}],
|
||
}
|
||
|
||
# 否则:删除标记章节、记录「通用知识续写」章节,直接回写作
|
||
outline = state.get("current_outline") or {"title": "", "sections": []}
|
||
kept, loose_titles = [], []
|
||
for s in outline.get("sections", []):
|
||
title = s.get("section_title") or s.get("sectionTitle", "")
|
||
decision = by_title.get(title)
|
||
if decision == "delete":
|
||
continue
|
||
if decision == "loose":
|
||
loose_titles.append(title)
|
||
kept.append(s)
|
||
# 防御:用户把所有章节都删了 → 保留原大纲并全部按通用知识写,避免空文章
|
||
if not kept:
|
||
kept = outline.get("sections", [])
|
||
loose_titles = [s.get("section_title") or s.get("sectionTitle", "") for s in kept]
|
||
return {
|
||
"current_outline": {**outline, "sections": kept},
|
||
"last_user_intervention": {**user_input, "_loose_sections": loose_titles},
|
||
"status": "writing_draft",
|
||
"blocked_sections": None,
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "已按你的选择继续写作…"}],
|
||
}
|
||
|
||
|
||
async def pause_review_confirm(state: AIWritingState, config: RunnableConfig) -> dict:
|
||
"""PAUSE-4:编辑打回确认"""
|
||
user_input: dict = interrupt({
|
||
"pause_point": "review_confirm",
|
||
"message": "编辑认为文章需要修改,请决定下一步操作。",
|
||
"review_result": state["review_result"],
|
||
"revision_count": state.get("revision_count", 0),
|
||
"max_revisions": state.get("max_revisions", 3),
|
||
})
|
||
|
||
if user_input.get("overrideVerdict") or user_input.get("action") == "force_finalize":
|
||
await _mark_session_done(config)
|
||
return {
|
||
"last_user_intervention": user_input,
|
||
"status": "done",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "用户强制定稿"}],
|
||
}
|
||
|
||
action = user_input.get("action", "accept_review")
|
||
extra_notes = user_input.get("userRevisionNotes", "")
|
||
return {
|
||
"last_user_intervention": {**user_input, "_extra_notes": extra_notes, "action": action},
|
||
"revision_count": state.get("revision_count", 0) + 1,
|
||
"status": "revising",
|
||
"progress_events": [{"type": "resumed", "agent_name": "系统", "message": "按编辑意见重写"}],
|
||
}
|
||
|
||
|
||
# ── 条件路由 ──────────────────────────────────────────────────────────────────
|
||
|
||
def after_researcher(state: AIWritingState) -> str:
|
||
"""检索完成后:补充检索(素材不足求助)直接回写作,否则去素材确认。"""
|
||
return "writer_draft" if state.get("status") == "writing_draft" else "pause_material"
|
||
|
||
|
||
def after_material_pause(state: AIWritingState) -> str:
|
||
status = state["status"]
|
||
if status == "researching":
|
||
return "researcher"
|
||
# 用户自带大纲:researcher 已在搜集阶段整理好大纲 → 跳过 writer_outline 生成步,直接确认大纲。
|
||
if status == "awaiting_outline_confirm":
|
||
return "pause_outline"
|
||
return "writer_outline"
|
||
|
||
|
||
def after_outline_pause(state: AIWritingState) -> str:
|
||
return "writer_outline" if state["status"] == "writing_outline" else "writer_draft"
|
||
|
||
|
||
def after_writer_draft(state: AIWritingState) -> str:
|
||
"""写作完成后:素材不足则去求助暂停点,否则去草稿确认。
|
||
|
||
写作节点兜底成 ``status:error``(如模型创建失败)时直接结束本次 run —— 否则会
|
||
流到 pause_draft,而此时 ``current_draft`` 为 None,``pause_draft_confirm`` 取
|
||
``draft["title"]`` 会二次崩溃('NoneType' object is not subscriptable)。错误事件
|
||
已由写作节点 emit,run 优雅结束即可,用户重试不受影响。
|
||
"""
|
||
if state["status"] == "error":
|
||
return "end"
|
||
return "pause_section_help" if state["status"] == "awaiting_section_help" else "pause_draft"
|
||
|
||
|
||
def after_section_help_pause(state: AIWritingState) -> str:
|
||
"""求助暂停点之后:补充检索回 researcher,否则回写作。"""
|
||
return "researcher" if state["status"] == "researching" else "writer_draft"
|
||
|
||
|
||
def after_draft_pause(state: AIWritingState) -> str:
|
||
status = state["status"]
|
||
if status == "done":
|
||
return "end"
|
||
if status == "revising":
|
||
return "writer_draft"
|
||
return "editor"
|
||
|
||
|
||
def after_review(state: AIWritingState) -> str:
|
||
review = state.get("review_result") or {}
|
||
if review.get("verdict") == "pass":
|
||
return "pause_draft"
|
||
rev_count = state.get("revision_count", 0)
|
||
max_rev = state.get("max_revisions", 3)
|
||
if rev_count >= max_rev:
|
||
return "pause_draft" # 超出上限,强制到草稿确认
|
||
return "pause_review"
|
||
|
||
|
||
def after_review_pause(state: AIWritingState) -> str:
|
||
return "end" if state["status"] == "done" else "writer_draft"
|
||
|
||
|
||
|
||
# ── LangGraph Server 入口 ────────────────────────────────────────────────────
|
||
#
|
||
# compile 出裸 graph(不带 checkpointer),由 LangGraph Server 注入主 checkpointer
|
||
# (langgraph.json 里 ``make_checkpointer`` 配的那个)。AI 写作的 thread 状态与
|
||
# 主聊天共享同一份 saver —— LangGraph schema 用 ``(thread_id, checkpoint_ns)`` 做
|
||
# 主键,UUID 不会撞。
|
||
def make_ai_writing_graph(config=None):
|
||
"""LangGraph Server graph factory:不接收 checkpointer,由 Server 注入。
|
||
|
||
Args:
|
||
config: LangGraph Server 调用工厂时注入的 RunnableConfig(暂未消费)。
|
||
"""
|
||
g = StateGraph(AIWritingState)
|
||
|
||
g.add_node("sample_analyzer", sample_analyzer_node)
|
||
g.add_node("intent_parser", intent_parser_node)
|
||
g.add_node("researcher", researcher_node)
|
||
g.add_node("pause_material", pause_material_confirm)
|
||
g.add_node("writer_outline", writer_outline_node)
|
||
g.add_node("pause_outline", pause_outline_confirm)
|
||
g.add_node("writer_draft", writer_draft_node)
|
||
g.add_node("pause_section_help", pause_section_help)
|
||
g.add_node("pause_draft", pause_draft_confirm)
|
||
g.add_node("editor", editor_node)
|
||
g.add_node("pause_review", pause_review_confirm)
|
||
|
||
g.set_entry_point("sample_analyzer")
|
||
g.add_edge("sample_analyzer", "intent_parser")
|
||
g.add_edge("intent_parser", "researcher")
|
||
g.add_conditional_edges("researcher", after_researcher,
|
||
{"pause_material": "pause_material", "writer_draft": "writer_draft", "writer_outline": "writer_outline"})
|
||
g.add_conditional_edges("pause_material", after_material_pause,
|
||
{"researcher": "researcher", "writer_outline": "writer_outline", "pause_outline": "pause_outline"})
|
||
g.add_edge("writer_outline", "pause_outline")
|
||
g.add_conditional_edges("pause_outline", after_outline_pause,
|
||
{"writer_outline": "writer_outline", "writer_draft": "writer_draft"})
|
||
g.add_conditional_edges("writer_draft", after_writer_draft,
|
||
{"pause_section_help": "pause_section_help", "pause_draft": "pause_draft", "end": END})
|
||
g.add_conditional_edges("pause_section_help", after_section_help_pause,
|
||
{"researcher": "researcher", "writer_draft": "writer_draft"})
|
||
g.add_conditional_edges("pause_draft", after_draft_pause,
|
||
{"end": END, "writer_draft": "writer_draft", "editor": "editor"})
|
||
g.add_conditional_edges("editor", after_review,
|
||
{"pause_draft": "pause_draft", "pause_review": "pause_review"})
|
||
g.add_conditional_edges("pause_review", after_review_pause,
|
||
{"end": END, "writer_draft": "writer_draft"})
|
||
|
||
return g.compile()
|