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

547 lines
25 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 写作「后台挂起」执行器。
把 AI 写作 LangGraph 图作为**后台 asyncio 任务**驱动到定稿,survive 发起请求的
生命周期(离开页面 / 刷新 / 切换账号都不打断)。
与圆桌后台作业(``roundtable_job_executor``)的根本区别:AI 写作的编排**就是
LangGraph 图本身**(4 个 ``interrupt()`` 干预点),所以这里不另起编排引擎,而是
直接驱动同一张图——用 app 启动时创建的**共享 checkpointer**(``app.state.checkpointer``,
即运行时 worker 跑该图用的同一个 saver,见 ``runtime/runs/worker.py`` 的
``agent.checkpointer = ctx.checkpointer``)续上前端已经写到 checkpoint 的线程状态,
在每个 ``interrupt()`` 处用「默认放行」动作自动 resume,一路推进到定稿:
素材确认 → confirm
大纲确认 → confirm
缺素材求助 → 全部「通用知识续写」(loose,不回检索、不丢章节)
草稿确认 → 未审过先 to_editor,审过(通过/用完次数)再 finalize
编辑打回 → accept_review(按编辑意见重写),直到通过或用完 max_revisions
图节点本就把进度落库(``writer_draft`` 写 draft、``editor`` 写 review_result、
``pause_*`` finalize 时 ``_mark_session_done`` 写 status=done),所以后台跑完用户
从历史里重新打开会话即可看到成稿——这里无需额外的产物表。
"""
from __future__ import annotations
import asyncio
import logging
import os
from langgraph.types import Command
from deerflow.runtime.user_context import AUTO, reset_current_user, set_current_user
logger = logging.getLogger(__name__)
def _slots_per_user() -> int:
"""每个用户同时真正在跑的后台写作上限(默认 2)。
根因与圆桌同源:checkpointer 多为单连接 SQLite(一把 asyncio.Lock 串行全进程
checkpoint 读写),后台作业无并发上限会把多张图的模型轮次压到同一把锁 + 同一
离线模型端点,拖慢前台。按用户限流后随便挂多少都能建,只是同时真跑的 ≤ N。
前台实时路径走 LangGraph 运行时 worker,不经这里,永不受限。
"""
raw = os.getenv("AI_WRITING_BACKGROUND_SLOTS_PER_USER", "").strip()
if not raw:
return 2
try:
return max(1, int(raw))
except ValueError:
return 2
# 每用户一把信号量(懒建)。键为 user_id(None → "default")。
_user_semaphores: dict[str, asyncio.Semaphore] = {}
def _user_semaphore(user_id: str | None) -> asyncio.Semaphore:
key = user_id or "default"
sem = _user_semaphores.get(key)
if sem is None:
sem = asyncio.Semaphore(_slots_per_user())
_user_semaphores[key] = sem
return sem
class _JobUser:
"""最小 CurrentUser:满足 user_context 的 ``.id`` 结构协议。"""
def __init__(self, user_id: str | None) -> None:
self.id = user_id
# 后台运行中的会话状态值(历史下拉据此显示「挂起中」徽标)。
STATUS_BACKGROUND = "background"
# 单个会话最多自动推进的步数(resume / 续跑各算一步)。安全上限,防图拓扑异常死循环。
# 正常一篇:素材→大纲→草稿→审核(最多 max_revisions 轮)→定稿,远小于这个数。
_MAX_DRIVE_STEPS = 60
# 「等待可操作态」轮询:用户在节点执行中(如素材搜集)点后台挂起时,运行时 worker
# 可能还在跑那个节点。此时**绝不能**并发 ainvoke 同一线程(会撞坏 checkpoint),先观察:
# - 出现 interrupt / 走到 done → 立刻接管;
# - 停在节点间且 checkpoint 连续 _INFLIGHT_MAX_STABLE 次不再推进 → 判定无 worker 在跑
# (被遗弃 / 进程重启续跑),可安全自驱。
_INFLIGHT_POLL_SECONDS = 1.5
_INFLIGHT_MAX_STABLE = 3
# 等待极长节点(LLM 慢)的轮询上限,≈ _INFLIGHT_MAX_WAIT_POLLS * _INFLIGHT_POLL_SECONDS 秒。
_INFLIGHT_MAX_WAIT_POLLS = 800
# 进度弹窗:把「下一个待执行节点 / 当前 interrupt 暂停点」映射成前端流程图的规范化阶段 key。
_NODE_TO_STAGE = {
"intent_parser": "intent",
"researcher": "research",
"pause_material": "material_confirm",
"writer_outline": "outline",
"pause_outline": "outline_confirm",
"writer_draft": "draft",
"pause_section_help": "draft",
"pause_draft": "draft_confirm",
"editor": "review",
"pause_review": "review",
}
_PAUSE_TO_STAGE = {
"material_confirm": "material_confirm",
"outline_confirm": "outline_confirm",
"section_help": "draft",
"draft_confirm": "draft_confirm",
"review_confirm": "review",
}
# 暂停点附带的细化说明(流水线主链之外的分支)。
_STAGE_NOTE = {
"section_help": "素材不足求助",
"review_confirm": "编辑打回",
}
# resumed 事件 message → (pause_point, action)。与 graph.py 各 pause 节点 emit 的 resumed
# 文案严格对齐(见 graph.py pause_material/outline/draft/section_help/review)。从 checkpoint
# 重建 await_user/干预卡时据此分类——await_user 与 completedInterventions 都源自同一批 resumed
# 事件,永远 1:1 对齐(不再依赖前端可能滞后的已存 interventions,根治错位 off-by-one)。
def _classify_resumed(message: str) -> tuple[str | None, str | None]:
m = message or ""
if "确认素材" in m:
return "material_confirm", "confirm"
if "匹配技能检索" in m or "重新检索" in m or "重新搜索" in m:
return "material_confirm", "re_search"
if "确认大纲" in m:
return "outline_confirm", "confirm"
if "重新规划大纲" in m:
return "outline_confirm", "re_outline"
if "提交编辑审核" in m:
return "draft_confirm", "to_editor"
if "强制定稿" in m:
return "review_confirm", "force_finalize"
if "确认定稿" in m:
return "draft_confirm", "finalize"
if "按用户意见重写" in m:
return "draft_confirm", "user_revise"
if "按编辑意见重写" in m:
return "review_confirm", "accept_review"
if "补充检索素材" in m:
return "section_help", "supplement"
if "继续写作" in m:
return "section_help", "loose"
return None, None
# 暂停提示文案(await_user 的 message)——对齐 graph.py 各 interrupt() 的 message,
# 让历史回显里暂停点显示与实时一致的提示语。
_PAUSE_MESSAGES = {
"material_confirm": "素材收集专家已完成素材搜集,请确认素材后继续。",
"outline_confirm": "作家已完成大纲规划,请确认或编辑大纲后继续。",
"draft_confirm": "草稿已完成,请确认或给出修改意见。",
"section_help": "部分章节缺少足够素材支撑写作,请决定如何处理。",
"review_confirm": "编辑认为文章需要修改,请决定下一步操作。",
}
def _build_transcript_from_checkpoint(values: dict) -> tuple[list[dict], list[dict]]:
"""从最终 checkpoint 重建 (progressEvents, completedInterventions),二者自洽 1:1 对齐。
每条 ``resumed`` 事件 = 一个已处理的暂停点:在其前插一条 ``await_user``(带暂停提示
文案),并据 resumed 文案分类出 pausePoint+action 生成**带富数据**的干预卡(素材列表/
大纲/草稿标题字数/审核结果取自最终 state;snake_case 字段由前端归一化)。素材确认前
再插一条 ``materials_ready``。其余里程碑(intent_ready/step_*/outline_ready/draft_ready/
review_ready/notice/error)节点本就 append 进 progress_events,原位即正确顺序。
时间线靠 ``await_user`` 锚点按序取 completedInterventions 渲染干预卡——这里两者都源自
同一批 resumed,数量与顺序天然一致,根治了「卡片错位 / 只显示标题没数据」的 bug。
"""
raw_events = list(values.get("progress_events") or [])
pkg = values.get("material_package") or {}
outline = values.get("current_outline")
draft = values.get("current_draft") or {}
review = values.get("review_result")
status = values.get("status")
events: list[dict] = []
interventions: list[dict] = []
materials_inserted = not pkg
for ev in raw_events:
if isinstance(ev, dict) and ev.get("type") == "resumed":
pp, action = _classify_resumed(ev.get("message") or "")
if pp == "material_confirm" and not materials_inserted:
events.append({
"type": "materials_ready",
"materials": pkg.get("materials") or [],
"keywords": pkg.get("keywords") or [],
"summary": pkg.get("summary") or "",
})
materials_inserted = True
events.append({"type": "await_user", "pause_point": pp, "message": _PAUSE_MESSAGES.get(pp, "")})
rec: dict = {"pausePoint": pp, "action": action}
if pp == "material_confirm":
rec["materials"] = pkg.get("materials") or []
elif pp == "outline_confirm":
if outline:
rec["outline"] = outline
elif pp == "draft_confirm":
rec["draftTitle"] = draft.get("title") or ""
rec["draftWordCount"] = int(draft.get("word_count") or 0)
if review:
rec["reviewResult"] = review
elif pp == "review_confirm":
if review:
rec["reviewResult"] = review
interventions.append(rec)
events.append(ev)
if not materials_inserted:
# 没有 material_confirm 暂停(罕见)但有素材包 → 插到第一条 outline_ready 前,否则附末尾。
mr = {"type": "materials_ready", "materials": pkg.get("materials") or [],
"keywords": pkg.get("keywords") or [], "summary": pkg.get("summary") or ""}
idx = next((i for i, e in enumerate(events) if isinstance(e, dict) and e.get("type") == "outline_ready"), None)
events.insert(idx, mr) if idx is not None else events.append(mr)
if status == "done" and not any(isinstance(e, dict) and e.get("type") == "done" for e in events):
events.append({"type": "done", "message": "写作完成"})
return events, interventions
def _checkpoint_id(snapshot) -> str | None:
"""取快照的 checkpoint_id,用于判断图状态是否在被(worker)推进。"""
cfg = getattr(snapshot, "config", None) or {}
try:
return (cfg.get("configurable") or {}).get("checkpoint_id")
except AttributeError:
return None
def _auto_resume_payload(pause_point: str | None, interrupt_value: dict, values: dict) -> dict:
"""给定当前 ``interrupt()`` 暂停点 + 图状态,返回「默认放行」的 resume 入参。
设计目标:一路推进到定稿且**不绕回检索 / 不无限重审**。判定逻辑见模块 docstring。
"""
if pause_point == "material_confirm":
# 不传 approvedMaterialIds → 节点默认采用全部素材。
return {"action": "confirm"}
if pause_point == "outline_confirm":
# 不传 editedOutline → 节点采用当前大纲。
return {"action": "confirm"}
if pause_point == "section_help":
# 严格模式缺素材求助:全部章节按「通用知识续写」(loose)处理——既不回检索
# (supplement 会绕回 researcher 再触发暂停),也不删章节,直接续写到底。
blocked = (interrupt_value or {}).get("blocked_sections") or values.get("blocked_sections") or []
decisions = []
for b in blocked:
title = (b or {}).get("section_title") or (b or {}).get("sectionTitle") or ""
if title:
decisions.append({"sectionTitle": title, "decision": "loose"})
return {"action": "section_help", "sectionDecisions": decisions}
if pause_point == "draft_confirm":
review = values.get("review_result")
rev = int(values.get("revision_count") or 0)
max_rev = int(values.get("max_revisions") or 3)
if not review:
# 还没经过编辑审核 → 先送审。
return {"action": "to_editor"}
if (review or {}).get("verdict") == "pass":
# 审核通过 → 定稿。
return {"action": "finalize"}
if rev >= max_rev:
# 用完修改次数仍未通过(after_review 强制落到草稿确认)→ 接受现稿定稿。
return {"action": "finalize"}
# 刚按编辑意见改过、review 还是旧的不通过结果 → 再次送审,拿新一轮审核。
return {"action": "to_editor"}
if pause_point == "review_confirm":
# 编辑打回(仍有修改额度)→ 按编辑意见重写,重写完会再次 to_editor 复审。
return {"action": "accept_review"}
# 未知暂停点:保守放行。
return {"action": "confirm"}
def _pending_interrupt(snapshot) -> dict | None:
"""从状态快照里取第一个待处理 ``interrupt()`` 的 value(dict),无则 None。"""
for task in getattr(snapshot, "tasks", None) or ():
for intr in getattr(task, "interrupts", None) or ():
value = getattr(intr, "value", None)
if isinstance(value, dict):
return value
return {"value": value}
return None
class AIWritingJobExecutor:
"""驱动 AI 写作图在后台跑到定稿。按 thread_id 幂等、按用户限流。"""
def __init__(self, app) -> None:
# 闭包持有 app,运行时再读 ``app.state.checkpointer`` / session store——它们在
# lifespan 里创建,运行时读保证拿到最新实例(与 worker 用同一个 checkpointer)。
self._app = app
self._tasks: dict[str, asyncio.Task] = {}
# ── 公开 API ──────────────────────────────────────────────────────────────
def start(self, *, thread_id: str, user_id: str | None) -> bool:
"""启动一个后台驱动任务。同一 thread 已在跑则忽略(幂等)。返回是否新启动。"""
existing = self._tasks.get(thread_id)
if existing is not None and not existing.done():
return False
task = asyncio.create_task(self._run(thread_id, user_id))
self._tasks[thread_id] = task
task.add_done_callback(lambda t, tid=thread_id: self._tasks.pop(tid, None))
return True
def is_running(self, thread_id: str) -> bool:
task = self._tasks.get(thread_id)
return task is not None and not task.done()
async def get_progress(self, thread_id: str) -> dict:
"""读取该写作线程当前所处阶段(从 graph checkpoint 推导),供「查看进度」弹窗。
返回规范化的 ``stage`` key(见 ``_NODE_TO_STAGE`` / ``_PAUSE_TO_STAGE``,与前端
流程图阶段一一对应)+ 可选 ``note`` + 审核进度。无 checkpoint / 读失败 → stage=None。
"""
empty = {"stage": None, "note": None, "revisionCount": 0, "maxRevisions": 3, "reviewVerdict": None}
checkpointer = self._checkpointer()
if checkpointer is None:
return empty
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": thread_id}})
except Exception: # noqa: BLE001 — 读不到 checkpoint 不致命
return empty
values = getattr(snapshot, "values", None) or {}
review = values.get("review_result") or {}
out = {
"stage": None,
"note": None,
"revisionCount": int(values.get("revision_count") or 0),
"maxRevisions": int(values.get("max_revisions") or 3),
"reviewVerdict": review.get("verdict"),
}
if not values:
return out # 还没建 checkpoint
if values.get("status") == "done":
out["stage"] = "done"
return out
interrupt_value = _pending_interrupt(snapshot)
if interrupt_value is not None:
pp = interrupt_value.get("pause_point")
out["stage"] = _PAUSE_TO_STAGE.get(pp, pp)
out["note"] = _STAGE_NOTE.get(pp)
return out
next_nodes = list(getattr(snapshot, "next", None) or ())
if next_nodes:
node = next_nodes[0]
out["stage"] = _NODE_TO_STAGE.get(node, node)
return out
# 无 interrupt、无 next → 已结束。
out["stage"] = "done"
return out
def cancel(self, thread_id: str) -> bool:
task = self._tasks.get(thread_id)
if task is None or task.done():
return False
task.cancel()
return True
async def await_job(self, thread_id: str) -> None:
"""等待某后台任务结束(测试用;生产不调用)。"""
task = self._tasks.get(thread_id)
if task is not None:
try:
await task
except (Exception, asyncio.CancelledError): # noqa: BLE001
pass
def reconcile(self, *, sessions: list[dict]) -> int:
"""启动时把仍标记为 ``background`` 的会话重新拉起驱动(进程重启续跑)。
``sessions`` 为 ``[{id, user_id}, ...]``。返回拉起的任务数。
"""
launched = 0
for s in sessions:
tid = s.get("id")
if not tid:
continue
if self.start(thread_id=tid, user_id=s.get("user_id")):
launched += 1
if launched:
logger.info("[ai-writing-bg] reconcile relaunched %d background session(s)", launched)
return launched
# ── 内部 ────────────────────────────────────────────────────────────────
def _checkpointer(self):
return getattr(self._app.state, "checkpointer", None)
def _store(self):
return getattr(self._app.state, "ai_writing_session_store", None)
async def _set_status(self, thread_id: str, status: str) -> None:
store = self._store()
if store is None:
return
try:
await store.update(thread_id, user_id=AUTO, status=status)
except Exception: # noqa: BLE001 — 落库失败不致命
logger.warning("[ai-writing-bg] set status=%s for %s failed", status, thread_id, exc_info=True)
async def _run(self, thread_id: str, user_id: str | None) -> None:
token = set_current_user(_JobUser(user_id))
try:
checkpointer = self._checkpointer()
if checkpointer is None:
logger.warning("[ai-writing-bg] no checkpointer; cannot drive %s", thread_id)
return
await self._set_status(thread_id, STATUS_BACKGROUND)
from deerflow.agents.ai_writing.graph import make_ai_writing_graph
graph = make_ai_writing_graph()
# 与运行时 worker 同款:把共享 checkpointer 挂到图上,续上同一线程的 checkpoint。
graph.checkpointer = checkpointer
config = {"configurable": {"thread_id": thread_id}, "recursion_limit": 100}
# 按用户限流:同用户最多 N 篇后台写作同时真跑(见 _slots_per_user)。
async with _user_semaphore(user_id):
done, _interventions = await self._drive(graph, config)
# 把后台跑出的对话时间线写进 transcript(与前端直接写作的存取一致)。
await self._save_transcript(thread_id, graph, config)
if not done:
# 图已停在非 done 终态(无 interrupt 也无 next):兜底标 done,避免会话
# 永久卡在「挂起中」。正常路径下 finalize 已由节点 _mark_session_done 写 done。
await self._mark_terminal(thread_id, graph, config)
except asyncio.CancelledError:
logger.info("[ai-writing-bg] job %s cancelled", thread_id)
raise
except Exception: # noqa: BLE001 — 顶层兜底
logger.exception("[ai-writing-bg] job %s crashed", thread_id)
await self._set_status(thread_id, "error")
finally:
reset_current_user(token)
async def _wait_for_actionable(self, graph, config) -> tuple[object, str]:
"""等到图处于可安全操作的状态,返回 ``(snapshot, action)``。
action ∈ {``interrupt``(有干预点待 resume), ``done``(已结束),
``push``(停在节点间且无 worker 在推进,可自驱)}。
见 ``_INFLIGHT_*`` 注释:避免与运行时 worker 并发跑同一线程。
"""
last_ckpt: object = object()
stable = 0
for _ in range(_INFLIGHT_MAX_WAIT_POLLS):
snapshot = await graph.aget_state(config)
if (getattr(snapshot, "values", None) or {}).get("status") == "done":
return snapshot, "done"
if _pending_interrupt(snapshot) is not None:
return snapshot, "interrupt"
if not getattr(snapshot, "next", None):
return snapshot, "done" # 无中断、无待执行 → 图已结束
ckpt = _checkpoint_id(snapshot)
if ckpt != last_ckpt:
# checkpoint 在变 → 有 worker 正推进该节点,继续等它到下一个干预点。
last_ckpt = ckpt
stable = 0
else:
stable += 1
if stable >= _INFLIGHT_MAX_STABLE:
return snapshot, "push" # 连续不动 → 无 worker,自己推进
await asyncio.sleep(_INFLIGHT_POLL_SECONDS)
# 等待超时(极长节点)→ 保守按可推进处理。
return await graph.aget_state(config), "push"
async def _drive(self, graph, config) -> tuple[bool, list[dict]]:
"""推进图直到定稿(status==done)或无可推进。
返回 ``(是否走到 done, 本次自动处理的 interventions)``。每个干预点记录
``{pausePoint, action}``,供重建 transcript 时与 ``await_user`` 对齐。
"""
thread_id = (config.get("configurable") or {}).get("thread_id")
interventions: list[dict] = []
for _ in range(_MAX_DRIVE_STEPS):
snapshot, action = await self._wait_for_actionable(graph, config)
values = getattr(snapshot, "values", None) or {}
if action == "done":
return values.get("status") == "done", interventions
if action == "interrupt":
interrupt_value = _pending_interrupt(snapshot) or {}
pause_point = interrupt_value.get("pause_point")
payload = _auto_resume_payload(pause_point, interrupt_value, values)
interventions.append({"pausePoint": pause_point, "action": payload.get("action")})
logger.info("[ai-writing-bg] auto-resume thread=%s pause=%s action=%s",
thread_id, pause_point, payload.get("action"))
await graph.ainvoke(Command(resume=payload), config)
continue
# action == "push":停在节点间且无 worker 在跑 → 继续推进。
await graph.ainvoke(None, config)
logger.warning("[ai-writing-bg] drive hit max steps for thread=%s", thread_id)
return False, interventions
async def _save_transcript(self, thread_id: str, graph, config) -> None:
"""把后台跑出的对话时间线写进 ``transcript`` —— 与前端直接写作的存取完全一致。
完全从最终 checkpoint 自洽重建(``_build_transcript_from_checkpoint``):progressEvents
与 completedInterventions 都源自同一批 ``resumed`` 事件,永远 1:1 对齐——不再依赖前端
可能滞后的已存 interventions(那是「卡片错位」bug 的根因)。
"""
store = self._store()
if store is None:
return
try:
snapshot = await graph.aget_state(config)
except Exception: # noqa: BLE001
return
values = getattr(snapshot, "values", None) or {}
events, interventions = _build_transcript_from_checkpoint(values)
transcript = {"progressEvents": events, "completedInterventions": interventions}
try:
await store.update(
thread_id, user_id=AUTO, transcript=transcript, completed_interventions=interventions
)
except Exception: # noqa: BLE001 — transcript 是锦上添花,存失败不影响作业
logger.warning("[ai-writing-bg] save transcript for %s failed", thread_id, exc_info=True)
async def _mark_terminal(self, thread_id: str, graph, config) -> None:
try:
snapshot = await graph.aget_state(config)
status = (getattr(snapshot, "values", None) or {}).get("status")
except Exception: # noqa: BLE001
status = None
# 走到 done 由节点写过了;否则若图已停止(无 next)兜底标 done,停在中途则保留。
if status == "done":
return
try:
snapshot = await graph.aget_state(config)
if not getattr(snapshot, "next", None):
await self._set_status(thread_id, "done")
except Exception: # noqa: BLE001
logger.warning("[ai-writing-bg] mark terminal for %s failed", thread_id, exc_info=True)