Files
inquiry_robot/inquiry-agent/agent/graph/runtime.py
T

198 lines
6.7 KiB
Python
Raw 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.
"""
Graph 运行时公开 API:start / resume / wake。
本文件职责:
- Worker 进程内持有已编译图;
- 按 thread_id 启停;无 sender_id 拒绝进图;
- wake 只接受结构化 WorkerWakePayload。
禁止:在 HTTP 回调线程调用 invoke/resume(会卡企微);同一 thread_id 写冲突用细粒度锁。
骨架:节点无业务,invoke 会在第一个 interrupt_before 停下。
"""
from __future__ import annotations
import logging
import threading
import uuid
from dataclasses import dataclass, field
from typing import Any, Optional
from agent.config import Settings, get_settings
from agent.graph.builder import compile_inquiry_graph
from agent.graph.checkpoint import CheckpointBundle, create_checkpointer
from agent.graph.wake import WorkerWakePayload
logger = logging.getLogger(__name__)
class GraphIdentityError(ValueError):
"""无 sender_id 或其他身份不满足时抛出。"""
class GraphWakeError(ValueError):
"""结构化 wake 校验失败。"""
@dataclass
class GraphRuntime:
"""
进程级 Graph 句柄。
生命周期:Worker 启动 create(),停止 close()。
并发:不同 thread_id 可并行;同一 thread_id 串行(_thread_locks)。
"""
settings: Settings
compiled: Any
checkpoint: CheckpointBundle
_thread_locks: dict[str, threading.Lock] = field(default_factory=dict)
_locks_guard: threading.Lock = field(default_factory=threading.Lock)
@classmethod
def create(cls, settings: Optional[Settings] = None) -> "GraphRuntime":
"""
编译图并挂 checkpointer。
副作用:可能连接 PG 并 setup checkpoint 表;失败策略见 create_checkpointer。
"""
cfg = settings or get_settings()
bundle = create_checkpointer(cfg)
compiled = compile_inquiry_graph(checkpointer=bundle.checkpointer)
logger.info(
"GraphRuntime 就绪 checkpoint_backend=%s ytd_env=%s",
bundle.backend,
cfg.ytd_env,
)
return cls(settings=cfg, compiled=compiled, checkpoint=bundle)
def close(self) -> None:
"""关闭 checkpointer 连接。"""
self.checkpoint.close()
def _lock_for(self, thread_id: str) -> threading.Lock:
"""取得同一 thread_id 的互斥锁(禁止全局大锁串行所有工单)。"""
with self._locks_guard:
lock = self._thread_locks.get(thread_id)
if lock is None:
lock = threading.Lock()
self._thread_locks[thread_id] = lock
return lock
@staticmethod
def _require_sender_id(sender_id: str) -> str:
sid = (sender_id or "").strip()
if not sid:
raise GraphIdentityError("无 sender_id,拒绝进入 Graph(prompt/09、14)")
return sid
def start_thread(
self,
*,
sender_id: str,
thread_id: Optional[str] = None,
inquiry_no: Optional[str] = None,
) -> dict[str, Any]:
"""
新开一条询价图(新 thread_id)。
调用线程:必须是 Worker HTTP 类槽,禁止企微回调线程。
副作用:写入 checkpoint;骨架会在首个 pause 中断。
返回:含 thread_id 与 invoke 原始结果摘要。
"""
sid = self._require_sender_id(sender_id)
tid = (thread_id or "").strip() or str(uuid.uuid4())
config = {"configurable": {"thread_id": tid}}
initial = {
"thread_id": tid,
"sender_id": sid,
"inquiry_no": inquiry_no,
"wait_version": 0,
"phase": "start",
}
lock = self._lock_for(tid)
with lock:
# 骨架:invoke 至 interrupt_before 第一个停点
result = self.compiled.invoke(initial, config=config)
logger.info("graph.start_thread thread_id=%s sender_id=%s", tid, sid)
return {"thread_id": tid, "state": result, "interrupted": True}
def resume_thread(
self,
*,
thread_id: str,
sender_id: str,
values: Optional[dict[str, Any]] = None,
) -> dict[str, Any]:
"""
从精确断点继续(入站「继续本单」路径的图侧入口)。
无业务路由:values 仅合并进状态;真正 Handler 后续再接。
"""
sid = self._require_sender_id(sender_id)
tid = (thread_id or "").strip()
if not tid:
raise GraphIdentityError("resume 缺少 thread_id")
config = {"configurable": {"thread_id": tid}}
lock = self._lock_for(tid)
with lock:
patch = dict(values or {})
patch["sender_id"] = sid
patch["thread_id"] = tid
# None 表示从 checkpoint 继续;有 patch 时先 update 再跑
if patch:
self.compiled.update_state(config, patch)
result = self.compiled.invoke(None, config=config)
logger.info("graph.resume_thread thread_id=%s", tid)
return {"thread_id": tid, "state": result}
def wake_from_worker(self, payload: WorkerWakePayload) -> dict[str, Any]:
"""
Worker 结构化唤醒(识别/PDF 完成)。
拒绝:缺 inquiry_no/quoteVersion/waitVersion;用户纯文字不得走本 API。
"""
try:
patch = payload.as_state_patch()
except ValueError as exc:
raise GraphWakeError(str(exc)) from exc
tid = payload.thread_id.strip()
config = {"configurable": {"thread_id": tid}}
lock = self._lock_for(tid)
with lock:
self.compiled.update_state(config, patch)
result = self.compiled.invoke(None, config=config)
logger.info(
"graph.wake_from_worker thread_id=%s kind=%s wait_version=%s",
tid,
payload.kind,
payload.wait_version,
)
return {"thread_id": tid, "state": result, "wake": patch["last_wake"]}
_runtime_singleton: Optional[GraphRuntime] = None
_runtime_guard = threading.Lock()
def get_graph_runtime(settings: Optional[Settings] = None) -> GraphRuntime:
"""
进程内单例 GraphRuntime(懒创建)。
供 Worker 与后续 jobs 槽使用;HTTP 进程不应依赖本单例去 invoke。
"""
global _runtime_singleton
with _runtime_guard:
if _runtime_singleton is None:
_runtime_singleton = GraphRuntime.create(settings)
return _runtime_singleton
def shutdown_graph_runtime() -> None:
"""进程退出时关闭单例。"""
global _runtime_singleton
with _runtime_guard:
if _runtime_singleton is not None:
_runtime_singleton.close()
_runtime_singleton = None