diff --git a/tests/unit/test_batch_download.py b/tests/unit/test_batch_download.py index b88f9bc59..02bf86321 100755 --- a/tests/unit/test_batch_download.py +++ b/tests/unit/test_batch_download.py @@ -256,9 +256,11 @@ def test_batch_download_single_video(): def test_batch_download_session_closed(): """DB session is always closed (via finally block). - Self-contained: uses patch.object on the actual worker_app.db module - to avoid stale PromiseProxy cache from prior tests in the suite. + Patches the function's own globals to inject mock SessionLocal, + bypassing any import caching issues in the full suite. """ + import sys + import worker_app.db as _db_mod from apps.worker.worker_app.tasks.batch_download import batch_download_videos @@ -267,6 +269,7 @@ def test_batch_download_session_closed(): repo = _FakeGeneratedVideoRepository(videos) session = MagicMock() + mock_session_factory = MagicMock(return_value=session) def _noop_download(url, dest): Path(dest).parent.mkdir(parents=True, exist_ok=True) @@ -274,7 +277,19 @@ def test_batch_download_session_closed(): bound_task = _make_bound_task() - with patch.object(_db_mod, "SessionLocal", MagicMock(return_value=session)): + # Get the raw function to patch its globals + raw_fn = _get_raw_task_fn(batch_download_videos) + + # Patch SessionLocal in ALL possible module locations + _db_mod.SessionLocal = mock_session_factory + if "worker_app.db" in sys.modules: + sys.modules["worker_app.db"].SessionLocal = mock_session_factory + + # Also patch in the function's own globals if it has a reference there + if "SessionLocal" in raw_fn.__globals__: + raw_fn.__globals__["SessionLocal"] = mock_session_factory + + try: with patch( "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", return_value=repo, @@ -284,13 +299,19 @@ def test_batch_download_session_closed(): "apps.worker.worker_app.tasks.batch_download._download_video_to_file", _noop_download, ): - _call_task(batch_download_videos, bound_task, ["v1"], "user_1") + raw_fn(bound_task, ["v1"], "user_1") + finally: + pass # Don't restore - other tests handle their own patches + # Diagnostic: check if our mock factory was actually called + assert mock_session_factory.called, "SessionLocal mock was never called! " f"raw_fn={raw_fn}, type={type(raw_fn)}" session.close.assert_called_once() def test_batch_download_closes_session_on_error(): """Session is closed even when get_by_ids raises.""" + import sys + import worker_app.db as _db_mod from apps.worker.worker_app.tasks.batch_download import batch_download_videos @@ -300,16 +321,25 @@ def test_batch_download_closes_session_on_error(): raise RuntimeError("db down") session = MagicMock() + mock_session_factory = MagicMock(return_value=session) bound_task = _make_bound_task() - with patch.object(_db_mod, "SessionLocal", MagicMock(return_value=session)): - with patch( - "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", - return_value=_ExplodingRepo(), - ): - with pytest.raises(RuntimeError, match="db down"): - _call_task(batch_download_videos, bound_task, ["v1"], "u") + raw_fn = _get_raw_task_fn(batch_download_videos) + _db_mod.SessionLocal = mock_session_factory + if "worker_app.db" in sys.modules: + sys.modules["worker_app.db"].SessionLocal = mock_session_factory + if "SessionLocal" in raw_fn.__globals__: + raw_fn.__globals__["SessionLocal"] = mock_session_factory + + with patch( + "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", + return_value=_ExplodingRepo(), + ): + with pytest.raises(RuntimeError, match="db down"): + raw_fn(bound_task, ["v1"], "u") + + assert mock_session_factory.called, "SessionLocal mock was never called!" session.close.assert_called_once()