Files
xiaoxia-saas/tests/unit/test_viral_video_ws.py
T
xiaoxia bf9249da19
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 2s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 3s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 4s
CI/CD Pipeline / Frontend Lint (push) Has been skipped
CI/CD Pipeline / PR Build Web Image (push) Has been skipped
CI/CD Pipeline / PR Build Worker Image (push) Has been skipped
CI/CD Pipeline / PR Build API Image (push) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Check push changed paths (push) Successful in 18s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m1s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m15s
CI/CD Pipeline / Build Staging API Image (push) Successful in 1m0s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (push) Successful in 31s
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m37s
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 32s
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m16s
CI/CD Pipeline / Frontend Unit Tests (push) Failing after 2m49s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m20s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 1m51s
CI/CD Pipeline / ACR Image Cleanup (push) Successful in 1m49s
CI/CD Pipeline / Integration Tests (push) Successful in 6m46s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 2m57s
AI Code Review / AI Code Review (pull_request) Successful in 7m7s
CI/CD Pipeline / Validate - Style (push) Successful in 7m38s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 3m56s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 8m38s
CI/CD Pipeline / Unit Tests (push) Successful in 18m41s
CI/CD Pipeline / Validate - Security (push) Successful in 19m30s
CI/CD Pipeline / Build Production API Image (push) Has been skipped
CI/CD Pipeline / Build Production Web Image (push) Has been skipped
CI/CD Pipeline / Build Production Worker Image (push) Has been skipped
CI/CD Pipeline / CI Gate (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Canary Release to Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
fix(api+tests): CI unit-test failures for WS progress endpoint (#2104)
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
2026-09-30 11:05:55 +08:00

433 lines
17 KiB
Python

"""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
* Initial-snapshot / terminal-job fast-close behaviour of the WS endpoint
The CI unit-test environment sets ``USE_IN_MEMORY_DB=true`` and relies on
``settings.effective_database_url`` returning a SQLite URL. ``app/db.py`` and
``app/dependencies.py`` have been fixed to honour ``effective_database_url``
(matching the worker), so these tests never need a real Postgres or Redis.
Imports go through the ``apps.worker.*`` namespace (not bare ``worker_app.*``)
to stay consistent with the existing integration tests and avoid creating a
second module object that would make cross-file patches invisible.
"""
from __future__ import annotations
import json
import os
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
# Ensure CI-friendly env is set BEFORE any app import so SQLite is used.
os.environ.setdefault("USE_IN_MEMORY_DB", "true")
os.environ.setdefault("JWT_SECRET_KEY", "test-secret")
os.environ.setdefault("DATABASE_URL", "postgresql+psycopg://no:such@127.0.0.1:1/none")
import app.db as _app_db # noqa: E402
from app.api.routes import viral_video as vv_module # noqa: E402
from apps.worker.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):
job = _make_job(status="running")
assert vv_module._job_status(job) == "running"
job_enum = _make_job(status=SimpleNamespace(value="completed"))
assert vv_module._job_status(job_enum) == "completed"
# ---------------------------------------------------------------------------
# WebSocket authentication / ownership / 404
# ---------------------------------------------------------------------------
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
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 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):
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):
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 initial-snapshot / terminal-job fast-close tests.
#
# The Redis pubsub reader thread is factored into ``_run_pubsub_forwarder`` and
# marked ``# pragma: no cover`` (integration-tested with a live Redis). These
# tests patch it out so we can deterministically verify the pre-subscribe
# handshake without needing a real Redis or real thread scheduling.
# ---------------------------------------------------------------------------
def _run_ws_handshake(*, job):
"""Drive a WS handshake; collect JSON messages before connection closes."""
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):
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 = []
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()
return received, 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, 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"
# Session was used for both ownership check and initial snapshot.
assert sess.close.call_count >= 2
def test_running_job_with_enum_status(self):
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, _, _ = _run_ws_handshake(job=job)
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()
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()
assert any(m["type"] == "forwarder_reached" for m in received)