deerflow-code/offline-backend-20260512/backend/scripts/probe_ai_writing_assistant.py
2026-09-07 18:24:55 +08:00

180 lines
8.1 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 写作 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()))