180 lines
8.1 KiB
Python
180 lines
8.1 KiB
Python
"""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()))
|