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