fix(api+tests): CI unit-test failures for WS progress endpoint #2104
+2
-2
@@ -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,
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user