ADK-gateway/backend/app/services/notifier.py
2026-08-05 16:33:27 +08:00

97 lines
3.6 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.

"""任务完成事件通知服务。
设计CLI 提交任务时可携带 cli_session_id并通过 SSE 长连接订阅
GET /api/cli/events?cli_session_id=xxx 事件流。任务到达终态
success / failed含取消、超时网关通过 Redis Pub/Sub 向
对应会话推送事件,避免 CLI 轮询。
为保证订阅晚于任务完成也不丢事件publish 时同时写入一份
回放缓存Redis List保留最近 N 条SSE 连接建立后先回放再实时推送。
"""
import json
import logging
from typing import Any
import redis.asyncio as aioredis
from app.models.schemas import TaskInfo
logger = logging.getLogger(__name__)
CHANNEL_PREFIX = "task:event:cli:" # Redis Pub/Sub channel
REPLAY_KEY_PREFIX = "task:event:replay:" # 回放缓存 List key
REPLAY_TTL = 600
REPLAY_MAX = 200
class TaskNotifier:
def __init__(self, redis: aioredis.Redis):
self.redis = redis
@staticmethod
def channel(cli_session_id: str) -> str:
return f"{CHANNEL_PREFIX}{cli_session_id}"
@staticmethod
def _replay_key(cli_session_id: str) -> str:
return f"{REPLAY_KEY_PREFIX}{cli_session_id}"
# ---------- 发布 ----------
async def publish_task_done(self, task: TaskInfo) -> None:
"""任务到达终态时调用:写回放缓存并广播到该 CLI 会话。"""
if not task.cli_session_id:
return
payload = json.dumps(
{
"event": "task_done",
"request_id": task.request_id,
"status": task.status.value,
"result": task.result,
"error_info": task.error_info,
},
ensure_ascii=False,
)
key = self._replay_key(task.cli_session_id)
channel = self.channel(task.cli_session_id)
await self.redis.lpush(key, payload)
await self.redis.ltrim(key, 0, REPLAY_MAX - 1)
await self.redis.expire(key, REPLAY_TTL)
await self.redis.publish(channel, payload)
logger.info("notify task done session=%s request=%s status=%s",
task.cli_session_id, task.request_id, task.status.value)
async def publish_task_stop(self, cli_session_id: str, request_id: str) -> None:
"""取消任务时向 agent 下发停止指令(复用同一 CLI 会话通道event=task_stop
允许未携带 cli_session_id 时跳过(此时无法通过 SSE 通知,只能等待 agent 超时)。
"""
if not cli_session_id:
return
payload = json.dumps(
{
"event": "task_stop",
"request_id": request_id,
},
ensure_ascii=False,
)
channel = self.channel(cli_session_id)
await self.redis.publish(channel, payload)
logger.info("notify task stop session=%s request=%s", cli_session_id, request_id)
# ---------- 消费 ----------
async def replay(self, cli_session_id: str) -> list[str]:
"""返回该会话的历史完成事件(新 → 旧),供 SSE 连接回放。"""
items = await self.redis.lrange(self._replay_key(cli_session_id), 0, -1)
return [s.decode() if isinstance(s, bytes) else s for s in items]
async def drain_replay(self, cli_session_id: str) -> None:
"""SSE 连接回放结束后清空回放缓存(已消费)。"""
await self.redis.delete(self._replay_key(cli_session_id))
def format_sse(data: str | dict[str, Any]) -> str:
"""格式化为 SSE 帧。"""
if isinstance(data, dict):
data = json.dumps(data, ensure_ascii=False)
return f"data: {data}\n\n"