新增了插件相关的完整领域模型、应用服务、基础设施实现,包括: 1. 插件状态、注册模式、来源等基础枚举和数据结构 2. 插件清单解析、发现、加载工具类 3. 插件注册表领域服务和内存存储实现 4. 插件相关的命令、查询、事件定义 5. 插件REST API接口和DTO映射 6. 集成了原有通道适配器到插件系统 7. 新增内置插件注册和自动发现能力
104 lines
3.2 KiB
Python
104 lines
3.2 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
import re
|
|
from uuid import uuid4
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
|
from pydantic import BaseModel, Field
|
|
|
|
from yuxi.channel.container import ChannelContainer, get_channel
|
|
from yuxi.channel.interfaces.rest.auth.depends import channel_auth_depends
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
router = APIRouter()
|
|
|
|
_SESSION_ID_PATTERN = re.compile(r"^[a-zA-Z0-9\-_]{1,128}$")
|
|
|
|
_REDEEM_SCRIPT = """
|
|
local v = redis.call('GET', KEYS[1])
|
|
if v then
|
|
redis.call('DEL', KEYS[1])
|
|
end
|
|
return v
|
|
"""
|
|
|
|
|
|
class SseTicketRequest(BaseModel):
|
|
session_id: str = Field(..., min_length=1, max_length=128)
|
|
|
|
|
|
class SseTicketResponse(BaseModel):
|
|
ticket: str
|
|
expires_in: int = 30
|
|
|
|
|
|
async def _redeem_ticket(channel: ChannelContainer, ticket: str) -> str | None:
|
|
if not channel.cache_port:
|
|
return None
|
|
cache_key = f"sse:ticket:{ticket}"
|
|
try:
|
|
result = await channel.cache_port.eval(
|
|
_REDEEM_SCRIPT,
|
|
keys=[cache_key],
|
|
args=[],
|
|
)
|
|
if isinstance(result, str):
|
|
return result
|
|
if isinstance(result, (list, tuple)) and len(result) > 0:
|
|
return result[0] if isinstance(result[0], str) else None
|
|
return None
|
|
except Exception:
|
|
session_id = await channel.cache_port.get(cache_key)
|
|
if session_id:
|
|
await channel.cache_port.delete(cache_key)
|
|
return session_id
|
|
|
|
|
|
@router.post("/channel/sse/ticket", response_model=SseTicketResponse)
|
|
async def create_sse_ticket(
|
|
req: SseTicketRequest,
|
|
channel: ChannelContainer = Depends(get_channel),
|
|
_: bool = Depends(channel_auth_depends),
|
|
):
|
|
ticket = uuid4().hex
|
|
if channel.cache_port:
|
|
await channel.cache_port.set(
|
|
f"sse:ticket:{ticket}",
|
|
req.session_id,
|
|
ex=30,
|
|
)
|
|
return SseTicketResponse(ticket=ticket, expires_in=30)
|
|
|
|
|
|
@router.get("/channel/sse")
|
|
async def subscribe_sse(
|
|
request: Request,
|
|
ticket: str | None = Query(None, min_length=1, max_length=64),
|
|
token: str | None = Query(None, alias="token", min_length=1),
|
|
session_id: str | None = Query(None, min_length=1),
|
|
channel: ChannelContainer = Depends(get_channel),
|
|
):
|
|
resolved_session_id: str | None = None
|
|
|
|
if ticket:
|
|
resolved_session_id = await _redeem_ticket(channel, ticket)
|
|
if not resolved_session_id:
|
|
raise HTTPException(status_code=401, detail="invalid or expired ticket")
|
|
elif token:
|
|
passed, reason = await channel.auth_service.authenticate(f"Bearer {token}", client_id="mgmt")
|
|
if not passed:
|
|
if "rate limited" in reason:
|
|
raise HTTPException(status_code=429, detail=reason)
|
|
raise HTTPException(status_code=401, detail=reason or "authorization required")
|
|
resolved_session_id = session_id
|
|
else:
|
|
raise HTTPException(status_code=401, detail="ticket or token required")
|
|
|
|
if not resolved_session_id or not _SESSION_ID_PATTERN.match(resolved_session_id):
|
|
raise HTTPException(status_code=400, detail="invalid session_id format")
|
|
|
|
sse_endpoint = channel.require("sse_endpoint")
|
|
return await sse_endpoint.subscribe(resolved_session_id, request)
|