"""任务完成事件通知服务。 设计: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 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"