88 lines
3.0 KiB
Python
88 lines
3.0 KiB
Python
"""CLI 接口:任务提交/查询/取消/完成事件订阅。"""
|
||
from fastapi import APIRouter, Depends, Request
|
||
from fastapi.responses import StreamingResponse
|
||
|
||
from app.models.schemas import TaskInfo, TaskSubmit
|
||
from app.services.notifier import TaskNotifier
|
||
from app.services.security import require_auth
|
||
from app.services.task_service import TaskService
|
||
|
||
router = APIRouter()
|
||
|
||
|
||
def get_task_service(request: Request) -> TaskService:
|
||
return TaskService(request.app.state.redis)
|
||
|
||
|
||
def get_notifier(request: Request) -> TaskNotifier:
|
||
return TaskNotifier(request.app.state.redis)
|
||
|
||
|
||
@router.get("/events", summary="订阅任务完成事件(SSE 长连接)")
|
||
async def subscribe_events(
|
||
request: Request,
|
||
cli_session_id: str,
|
||
auth: str = "",
|
||
notifier: TaskNotifier = Depends(get_notifier),
|
||
):
|
||
"""任务到达终态时网关主动推送完成事件,无需轮询。
|
||
|
||
返回 SSE 流(text/event-stream),每帧形如:
|
||
data: {"event": "task_done", "request_id": "...", "status": "success", "result": {...}}
|
||
"""
|
||
require_auth(auth)
|
||
|
||
channel = notifier.channel(cli_session_id)
|
||
pubsub = request.app.state.redis.pubsub()
|
||
await pubsub.subscribe(channel)
|
||
replay = await notifier.replay(cli_session_id)
|
||
|
||
async def event_gen():
|
||
try:
|
||
# 先回放已完成的(连接晚于任务完成的场景)
|
||
for evt in replay:
|
||
yield f"data: {evt}\n\n"
|
||
# 再实时监听
|
||
async for message in pubsub.listen():
|
||
if await request.is_disconnected():
|
||
break
|
||
if message["type"] != "message":
|
||
continue
|
||
data = message["data"]
|
||
if isinstance(data, bytes):
|
||
data = data.decode()
|
||
yield f"data: {data}\n\n"
|
||
finally:
|
||
await pubsub.unsubscribe(channel)
|
||
await pubsub.aclose()
|
||
|
||
return StreamingResponse(
|
||
event_gen(),
|
||
media_type="text/event-stream",
|
||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||
)
|
||
|
||
|
||
@router.post("/tasks", response_model=TaskInfo, summary="提交任务")
|
||
async def submit_task(body: TaskSubmit, svc: TaskService = Depends(get_task_service)):
|
||
require_auth(body.auth)
|
||
return await svc.submit(body)
|
||
|
||
|
||
@router.get("/tasks/{request_id}", response_model=TaskInfo, summary="查询任务")
|
||
async def get_task(request_id: str, svc: TaskService = Depends(get_task_service)):
|
||
from fastapi import HTTPException
|
||
|
||
task = await svc.get(request_id)
|
||
if not task:
|
||
raise HTTPException(status_code=404, detail="task not found")
|
||
return task
|
||
|
||
|
||
@router.post("/tasks/{request_id}/cancel", summary="取消任务")
|
||
async def cancel_task(request_id: str, svc: TaskService = Depends(get_task_service)):
|
||
from fastapi import HTTPException
|
||
|
||
if not await svc.cancel(request_id):
|
||
raise HTTPException(status_code=400, detail="cannot cancel task")
|
||
return {"ok": True, "request_id": request_id} |