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
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com> Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
433 lines
17 KiB
Python
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)
|