547 lines
25 KiB
Python
547 lines
25 KiB
Python
"""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)
|