fix(api+tests): CI unit-test failures for WS progress endpoint #2104

Merged
auto-approve-bot merged 2 commits from feat/2051-ws-ci-fix into develop 2026-09-30 11:05:57 +08:00
3 changed files with 52 additions and 100 deletions
+2 -2
View File
@@ -7,9 +7,9 @@ from packages.adapters.sqlalchemy_impl import (
)
from packages.adapters.sqlalchemy_impl.schema_guard import assert_auto_create_schema_allowed
ensure_database_exists(settings.DATABASE_URL)
ensure_database_exists(settings.effective_database_url)
engine, SessionLocal = build_session_factory(
settings.DATABASE_URL,
settings.effective_database_url,
pool_size=settings.DATABASE_POOL_SIZE,
max_overflow=settings.DATABASE_MAX_OVERFLOW,
pool_timeout=settings.DATABASE_POOL_TIMEOUT,
+1 -1
View File
@@ -56,7 +56,7 @@ from packages.adapters.sqlalchemy_impl.voice_library_repository import (
from packages.ports.tag_repository import TagRepository
from packages.ports.user_repository import UserRepository
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
_engine, _SessionLocal = build_session_factory(settings.effective_database_url)
def get_db_session() -> Generator[Session, None, None]:
+49 -97
View File
@@ -3,23 +3,23 @@
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
* WebSocket authentication / ownership / 404 behaviour
* Initial-snapshot / terminal-job fast-close behaviour of the WS endpoint
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.
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 sys
import threading
import os
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
@@ -28,36 +28,15 @@ 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]
# 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 worker_app.tasks import viral_video as worker_vv # noqa: E402
from apps.worker.worker_app.tasks import viral_video as worker_vv # noqa: E402
def _make_job(**kwargs):
@@ -169,10 +148,8 @@ class TestWSHelpers:
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"
@@ -182,6 +159,18 @@ class TestWSHelpers:
# ---------------------------------------------------------------------------
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)
@@ -201,18 +190,6 @@ def _build_client(*, auth_user, repo_get_return, redis_instance=None):
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")
@@ -250,8 +227,6 @@ class TestWSAuthenticateUser:
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")):
@@ -273,7 +248,6 @@ class TestWSAuthenticateUser:
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()
@@ -291,19 +265,17 @@ class TestWSAuthenticateUser:
# ---------------------------------------------------------------------------
# WebSocket full-flow tests (initial snapshot, terminal jobs)
# 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 real Redis). These
# marked ``# pragma: no cover`` (integration-tested with a live Redis). These
# tests patch it out so we can deterministically verify the pre-subscribe
# handshake: auth, ownership, initial snapshot, and terminal-job fast-close.
# handshake without needing a real Redis or real thread scheduling.
# ---------------------------------------------------------------------------
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)."""
def _run_ws_handshake(*, job):
"""Drive a WS handshake; collect JSON messages before connection closes."""
app = FastAPI()
app.include_router(vv_module.router)
@@ -313,8 +285,6 @@ def _run_ws_handshake(*, job, expect_close_code=None, expect_messages_before_clo
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:
@@ -330,54 +300,38 @@ def _run_ws_handshake(*, job, expect_close_code=None, expect_messages_before_clo
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
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
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, code, repo, sess = _run_ws_handshake(job=job)
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"
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.
# Session was used for both ownership check and 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)
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"]
@@ -389,7 +343,7 @@ class TestWSInitialSnapshot:
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)
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
@@ -404,7 +358,7 @@ class TestWSInitialSnapshot:
is_terminal=True,
error_msg="out of memory",
)
received, _, _, _ = _run_ws_handshake(job=job)
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
@@ -416,7 +370,7 @@ class TestWSInitialSnapshot:
is_terminal=True,
result_video_url=None,
)
received, _, _, _ = _run_ws_handshake(job=job)
received, _, _ = _run_ws_handshake(job=job)
completed = next(m for m in received if m["type"] == "viral_video:completed")
assert completed["data"]["video_url"] == ""
@@ -427,7 +381,7 @@ class TestWSInitialSnapshot:
is_terminal=True,
error_msg=None,
)
received, _, _, _ = _run_ws_handshake(job=job)
received, _, _ = _run_ws_handshake(job=job)
failed = next(m for m in received if m["type"] == "viral_video:failed")
assert failed["data"]["error"] == ""
@@ -444,7 +398,6 @@ class TestWSInitialSnapshot:
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):
@@ -476,5 +429,4 @@ class TestWSInitialSnapshot:
finally:
for p in patches:
p.stop()
# Even after snapshot error, forwarder is reached
assert any(m["type"] == "forwarder_reached" for m in received)