diff --git a/apps/api/app/db.py b/apps/api/app/db.py index c9a879910..3819b672c 100644 --- a/apps/api/app/db.py +++ b/apps/api/app/db.py @@ -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, diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index ce622df26..12b2d0b43 100644 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -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]: diff --git a/tests/unit/test_viral_video_ws.py b/tests/unit/test_viral_video_ws.py index 9d8594762..67185f129 100644 --- a/tests/unit/test_viral_video_ws.py +++ b/tests/unit/test_viral_video_ws.py @@ -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)