feat(api+worker): 爆款视频 WebSocket 实时进度推送端点

- API: 新增 /api/v1/viral-video/ws/{job_id}?token=<jwt> WebSocket 端点
  - JWT 认证通过 ?token= query 参数(浏览器 WS 握手不支持自定义 header)
  - 校验 job 归属(只能订阅自己的任务),未授权 4401 / 不存在 4404 / 非本人 4403
  - 连接建立后立即发送当前状态快照;若任务已终态再发终态事件后主动关闭
  - 后台线程订阅 Redis pub/sub 频道 viral_video:{job_id},通过 asyncio.Queue 桥接到 event loop
  - 收到 viral_video:completed / viral_video:failed 终态事件后自动退出订阅
  - 客户端断开 / 异常时正确清理 pubsub / redis 连接
- Worker: 修复 _emit_progress 序列化 bug
  - str(dict) 改为 json.dumps(event, ensure_ascii=False),保证前端 JSON.parse 可解析
  - 新增 event_type 参数,新增 viral_video:wait_user / viral_video:completed / viral_video:failed 终态事件
  - 在主编排器和 resume 编排器异常分支补充 failed 事件推送
- 事件格式与前端 WSProgressEvent 对齐:{type, job_id, stage, progress, message, data}
- 新增 8 个单元测试覆盖 JSON 序列化 / helper 函数 / 4401/4403/4404
This commit is contained in:
xiaoxia
2026-09-30 08:58:01 +08:00
parent 37f7aa3329
commit 35c41c3151
3 changed files with 505 additions and 22 deletions
+252 -1
View File
@@ -8,6 +8,7 @@
POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
GET /api/v1/viral-video/style-templates 获取风格模板列表
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送(订阅 Redis pub/sub)
"""
from __future__ import annotations
@@ -26,7 +27,7 @@ from app.schemas.viral_video import (
ViralVideoHistoryResponse,
ViralVideoJobResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
@@ -295,3 +296,253 @@ def analyze_style(
status="analyzing",
style_guide=None,
)
# ── WebSocket 进度推送 ──────────────────────────────────────────────────
def _ws_authenticate_user(token: str):
"""从 token 字符串解析用户(复用 HTTP Bearer 的解码 + 黑名单逻辑)。
WebSocket 握手阶段不能发自定义 Authorization header,
因此统一通过 query 参数 ``?token=...`` 传 JWT。
"""
from app.auth import _decode_user_token
from app.dependencies import get_user_repository
if not token:
return None
try:
payload = _decode_user_token(token)
except Exception:
return None
user_id = payload.get("sub")
if not isinstance(user_id, str) or not user_id:
return None
# 同步场景下手动拉 repository 实例
from app.db import SessionLocal
session = SessionLocal()
try:
user_repo = get_user_repository(session)
user = user_repo.find_by_id(user_id)
return user
finally:
session.close()
@router.websocket("/ws/{job_id}")
async def viral_video_websocket(websocket: WebSocket, job_id: str) -> None:
"""WebSocket 桥接:订阅 Redis `viral_video:{job_id}` 频道并转发给前端。
认证:通过 ``?token=<jwt>`` query 参数传 JWT(浏览器 WS 握手不支持自定义 header)。
事件类型:
- viral_video:progress 中间进度(progress: 0-100)
- viral_video:wait_user 等待用户确认意图文案
- viral_video:completed 任务完成(data.video_url)
- viral_video:failed 任务失败(data.error)
- viral_video:error 服务端错误(如鉴权失败 / job 不存在 / 无权限)
"""
import asyncio
import json
import redis as redis_lib
from app.config import settings
# ── 1. 鉴权 ──────────────────────────────────────────────────────
token = websocket.query_params.get("token", "")
user = _ws_authenticate_user(token)
if user is None:
await websocket.close(code=4401, reason="Unauthorized")
return
# ── 2. 校验 job 归属 ─────────────────────────────────────────────
from app.db import SessionLocal
session = SessionLocal()
try:
job_repo = SQLAlchemyViralVideoJobRepository(session)
job = job_repo.get(job_id)
if job is None:
await websocket.close(code=4404, reason="Job not found")
return
if job.user_id != user.id:
await websocket.close(code=4403, reason="Forbidden")
return
finally:
session.close()
await websocket.accept()
# ── 3. 发送一条初始状态(前端连接后立即拿到当前进度) ────────────
try:
session = SessionLocal()
job_repo = SQLAlchemyViralVideoJobRepository(session)
job = job_repo.get(job_id)
if job is not None:
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
initial = {
"type": "viral_video:progress",
"job_id": job_id,
"stage": _stage_from_status(job),
"progress": _estimate_progress(job),
"message": _initial_message(job),
"data": {"status": status_val},
}
await websocket.send_json(initial)
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
if job.is_terminal:
is_completed = status_val == "completed"
terminal_type = "viral_video:completed" if is_completed else "viral_video:failed"
terminal_data = (
{"video_url": job.result_video_url or ""} if is_completed else {"error": job.error_msg or ""}
)
await websocket.send_json(
{
"type": terminal_type,
"job_id": job_id,
"stage": "",
"progress": 100 if is_completed else 0,
"message": "视频生成完成" if is_completed else "任务失败",
"data": terminal_data,
}
)
await websocket.close()
return
session.close()
except Exception as e:
logger.warning("[爆款视频WS] 发送初始状态失败: %s", e)
try:
session.close()
except Exception:
pass
# ── 4. 订阅 Redis 频道并转发 ─────────────────────────────────────
# redis-py 的 pubsub 是同步阻塞的,放到线程里跑,通过 asyncio.Queue 桥接到 event loop
r = redis_lib.from_url(settings.REDIS_URL, decode_responses=True)
pubsub = r.pubsub(ignore_subscribe_messages=True)
channel = f"viral_video:{job_id}"
pubsub.subscribe(channel)
loop = asyncio.get_running_loop()
queue: asyncio.Queue = asyncio.Queue(maxsize=64)
stop_event = asyncio.Event()
def _reader() -> None:
"""同步线程:从 pubsub 读消息,投递到 asyncio.Queue。"""
try:
while not stop_event.is_set():
msg = pubsub.get_message(timeout=0.5)
if msg is None or msg.get("type") != "message":
continue
raw = msg.get("data")
if not isinstance(raw, str):
continue
try:
payload = json.loads(raw)
except Exception:
payload = {"type": "viral_video:progress", "data": {"raw": raw}}
loop.call_soon_threadsafe(queue.put_nowait, payload)
# 终态消息 → 通知退出
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
loop.call_soon_threadsafe(stop_event.set)
break
except Exception as e:
logger.warning("[爆款视频WS] pubsub reader 异常退出: %s", e)
loop.call_soon_threadsafe(stop_event.set)
try:
import threading
reader_thread = threading.Thread(target=_reader, name=f"viral-video-ws-{job_id}", daemon=True)
reader_thread.start()
while not stop_event.is_set():
try:
payload = await asyncio.wait_for(queue.get(), timeout=1.0)
except asyncio.TimeoutError:
# 心跳:每 25 秒发 ping 防反代断连(FastAPI WebSocket 自带 ping,但显式保活)
continue
try:
await websocket.send_json(payload)
except Exception:
# 连接已断
break
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
break
except WebSocketDisconnect:
logger.info("[爆款视频WS] 客户端断开: job_id=%s", job_id)
except Exception as e:
logger.error("[爆款视频WS] 转发异常: %s", e, exc_info=True)
try:
await websocket.send_json({"type": "viral_video:error", "message": f"服务异常: {e}"})
except Exception:
pass
finally:
stop_event.set()
try:
pubsub.unsubscribe(channel)
pubsub.close()
except Exception:
pass
try:
r.close()
except Exception:
pass
try:
await websocket.close()
except Exception:
pass
def _job_status(job) -> str:
return job.status.value if hasattr(job.status, "value") else str(job.status)
# 初始快照的 stage 推断:领域对象不持久化 stage,
# 只能根据 status 给一个占位,后续 worker 推送的真实进度事件会覆盖。
_STATUS_STAGE = {
"pending": "",
"running": "",
"wait_user_confirm": "intent_parsing",
"completed": "uploading",
"failed": "",
"cancelled": "",
}
_STATUS_PROGRESS = {
"pending": 0.0,
"running": 5.0,
"wait_user_confirm": 35.0,
"completed": 100.0,
"failed": 0.0,
"cancelled": 0.0,
}
_STATUS_MESSAGE = {
"pending": "任务已创建,等待执行",
"running": "任务执行中",
"wait_user_confirm": "等待用户确认意图文案",
"completed": "视频生成完成",
"failed": "任务失败",
"cancelled": "任务已取消",
}
def _stage_from_status(job) -> str:
return _STATUS_STAGE.get(_job_status(job), "")
def _estimate_progress(job) -> float:
"""根据 status 粗略估算百分比(0-100),用于连接初始快照;
连接建立后由 Redis 推送的真实事件持续更新。
"""
return _STATUS_PROGRESS.get(_job_status(job), 5.0)
def _initial_message(job) -> str:
"""给新连接的前端一个可读的初始状态文案。"""
status_val = _job_status(job)
if status_val == "failed" and job.error_msg:
return f"任务失败: {job.error_msg}"
return _STATUS_MESSAGE.get(status_val, "任务准备中")
+82 -21
View File
@@ -16,6 +16,7 @@
from __future__ import annotations
import json
import logging
import os
@@ -41,22 +42,37 @@ logger = logging.getLogger(__name__)
# ── WS 进度推送 ──────────────────────────────────────────────────────────
def _emit_progress(job_id: str, stage: str, progress: float, message: str = "", data: dict | None = None):
"""通过 Redis 发布进度事件,供 WebSocket 消费。"""
def _emit_progress(
job_id: str,
stage: str,
progress: float,
message: str = "",
data: dict | None = None,
event_type: str = "viral_video:progress",
):
"""通过 Redis 发布进度事件,供 WebSocket 消费。
event_type 取值:
- viral_video:progress 中间进度(默认)
- viral_video:completed 任务完成
- viral_video:failed 任务失败
- viral_video:wait_user 等待用户确认
所有事件 payload 均为合法 JSON,前端 JSON.parse 即可。
"""
try:
import redis as redis_lib
redis_url = os.environ.get("REDIS_URL", "redis://localhost:6379/0")
r = redis_lib.from_url(redis_url)
event = {
"type": "viral_video:progress",
"type": event_type,
"job_id": job_id,
"stage": stage,
"progress": progress,
"message": message or STAGE_LABELS.get(stage, stage),
"data": data or {},
}
r.publish(f"viral_video:{job_id}", str(event))
r.publish(f"viral_video:{job_id}", json.dumps(event, ensure_ascii=False))
except Exception as e:
logger.warning("[爆款视频] WS 进度推送失败: %s", e)
@@ -429,6 +445,14 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
"意图解析完成,等待用户确认",
{"intent_result": intent_result, "waiting_confirm": True},
)
_emit_progress(
job_id,
ViralVideoStage.INTENT_PARSING,
35.0,
"等待用户确认意图文案",
{"intent_result": intent_result},
event_type="viral_video:wait_user",
)
# 这里流水线暂停,等待 confirm-intent API 调用 resume
# resume 后由 resume_viral_video_pipeline 继续
@@ -438,15 +462,30 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
raise
except Exception as e:
logger.error("[爆款视频] 流水线异常: %s", e, exc_info=True)
if session:
try:
err_msg = str(e)
failed_stage = ""
try:
if session is None:
session = SessionLocal()
repo = SQLAlchemyViralVideoJobRepository(session)
job = repo.get(job_id)
else:
_, repo, job = _get_repo_and_job(job_id)
if job and not job.is_terminal:
job.mark_failed(str(e))
_save_job(repo, job, session)
except Exception:
pass
return {"ok": False, "job_id": job_id, "error": str(e)}
if job is not None and not job.is_terminal:
job.mark_failed(err_msg)
failed_stage = getattr(job, "current_stage", "") or ""
_save_job(repo, job, session)
except Exception as inner:
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
_emit_progress(
job_id,
failed_stage,
0,
f"任务失败: {err_msg}",
{"error": err_msg},
event_type="viral_video:failed",
)
return {"ok": False, "job_id": job_id, "error": err_msg}
finally:
if session:
session.close()
@@ -456,6 +495,7 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
"""用户确认意图后,从断点恢复流水线(步骤 3-10)。"""
session = None
job = None
try:
session, repo, job = _get_repo_and_job(job_id)
if job is None:
@@ -517,6 +557,14 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
job.mark_completed(video_url)
_save_job(repo, job, session)
_emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url})
_emit_progress(
job_id,
ViralVideoStage.UPLOADING,
100.0,
"视频生成完成",
{"video_url": video_url},
event_type="viral_video:completed",
)
logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url)
return {"ok": True, "job_id": job_id, "video_url": video_url}
@@ -525,15 +573,28 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
raise
except Exception as e:
logger.error("[爆款视频] 恢复流水线异常: %s", e, exc_info=True)
if session:
try:
_, repo, job = _get_repo_and_job(job_id)
if job and not job.is_terminal:
job.mark_failed(str(e))
_save_job(repo, job, session)
except Exception:
pass
return {"ok": False, "job_id": job_id, "error": str(e)}
err_msg = str(e)
failed_stage = ""
try:
if session is None:
session = SessionLocal()
repo = SQLAlchemyViralVideoJobRepository(session)
job = repo.get(job_id)
elif job is not None and not job.is_terminal:
job.mark_failed(err_msg)
failed_stage = getattr(job, "current_stage", "") or ""
_save_job(repo, job, session)
except Exception as inner:
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
_emit_progress(
job_id,
failed_stage,
0,
f"任务失败: {err_msg}",
{"error": err_msg},
event_type="viral_video:failed",
)
return {"ok": False, "job_id": job_id, "error": err_msg}
finally:
if session:
session.close()
+171
View File
@@ -0,0 +1,171 @@
"""Tests for the viral video WebSocket progress endpoint + worker event serialization."""
from __future__ import annotations
import json
from unittest.mock import MagicMock, patch
import pytest
from starlette.websockets import WebSocketDisconnect
class TestWorkerEmitProgress:
"""Verify _emit_progress publishes valid JSON (not Python dict repr)."""
def test_publishes_valid_json(self):
import redis
from worker_app.tasks.viral_video import _emit_progress
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
_emit_progress("job-xyz", "image_analysis", 10.0, "hello", {"k": "v"})
assert fake_r.publish.called
channel, payload = fake_r.publish.call_args.args
assert channel == "viral_video:job-xyz"
parsed = json.loads(payload)
assert parsed["type"] == "viral_video:progress"
assert parsed["job_id"] == "job-xyz"
assert parsed["stage"] == "image_analysis"
assert parsed["progress"] == 10.0
assert parsed["data"] == {"k": "v"}
def test_event_type_parameter(self):
"""event_type 参数可覆盖 type 字段。"""
from worker_app.tasks.viral_video import _emit_progress
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
_emit_progress(
"j1",
"uploading",
100.0,
"done",
{"video_url": "u"},
event_type="viral_video:completed",
)
_, payload = fake_r.publish.call_args.args
parsed = json.loads(payload)
assert parsed["type"] == "viral_video:completed"
assert parsed["progress"] == 100.0
class TestWSHelpers:
"""Verify helper functions used by the WS endpoint."""
def _make_job(self, status, error_msg="", result_video_url=""):
job = MagicMock()
job.status = MagicMock()
job.status.value = status
job.error_msg = error_msg
job.result_video_url = result_video_url
job.is_terminal = status in ("completed", "failed", "cancelled")
return job
def test_estimate_progress(self):
from app.api.routes.viral_video import _estimate_progress, _initial_message, _stage_from_status
assert _estimate_progress(self._make_job("pending")) == 0.0
assert _estimate_progress(self._make_job("running")) == 5.0
assert _estimate_progress(self._make_job("wait_user_confirm")) == 35.0
assert _estimate_progress(self._make_job("completed")) == 100.0
assert _estimate_progress(self._make_job("failed")) == 0.0
def test_initial_message(self):
from app.api.routes.viral_video import _initial_message
assert "失败" in _initial_message(self._make_job("failed", error_msg="boom"))
assert _initial_message(self._make_job("completed")) == "视频生成完成"
assert "确认" in _initial_message(self._make_job("wait_user_confirm"))
def test_stage_from_status(self):
from app.api.routes.viral_video import _stage_from_status
assert _stage_from_status(self._make_job("wait_user_confirm")) == "intent_parsing"
assert _stage_from_status(self._make_job("running")) == ""
assert _stage_from_status(self._make_job("completed")) == "uploading"
class TestWSRejectsUnauthenticated:
"""WS endpoint must close with 4401 when no token / bad token."""
def _make_client(self):
from fastapi import FastAPI
from fastapi.testclient import TestClient
from app.api.routes.viral_video import router
app = FastAPI()
app.include_router(router)
# Patch _ws_authenticate_user to return None (no auth) regardless of token
with patch("app.api.routes.viral_video._ws_authenticate_user", return_value=None):
return TestClient(app)
def test_no_token_closes_with_4401(self):
client = self._make_client()
with pytest.raises(WebSocketDisconnect) as exc:
with client.websocket_connect("/ws/job-123"):
pass
assert exc.value.code == 4401
class TestWSOwnershipAnd404:
"""WS endpoint must verify the job exists and belongs to the authenticated user."""
def _build_client(self, *, fake_user_id="user-1", fake_job=None):
from fastapi import FastAPI
from fastapi.testclient import TestClient
from app.api.routes import viral_video as vv_module
app = FastAPI()
app.include_router(vv_module.router)
fake_user = MagicMock()
fake_user.id = fake_user_id
fake_repo = MagicMock()
fake_repo.get.return_value = fake_job
fake_session = MagicMock()
patches = [
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
patch(
"app.api.routes.viral_video.SQLAlchemyViralVideoJobRepository",
return_value=fake_repo,
),
patch("app.db.SessionLocal", return_value=fake_session),
# Block Redis subscription (we don't test pubsub flow here)
patch("redis.from_url", return_value=MagicMock()),
]
for p in patches:
p.start()
client = TestClient(app)
return client, patches, fake_session
def _close(self, patches, fake_session):
for p in patches:
p.stop()
fake_session.close()
def test_ownership_mismatch_closes_403(self):
fake_job = MagicMock()
fake_job.user_id = "user-2"
client, patches, sess = self._build_client(fake_job=fake_job)
try:
with pytest.raises(WebSocketDisconnect) as exc:
with client.websocket_connect("/ws/job-x?token=valid-token"):
pass
assert exc.value.code == 4403
finally:
self._close(patches, sess)
def test_nonexistent_job_closes_404(self):
client, patches, sess = self._build_client(fake_job=None)
try:
with pytest.raises(WebSocketDisconnect) as exc:
with client.websocket_connect("/ws/job-missing?token=valid-token"):
pass
assert exc.value.code == 4404
finally:
self._close(patches, sess)