ForcePilot/backend/package/yuxi/channel/interfaces/sse/endpoint.py
Kris c61d5f0163 feat: 完成通道服务多轮功能迭代
本次提交完成了一系列核心功能迭代与优化:
1.  新增并完善了多个领域模型与端口定义,补充了`__all__`导出规范
2.  优化了会话、绑定、出箱等模块的数据模型,修复了时间字段类型不一致问题
3.  新增了代理ID解析、缓存发布等接口,扩展了系统能力
4.  重构了去重中间件逻辑,优化了空内容校验规则
5.  新增了认证中间件的匿名访问支持,完善了鉴权流程
6.  优化了SSE连接管理,增加了单会话连接上限限制
7.  重构了消息日志与仓储相关代码,将数据类迁移至对应模型目录
8.  新增了重复绑定校验、绑定更新接口,完善了绑定服务逻辑
9.  优化了健康检查逻辑,新增了环境变量控制启动时间线展示
10. 重构了出箱重试工作线程,使用缓存端口替代直接redis操作,新增了消息处理标记逻辑
11. 完善了飞书、Web、钩子等通道的翻译器逻辑,补充了账户ID传递
12. 新增了多种自定义异常类型,优化了异常映射与错误处理流程
13. 完善了配置热重载逻辑,同步认证凭证与校验器配置
14. 重构了Redis缓存实现,增加了异常捕获与包装
2026-05-31 21:42:03 +08:00

166 lines
6.1 KiB
Python

from __future__ import annotations
import asyncio
import json
import logging
import time
from dataclasses import dataclass
from fastapi import HTTPException, Request
from fastapi.responses import StreamingResponse
from yuxi.channel.domain.port.metrics_port import MetricsPort
from yuxi.channel.domain.port.sse_push_port import SsePushPort
logger = logging.getLogger(__name__)
@dataclass
class _SseConnection:
session_id: str
queue: asyncio.Queue
last_active_at: float = 0.0
class SseEndpoint(SsePushPort):
STALE_THRESHOLD_SECONDS = 300
CLEANUP_INTERVAL_SECONDS = 60
MAX_CONNECTIONS_PER_SESSION = 5
def __init__(self, *, max_connections: int = 1000, metrics: MetricsPort | None = None) -> None:
self._connections: dict[str, list[_SseConnection]] = {}
self._max_connections = max_connections
self._cleanup_task: asyncio.Task | None = None
self._metrics = metrics
@property
def connection_count(self) -> int:
return sum(len(conns) for conns in self._connections.values())
async def start(self) -> None:
logger.info("SSE endpoint started, max_connections=%d", self._max_connections)
self._cleanup_task = asyncio.create_task(self._periodic_cleanup())
async def stop(self) -> None:
if self._cleanup_task:
self._cleanup_task.cancel()
try:
await self._cleanup_task
except asyncio.CancelledError:
pass
self._cleanup_task = None
for conns in self._connections.values():
for conn in conns:
try:
conn.queue.put_nowait({"type": "shutdown", "reason": "server_stopping"})
except asyncio.QueueFull:
pass
await asyncio.sleep(1.0)
self._connections.clear()
async def subscribe(self, session_id: str, request: Request) -> StreamingResponse:
if self.connection_count >= self._max_connections:
await self._cleanup_stale()
if self.connection_count >= self._max_connections:
raise HTTPException(status_code=503, detail="SSE connection limit reached")
session_conns = self._connections.get(session_id, [])
if len(session_conns) >= self.MAX_CONNECTIONS_PER_SESSION:
raise HTTPException(status_code=429, detail="too many SSE connections for this session")
now = time.monotonic()
queue: asyncio.Queue = asyncio.Queue(maxsize=256)
conn = _SseConnection(session_id=session_id, queue=queue, last_active_at=now)
self._connections.setdefault(session_id, []).append(conn)
await self._update_connection_count()
async def event_stream():
try:
yield f"event: connected\ndata: {json.dumps({'session_id': session_id}, ensure_ascii=False)}\n\n"
while True:
if await request.is_disconnected():
break
try:
data = await asyncio.wait_for(queue.get(), timeout=30.0)
conn.last_active_at = time.monotonic()
yield f"event: message\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
except TimeoutError:
conn.last_active_at = time.monotonic()
yield "event: ping\ndata: {}\n\n"
finally:
conns = self._connections.get(session_id)
if conns:
try:
conns.remove(conn)
except ValueError:
pass
if not conns:
self._connections.pop(session_id, None)
await self._update_connection_count()
return StreamingResponse(event_stream(), media_type="text/event-stream")
async def push_event(self, session_id: str, data: dict) -> bool:
conns = self._connections.get(session_id)
if not conns:
return False
delivered = False
for conn in conns:
try:
conn.queue.put_nowait(data)
conn.last_active_at = time.monotonic()
delivered = True
except asyncio.QueueFull:
pass
return delivered
async def broadcast_shutdown(self) -> None:
for conns in self._connections.values():
for conn in conns:
try:
conn.queue.put_nowait({"type": "shutdown", "reason": "gateway_shutting_down"})
except asyncio.QueueFull:
pass
async def _periodic_cleanup(self) -> None:
try:
while True:
await asyncio.sleep(self.CLEANUP_INTERVAL_SECONDS)
await self._cleanup_stale()
except asyncio.CancelledError:
pass
async def _update_connection_count(self) -> None:
if self._metrics:
try:
await self._metrics.set_sse_connections(self.connection_count)
except Exception:
pass
async def _cleanup_stale(self) -> None:
now = time.monotonic()
for sid in list(self._connections):
conns = self._connections[sid]
stale = [c for c in conns if now - c.last_active_at > self.STALE_THRESHOLD_SECONDS]
for c in stale:
conns.remove(c)
if not conns:
del self._connections[sid]
if self.connection_count >= self._max_connections:
all_conns: list[tuple[str, _SseConnection]] = []
for sid, conns in self._connections.items():
for c in conns:
all_conns.append((sid, c))
all_conns.sort(key=lambda x: x[1].last_active_at)
evict_count = len(all_conns) // 4
for sid, conn in all_conns[:evict_count]:
conns = self._connections.get(sid)
if conns:
try:
conns.remove(conn)
except ValueError:
pass
if not conns:
self._connections.pop(sid, None)