"""AI 写作 assistant 冒烟脚本:验证 ai_writing graph 注册到 LangGraph Server 后, interrupt/Command(resume=...)/state 持久化/custom stream 全链路是否跑通。 阶段 0 用名为 ai_writing_probe;阶段 1 起改为正式 assistant 名 ai_writing,本脚本随之切换。 前置条件: 1. langgraph.json 已注册 ai_writing(见 ai-writing-优化.md 阶段 1)。 2. 后端已通过 `make dev` 启动,Gateway + LangGraph Server 监听 8001。 3. config.yaml 配好至少一个可用 model,researcher 用到的 search_provider 也已配置 —— 否则 researcher 节点会失败(这是业务依赖,不是迁移问题,仍然算「跑通」)。 运行: cd offline-backend-20260512/backend PYTHONPATH=. uv run python scripts/probe_ai_writing_assistant.py 通过条件(与 ai-writing-优化.md §阶段 0 §通过条件 对齐): - assistants 列表里能看到 ai_writing_probe - 跑到 pause_material_confirm 能拿到 interrupt 事件 - Command(resume={"action": "confirm", ...}) 能成功续跑 - 重建客户端后 client.threads.get_state(thread_id) 仍能拿到上一阶段 state """ from __future__ import annotations import asyncio import os import sys from typing import Any try: from langgraph_sdk import get_client except ImportError: print("ERROR: langgraph-sdk 未安装。run: uv pip install langgraph-sdk", file=sys.stderr) sys.exit(1) ASSISTANT_ID = "ai_writing" DEFAULT_URL = os.environ.get("DEERFLOW_LANGGRAPH_URL", "http://localhost:8001/api/langgraph") def _banner(text: str) -> None: print(f"\n{'=' * 8} {text} {'=' * 8}") def _summarize_event(event: Any) -> str: """裁剪日志,避免一屏全是大段 markdown / 素材列表。""" s = repr(event) return s if len(s) <= 240 else s[:240] + "…" async def _stream_until_interrupt( client, thread_id: str, *, input_payload: dict | None = None, command: dict | None = None ) -> dict | None: """跑一次 runs.stream,捕获第一个 interrupt(含 pause payload)后返回,否则返回 None。 收到 interrupt 时把 payload 一并打出来,便于人工核对 4 个 pause 点的 schema。 """ interrupt_payload: dict | None = None kwargs: dict[str, Any] = { "stream_mode": ["values", "messages-tuple", "custom", "updates"], } if input_payload is not None: kwargs["input"] = input_payload if command is not None: kwargs["command"] = command async for chunk in client.runs.stream(thread_id, ASSISTANT_ID, **kwargs): # langgraph-sdk 把不同 mode 的事件混在同一 stream 里,按 chunk.event 分类。 event_name = getattr(chunk, "event", None) or ( chunk.get("event") if isinstance(chunk, dict) else None ) data = getattr(chunk, "data", None) or ( chunk.get("data") if isinstance(chunk, dict) else None ) print(f"[stream] event={event_name} data={_summarize_event(data)}") # SDK 在遇到 interrupt() 时会发一个特殊事件 —— 不同版本字段名不同,做下兼容。 if event_name == "updates" and isinstance(data, dict): for _node, node_update in data.items(): if isinstance(node_update, dict) and "__interrupt__" in node_update: interrupt_payload = node_update["__interrupt__"] if event_name in {"interrupt", "__interrupt__"}: interrupt_payload = data break return interrupt_payload async def main() -> int: print(f"LangGraph URL = {DEFAULT_URL}") client = get_client(url=DEFAULT_URL) # ── Step 1: 确认 assistant 已注册 ────────────────────────────────────────── _banner("Step 1: 列 assistant,确认 probe 已注册") assistants = await client.assistants.search() names = [a.get("name") or a.get("graph_id") for a in assistants] print(f"已注册 assistants: {names}") if not any(ASSISTANT_ID in (n or "") for n in names): print(f"FAIL: 未找到 {ASSISTANT_ID},请检查 langgraph.json 注册是否生效。") return 1 # ── Step 2: 创建 thread 并触发第一次 run ────────────────────────────────── _banner("Step 2: 创建 thread + 启动写作(应停在 pause_material)") thread = await client.threads.create() thread_id = thread["thread_id"] print(f"thread_id = {thread_id}") start_input = { "user_intent": "写一篇关于 LangGraph 多 worker 部署最佳实践的分析报道,目标读者是后端工程师,约 800 字", "article_type": "analysis", "target_audience": "后端工程师", "word_count_target": 800, "keyword_count": 4, "writing_mode": "loose", "max_revisions": 2, "status": "researching", "drafts": [], "progress_events": [], } payload = await _stream_until_interrupt(client, thread_id, input_payload=start_input) if not payload: print("FAIL: 未收到 interrupt 事件 —— pause_material 可能没停住,或 researcher 抛错。") return 2 print(f"\nPAUSE-1 payload = {_summarize_event(payload)}") # ── Step 3: 确认素材,应停在 pause_outline ──────────────────────────────── _banner("Step 3: Command(resume) 确认素材(应停在 pause_outline)") payload = await _stream_until_interrupt( client, thread_id, command={"resume": {"action": "confirm"}}, # 不指定 approvedMaterialIds → 默认全部通过 ) if not payload: print("FAIL: 未停在 pause_outline。") return 3 print(f"\nPAUSE-2 payload = {_summarize_event(payload)}") # ── Step 4: 确认大纲,应停在 pause_draft 或 pause_section_help ─────────── _banner("Step 4: Command(resume) 确认大纲(应停在 pause_draft 或 pause_section_help)") payload = await _stream_until_interrupt( client, thread_id, command={"resume": {"action": "confirm"}}, ) if not payload: print("FAIL: 未停在 pause_draft / pause_section_help。") return 4 print(f"\nPAUSE-3 payload = {_summarize_event(payload)}") # ── Step 5: 定稿 ───────────────────────────────────────────────────────── _banner("Step 5: Command(resume) finalize 草稿(应跑到 END)") payload = await _stream_until_interrupt( client, thread_id, command={"resume": {"action": "finalize"}}, ) if payload: print(f"NOTE: 收到额外 interrupt(可能编辑审核打回) payload = {_summarize_event(payload)}") print("尝试强制定稿…") payload = await _stream_until_interrupt( client, thread_id, command={"resume": {"action": "force_finalize", "overrideVerdict": True}}, ) if payload: print("WARN: 强制定稿后仍有 interrupt,graph 路由可能有意外分支。") # ── Step 6: 验证 state 持久化 ──────────────────────────────────────────── _banner("Step 6: 重新建客户端 client.threads.get_state(thread_id) 应能拿到最终 state") client2 = get_client(url=DEFAULT_URL) state = await client2.threads.get_state(thread_id) values = state.get("values") or {} final_status = values.get("status") has_draft = bool(values.get("current_draft")) print(f"final status = {final_status}, has_draft = {has_draft}, " f"next nodes = {state.get('next')}") if final_status not in {"done", "writing_draft", "revising"}: print("WARN: 最终 status 不是 done,请人工确认上面流程是否正常。") print("\nPROBE OK —— 阶段 0 验证通过的核心条件已成立(interrupt + resume + state 持久化)。") return 0 if __name__ == "__main__": sys.exit(asyncio.run(main()))