deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/ai_writing/graph.py
2026-09-07 18:24:55 +08:00

464 lines
22 KiB
Python
Raw 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 写作 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()