83 lines
2.5 KiB
Python
83 lines
2.5 KiB
Python
"""
|
||
Worker 侧:消费主账/文件完成等结构化 wake。
|
||
|
||
本文件职责:Redis Stream 入队/认领;在 HTTP 类槽内调 GraphRuntime.wake_from_worker。
|
||
禁止:在企微回调线程调用;禁止用户文字冒充本路径。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from typing import Any, Optional
|
||
|
||
from agent.graph.runtime import GraphRuntime
|
||
from agent.graph.wake import WorkerWakePayload
|
||
from agent.redis_coord.client import RedisClient
|
||
from agent.redis_coord.stream_queue import StreamQueue
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
STREAM_WAKE = "stream:wake"
|
||
WAKE_GROUP = "inquiry-wake-workers"
|
||
|
||
|
||
def wake_stream(client: RedisClient) -> StreamQueue:
|
||
"""结构化 wake 专用 Stream。"""
|
||
q = StreamQueue(client=client, stream_suffix=STREAM_WAKE, group=WAKE_GROUP)
|
||
q.ensure_group()
|
||
return q
|
||
|
||
|
||
def enqueue_wake(client: RedisClient, payload: WorkerWakePayload) -> str:
|
||
"""
|
||
入队 wake(HTTP bridge 快回路径)。
|
||
|
||
副作用:XADD;不跑 Graph。
|
||
"""
|
||
payload.validate()
|
||
q = wake_stream(client)
|
||
return q.enqueue(
|
||
payload={
|
||
"thread_id": payload.thread_id,
|
||
"inquiry_no": payload.inquiry_no,
|
||
"quote_version": int(payload.quote_version),
|
||
"wait_version": int(payload.wait_version),
|
||
"kind": payload.kind,
|
||
"detail": payload.detail or {},
|
||
},
|
||
idempotency_key=f"wake:{payload.thread_id}:{payload.wait_version}:{payload.kind}",
|
||
)
|
||
|
||
|
||
def claim_and_apply_wake(
|
||
*,
|
||
queue: StreamQueue,
|
||
graph: GraphRuntime,
|
||
consumer: str,
|
||
) -> bool:
|
||
"""
|
||
认领 1 条并 wake Graph。成功/失败均 ack(骨架防堵死;正式应进死信)。
|
||
|
||
返回:是否领到任务。
|
||
"""
|
||
claimed = queue.claim(consumer=consumer, count=1, block_ms=200)
|
||
if not claimed:
|
||
return False
|
||
entry_id, data = claimed[0]
|
||
try:
|
||
wake = WorkerWakePayload(
|
||
thread_id=str(data.get("thread_id") or ""),
|
||
inquiry_no=str(data.get("inquiry_no") or ""),
|
||
quote_version=int(data.get("quote_version") or 0),
|
||
wait_version=int(data.get("wait_version") or 0),
|
||
kind=str(data.get("kind") or ""),
|
||
detail=dict(data.get("detail") or {}),
|
||
)
|
||
graph.wake_from_worker(wake)
|
||
logger.info("wake 已处理 id=%s thread=%s kind=%s", entry_id, wake.thread_id, wake.kind)
|
||
except Exception: # noqa: BLE001
|
||
logger.exception("wake 处理失败 id=%s", entry_id)
|
||
finally:
|
||
queue.ack(entry_id)
|
||
return True
|