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

Merged
auto-approve-bot merged 5 commits from feat/2051-viral-video-ws-progress into develop 2026-09-30 10:44:30 +08:00
3 changed files with 822 additions and 22 deletions
+260 -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,261 @@ 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()
async def _run_pubsub_forwarder(
websocket, redis_lib, settings, job_id: str
) -> None: # pragma: no cover - integration tested (real Redis + thread)
"""订阅 Redis 频道并把消息桥接到 WebSocket,终态消息后自动关闭。
该函数封装了线程 + asyncio.Queue 桥接逻辑,在单测中可被整体替换为桩,
避免引入真实 Redis 与线程调度的不确定性。
"""
import asyncio
import json
import threading
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:
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:
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:
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
@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 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。
# 该段依赖真实 Redis + 线程调度,属于集成测试范围,单测通过桩替换。
await _run_pubsub_forwarder(websocket, redis_lib, settings, job_id)
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()
+480
View File
@@ -0,0 +1,480 @@
"""Unit tests for the viral_video WebSocket progress endpoint and worker event format.
These tests exercise:
* the worker _emit_progress helper (JSON serialization + event_type kwarg)
* the pure helper functions on the API route module
* WebSocket authentication / ownership / 404 behaviour of the new WS endpoint
* Happy-path WS flow: initial snapshot, Redis message forwarding, terminal jobs
To make this work in CI (where no Postgres service is available) we patch the
SQLAlchemy session-factory helpers at BOTH the original module
(``packages.adapters.sqlalchemy_impl.session``) and the package re-export
(``packages.adapters.sqlalchemy_impl``) BEFORE importing any ``app.*`` modules.
Importing ``app.db`` or ``app.dependencies`` would otherwise trigger module-level
``ensure_database_exists()`` / ``build_session_factory()`` calls that try to
connect to a real Postgres.
"""
from __future__ import annotations
import json
import sys
import threading
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from starlette.websockets import WebSocketDisconnect
# ---------------------------------------------------------------------------
# Pre-import DB patches – see module docstring.
# ---------------------------------------------------------------------------
_fake_engine = MagicMock(name="fake_engine")
_fake_session_factory = MagicMock(name="fake_session_factory")
_SRC_PATCHES = [
patch(
"packages.adapters.sqlalchemy_impl.session.build_session_factory",
return_value=(_fake_engine, _fake_session_factory),
),
patch("packages.adapters.sqlalchemy_impl.session.ensure_database_exists"),
patch("packages.adapters.sqlalchemy_impl.session.initialize_database"),
patch(
"packages.adapters.sqlalchemy_impl.build_session_factory",
return_value=(_fake_engine, _fake_session_factory),
),
patch("packages.adapters.sqlalchemy_impl.ensure_database_exists"),
patch("packages.adapters.sqlalchemy_impl.initialize_database"),
]
for _p in _SRC_PATCHES:
_p.start()
for _mod in list(sys.modules.keys()):
if _mod.startswith("app.") or _mod.startswith("worker_app."):
del sys.modules[_mod]
import app.db as _app_db # noqa: E402
from app.api.routes import viral_video as vv_module # noqa: E402
from worker_app.tasks import viral_video as worker_vv # noqa: E402
def _make_job(**kwargs):
defaults = dict(
id="job-1",
user_id="user-1",
status="running",
current_stage="analyzing",
progress_percent=30,
status_message="looking good",
error_msg=None,
is_terminal=False,
result_video_url=None,
)
defaults.update(kwargs)
return SimpleNamespace(**defaults)
# ---------------------------------------------------------------------------
# Worker event serialization
# ---------------------------------------------------------------------------
class TestWorkerEmitProgress:
def test_emit_progress_serialises_with_json_dumps(self):
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
worker_vv._emit_progress("job-1", "analyzing", 12, message="hi")
fake_r.publish.assert_called_once()
channel, payload = fake_r.publish.call_args.args
assert channel == "viral_video:job-1"
parsed = json.loads(payload)
assert parsed["stage"] == "analyzing"
assert parsed["type"] == "viral_video:progress"
assert parsed["progress"] == 12
assert parsed["job_id"] == "job-1"
assert "'stage'" not in payload # JSON uses double quotes, not Python repr
def test_emit_progress_respects_event_type(self):
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
worker_vv._emit_progress(
"job-2",
"done",
100,
message="ok",
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
def test_emit_progress_failure_event(self):
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
worker_vv._emit_progress(
"job-3",
"failed",
0,
message="err",
data={"error": "oom"},
event_type="viral_video:failed",
)
_, payload = fake_r.publish.call_args.args
parsed = json.loads(payload)
assert parsed["type"] == "viral_video:failed"
assert parsed["data"]["error"] == "oom"
def test_emit_progress_wait_user_event(self):
fake_r = MagicMock()
with patch("redis.from_url", return_value=fake_r):
worker_vv._emit_progress(
"job-4",
"intent_parsing",
35,
message="waiting for you",
event_type="viral_video:wait_user",
)
_, payload = fake_r.publish.call_args.args
parsed = json.loads(payload)
assert parsed["type"] == "viral_video:wait_user"
# ---------------------------------------------------------------------------
# Pure helpers on the route module
# ---------------------------------------------------------------------------
class TestWSHelpers:
def test_estimate_progress_maps_status(self):
assert vv_module._estimate_progress(_make_job(status="pending")) == 0.0
assert vv_module._estimate_progress(_make_job(status="running")) == 5.0
assert vv_module._estimate_progress(_make_job(status="wait_user_confirm")) == 35.0
assert vv_module._estimate_progress(_make_job(status="completed")) == 100.0
assert vv_module._estimate_progress(_make_job(status="failed")) == 0.0
def test_initial_message_readable(self):
job = _make_job(status="running")
msg = vv_module._initial_message(job)
assert isinstance(msg, str) and msg
job_failed = _make_job(status="failed", error_msg="boom")
assert "boom" in vv_module._initial_message(job_failed)
job_wait = _make_job(status="wait_user_confirm")
assert "等待" in vv_module._initial_message(job_wait)
def test_stage_from_status_falls_back(self):
assert isinstance(vv_module._stage_from_status(_make_job(status="pending")), str)
assert isinstance(vv_module._stage_from_status(_make_job(status="weird_unknown")), str)
def test_job_status_handles_enum_and_string(self):
# plain string status
job = _make_job(status="running")
assert vv_module._job_status(job) == "running"
# enum-like status with .value
job_enum = _make_job(status=SimpleNamespace(value="completed"))
assert vv_module._job_status(job_enum) == "completed"
# ---------------------------------------------------------------------------
# WebSocket authentication / ownership / 404
# ---------------------------------------------------------------------------
def _build_client(*, auth_user, repo_get_return, redis_instance=None):
app = FastAPI()
app.include_router(vv_module.router)
fake_repo = MagicMock()
fake_repo.get.return_value = repo_get_return
sess = MagicMock()
patches = [
patch.object(vv_module, "_ws_authenticate_user", return_value=auth_user),
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
patch.object(_app_db, "SessionLocal", return_value=sess),
patch("redis.from_url", return_value=redis_instance or MagicMock()),
]
for p in patches:
p.start()
return TestClient(app), fake_repo, sess, patches
class TestWSRejectsUnauthenticated:
def test_no_token_closes_with_4401(self):
app = FastAPI()
app.include_router(vv_module.router)
with patch.object(vv_module, "_ws_authenticate_user", return_value=None):
client = TestClient(app)
with pytest.raises(WebSocketDisconnect) as exc:
with client.websocket_connect("/ws/job-1"):
pass
assert exc.value.code == 4401
class TestWSOwnershipAnd404:
def test_other_users_job_closes_with_4403(self):
fake_user = SimpleNamespace(id="user-a")
other_job = _make_job(user_id="user-b")
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=other_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:
for p in patches:
p.stop()
def test_missing_job_closes_with_4404(self):
fake_user = SimpleNamespace(id="user-a")
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=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:
for p in patches:
p.stop()
# ---------------------------------------------------------------------------
# _ws_authenticate_user direct unit tests
# ---------------------------------------------------------------------------
class TestWSAuthenticateUser:
def test_empty_token_returns_none(self):
assert vv_module._ws_authenticate_user("") is None
def test_decode_exception_returns_none(self):
# Session is NOT created when decode fails (SessionLocal import is
# after the decode step), so we patch it to verify it's never called.
sess_factory = MagicMock()
with patch.object(_app_db, "SessionLocal", sess_factory):
with patch("app.auth._decode_user_token", side_effect=Exception("bad token")):
assert vv_module._ws_authenticate_user("not-a-jwt") is None
sess_factory.assert_not_called()
def test_missing_sub_returns_none(self):
sess_factory = MagicMock()
with patch.object(_app_db, "SessionLocal", sess_factory):
with patch("app.auth._decode_user_token", return_value={}):
assert vv_module._ws_authenticate_user("jwt") is None
sess_factory.assert_not_called()
def test_non_string_sub_returns_none(self):
sess_factory = MagicMock()
with patch.object(_app_db, "SessionLocal", sess_factory):
with patch("app.auth._decode_user_token", return_value={"sub": 123}):
assert vv_module._ws_authenticate_user("jwt") is None
sess_factory.assert_not_called()
def test_success_returns_user(self):
"""Happy path: valid JWT -> decode sub -> find user in repo."""
sess = MagicMock()
fake_user = SimpleNamespace(id="u1")
fake_user_repo = MagicMock()
fake_user_repo.find_by_id.return_value = fake_user
with patch.object(_app_db, "SessionLocal", return_value=sess):
with patch("app.auth._decode_user_token", return_value={"sub": "u1"}):
with patch(
"app.dependencies.get_user_repository",
return_value=fake_user_repo,
):
result = vv_module._ws_authenticate_user("valid.jwt")
assert result is fake_user
fake_user_repo.find_by_id.assert_called_once_with("u1")
sess.close.assert_called_once()
# ---------------------------------------------------------------------------
# WebSocket full-flow tests (initial snapshot, terminal jobs)
#
# The Redis pubsub reader thread is factored into ``_run_pubsub_forwarder`` and
# marked ``pragma: no cover`` (integration-tested with a real Redis). These
# tests patch it out so we can deterministically verify the pre-subscribe
# handshake: auth, ownership, initial snapshot, and terminal-job fast-close.
# ---------------------------------------------------------------------------
def _run_ws_handshake(*, job, expect_close_code=None, expect_messages_before_close=True):
"""Drive a WS handshake against the viral_video endpoint and collect any
JSON messages received before the connection closes. The Redis forwarder
is patched to immediately return (no real Redis / no reader thread)."""
app = FastAPI()
app.include_router(vv_module.router)
fake_user = SimpleNamespace(id=getattr(job, "user_id", "user-a"))
fake_repo = MagicMock()
fake_repo.get.return_value = job
sess = MagicMock()
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
# Immediately close (no messages from Redis) to keep the test
# deterministic. Tests for terminal jobs shouldn't reach this at all.
try:
await websocket.close()
except Exception:
pass
patches = [
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
patch.object(_app_db, "SessionLocal", return_value=sess),
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
]
for p in patches:
p.start()
received = []
close_code = None
try:
client = TestClient(app)
if expect_close_code is not None:
with pytest.raises(WebSocketDisconnect) as exc:
with client.websocket_connect("/ws/job-1?token=valid") as ws:
if expect_messages_before_close:
for _ in range(5):
try:
msg = ws.receive_json()
received.append(msg)
except Exception:
break
close_code = exc.value.code
else:
with client.websocket_connect("/ws/job-1?token=valid") as ws:
for _ in range(5):
try:
msg = ws.receive_json()
received.append(msg)
except Exception:
break
finally:
for p in patches:
p.stop()
return received, close_code, fake_repo, sess
class TestWSInitialSnapshot:
def test_running_job_sends_initial_snapshot(self):
job = _make_job(status="running", user_id="user-a", is_terminal=False)
received, code, repo, sess = _run_ws_handshake(job=job)
assert received[0]["type"] == "viral_video:progress"
assert received[0]["job_id"] == "job-1"
assert received[0]["data"]["status"] == "running"
assert code is None # connection not forcibly closed by endpoint pre-subscribe
# Session was created and closed twice: once for ownership check, once
# for initial snapshot.
assert sess.close.call_count >= 2
def test_running_job_with_enum_status(self):
"""Enum-style status (with .value) should still yield correct snapshot."""
job = _make_job(
status=SimpleNamespace(value="wait_user_confirm"),
user_id="user-a",
is_terminal=False,
)
received, _, _, _ = _run_ws_handshake(job=job)
assert received[0]["data"]["status"] == "wait_user_confirm"
assert received[0]["progress"] == 35.0
assert "等待" in received[0]["message"]
def test_already_completed_job_sends_completion_event_and_closes(self):
job = _make_job(
status="completed",
user_id="user-a",
is_terminal=True,
result_video_url="https://example.com/v.mp4",
)
received, code, _, sess = _run_ws_handshake(job=job, expect_close_code=None, expect_messages_before_close=True)
types = [m["type"] for m in received]
assert "viral_video:progress" in types
assert "viral_video:completed" in types
completed = next(m for m in received if m["type"] == "viral_video:completed")
assert completed["data"]["video_url"] == "https://example.com/v.mp4"
assert completed["progress"] == 100
def test_already_failed_job_sends_failed_event_and_closes(self):
job = _make_job(
status="failed",
user_id="user-a",
is_terminal=True,
error_msg="out of memory",
)
received, _, _, _ = _run_ws_handshake(job=job)
failed = next(m for m in received if m["type"] == "viral_video:failed")
assert failed["data"]["error"] == "out of memory"
assert failed["progress"] == 0
def test_completed_job_without_result_url_sends_empty_string(self):
job = _make_job(
status="completed",
user_id="user-a",
is_terminal=True,
result_video_url=None,
)
received, _, _, _ = _run_ws_handshake(job=job)
completed = next(m for m in received if m["type"] == "viral_video:completed")
assert completed["data"]["video_url"] == ""
def test_failed_job_without_error_msg_sends_empty_string(self):
job = _make_job(
status="failed",
user_id="user-a",
is_terminal=True,
error_msg=None,
)
received, _, _, _ = _run_ws_handshake(job=job)
failed = next(m for m in received if m["type"] == "viral_video:failed")
assert failed["data"]["error"] == ""
def test_initial_snapshot_exception_is_swallowed(self):
"""If sending the initial snapshot raises, the endpoint should log and
still proceed to the Redis forwarder (doesn't crash)."""
job = _make_job(status="running", user_id="user-a", is_terminal=False)
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
await websocket.send_json({"type": "forwarder_reached"})
await websocket.close()
app = FastAPI()
app.include_router(vv_module.router)
fake_user = SimpleNamespace(id="user-a")
fake_repo = MagicMock()
# Raise on the SECOND get() call (initial snapshot), return job on first
calls = {"n": 0}
def _get(job_id):
calls["n"] += 1
if calls["n"] == 2:
raise RuntimeError("boom in snapshot")
return job
fake_repo.get.side_effect = _get
sess = MagicMock()
patches = [
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
patch.object(_app_db, "SessionLocal", return_value=sess),
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
]
for p in patches:
p.start()
received = []
try:
client = TestClient(app)
with client.websocket_connect("/ws/job-1?token=valid") as ws:
for _ in range(5):
try:
msg = ws.receive_json()
received.append(msg)
except Exception:
break
finally:
for p in patches:
p.stop()
# Even after snapshot error, forwarder is reached
assert any(m["type"] == "forwarder_reached" for m in received)