198 lines
6.7 KiB
Python
198 lines
6.7 KiB
Python
"""
|
||
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
|