feat(api+worker): 爆款视频 WebSocket 实时进度推送端点(#2051) #2103
Executable → Regular
+260
-1
@@ -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, "任务准备中")
|
||||
|
||||
Executable → Regular
+82
-21
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user