"""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)