From 153d04d259f28a79675879afda80c9dc736e152a Mon Sep 17 00:00:00 2001 From: CI Bot Date: Sat, 22 Aug 2026 21:30:13 +0800 Subject: [PATCH] fix(test): resolve flaky test pollution from sys.modules mock leakage - test_generation_worker_fixes: use patch.dict(sys.modules) with a fresh module instead of patch.object, avoiding leakage from other test files that pre-register MagicMock for worker_app.db - test_health_routes: replace dotted-path @patch decorators with patch.object on the imported module to avoid FastAPI app attribute shadowing in the apps.api.app namespace --- tests/unit/test_generation_worker_fixes.py | 25 ++- tests/unit/test_health_routes.py | 171 ++++++++++----------- 2 files changed, 98 insertions(+), 98 deletions(-) diff --git a/tests/unit/test_generation_worker_fixes.py b/tests/unit/test_generation_worker_fixes.py index 70c43294c..edbd4ac2f 100644 --- a/tests/unit/test_generation_worker_fixes.py +++ b/tests/unit/test_generation_worker_fixes.py @@ -13,6 +13,17 @@ os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret") os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") +def _patch_session_local(mock_session): + """Patch worker_app.db.SessionLocal robustly even when other tests + have pre-registered a MagicMock for worker_app.db in sys.modules. + Uses patch.dict to inject a clean module so that + 'from worker_app.db import SessionLocal' resolves correctly.""" + from types import ModuleType + _fresh_db = ModuleType("worker_app.db") + _fresh_db.SessionLocal = lambda *a, **kw: mock_session + return patch.dict(sys.modules, {"worker_app.db": _fresh_db}) + + class TestLoadTemplateSegmentDurations: """_load_template_segment_durations 单元测试 (covers lines 198-226).""" @@ -40,8 +51,7 @@ class TestLoadTemplateSegmentDurations: mock_session = MagicMock() mock_session.query.return_value = mock_query - # Patch at the source module since it's imported inside the function - with patch("worker_app.db.SessionLocal", return_value=mock_session): + with _patch_session_local(mock_session): result = _load_template_segment_durations("tpl_123") assert result == [5.0, 8.0, 3.0] @@ -63,7 +73,7 @@ class TestLoadTemplateSegmentDurations: mock_session = MagicMock() mock_session.query.return_value = mock_query - with patch("worker_app.db.SessionLocal", return_value=mock_session): + with _patch_session_local(mock_session): result = _load_template_segment_durations("tpl_456") assert result == [5.0] @@ -72,7 +82,12 @@ class TestLoadTemplateSegmentDurations: """数据库异常返回空列表,不抛出。""" from worker_app.tasks.generation import _load_template_segment_durations - with patch("worker_app.db.SessionLocal", side_effect=Exception("DB down")): + from types import ModuleType + _err_db = ModuleType("worker_app.db") + def _raise(*a, **kw): + raise Exception("DB down") + _err_db.SessionLocal = _raise + with patch.dict(sys.modules, {"worker_app.db": _err_db}): result = _load_template_segment_durations("tpl_789") assert result == [] @@ -86,7 +101,7 @@ class TestLoadTemplateSegmentDurations: mock_session = MagicMock() mock_session.query.return_value = mock_query - with patch("worker_app.db.SessionLocal", return_value=mock_session): + with _patch_session_local(mock_session): result = _load_template_segment_durations("tpl_empty") assert result == [] diff --git a/tests/unit/test_health_routes.py b/tests/unit/test_health_routes.py index beb6543a7..bb15101eb 100644 --- a/tests/unit/test_health_routes.py +++ b/tests/unit/test_health_routes.py @@ -1,7 +1,6 @@ """Unit tests for apps/api/app/api/routes/health.py 覆盖 _check_database() 和 _check_migrations() 中 psycopg3 连接逻辑。 -确保增量覆盖率 ≥ 60%(目标覆盖 lines 52, 127)。 """ from unittest.mock import AsyncMock, MagicMock, patch @@ -9,63 +8,70 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +def _make_cursor(fetchone_result=None): + """Create a mock cursor with context manager support.""" + cur = MagicMock() + cur.__enter__ = MagicMock(return_value=cur) + cur.__exit__ = MagicMock(return_value=False) + if fetchone_result is not None: + cur.fetchone.return_value = fetchone_result + return cur + + +def _make_conn(cursor_result=None): + conn = MagicMock() + conn.cursor.return_value = cursor_result or _make_cursor() + return conn + + @pytest.mark.asyncio class TestCheckDatabase: - """Tests for _check_database() health check function.""" - @patch("apps.api.app.api.routes.health.settings") - @patch("apps.api.app.api.routes.health.psycopg.connect") - async def test_check_database_success(self, mock_connect, mock_settings): - """PostgreSQL 连接成功时返回 healthy。""" + async def test_check_database_success(self): + mock_cur = _make_cursor(fetchone_result=(1,)) + mock_conn = _make_conn(mock_cur) + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = False mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test" - # Mock connection and cursor - mock_conn = MagicMock() - mock_cursor = MagicMock() - mock_cursor.__enter__ = MagicMock(return_value=mock_cursor) - mock_cursor.__exit__ = MagicMock(return_value=False) - mock_cursor.fetchone.return_value = (1,) - mock_conn.cursor.return_value = mock_cursor - mock_connect.return_value = mock_conn + from apps.api.app.api.routes import health - from apps.api.app.api.routes.health import _check_database - - result = await _check_database() + with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings): + mock_psycopg.connect.return_value = mock_conn + result = await health._check_database() assert result["status"] == "healthy" assert result["type"] == "postgresql" assert result["message"] == "Database connection successful" - mock_connect.assert_called_once_with( + mock_psycopg.connect.assert_called_once_with( "postgresql+psycopg://test:test@localhost/test", connect_timeout=3 ) - mock_cursor.execute.assert_called_once_with("SELECT 1") + mock_cur.execute.assert_called_once_with("SELECT 1") mock_conn.close.assert_called_once() - @patch("apps.api.app.api.routes.health.settings") - @patch("apps.api.app.api.routes.health.psycopg.connect") - async def test_check_database_connection_failure(self, mock_connect, mock_settings): - """PostgreSQL 连接失败时返回 unhealthy。""" + async def test_check_database_connection_failure(self): + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = False mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test" - mock_connect.side_effect = Exception("connection refused") - from apps.api.app.api.routes.health import _check_database + from apps.api.app.api.routes import health - result = await _check_database() + with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings): + mock_psycopg.connect.side_effect = Exception("connection refused") + result = await health._check_database() assert result["status"] == "unhealthy" assert result["type"] == "postgresql" assert "connection refused" in result["message"] - @patch("apps.api.app.api.routes.health.settings") - async def test_check_database_in_memory(self, mock_settings): - """使用内存数据库时跳过 PostgreSQL 检查。""" + async def test_check_database_in_memory(self): + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = True - from apps.api.app.api.routes.health import _check_database + from apps.api.app.api.routes import health - result = await _check_database() + with patch.object(health, "settings", mock_settings): + result = await health._check_database() assert result["status"] == "healthy" assert result["type"] == "in_memory" @@ -73,80 +79,66 @@ class TestCheckDatabase: @pytest.mark.asyncio class TestCheckMigrations: - """Tests for _check_migrations() health check function.""" - @patch("apps.api.app.api.routes.health.settings") - @patch("apps.api.app.api.routes.health.psycopg.connect") - async def test_check_migrations_success(self, mock_connect, mock_settings): - """所有迁移表存在时返回 healthy。""" + async def test_check_migrations_success(self): + mock_cur = _make_cursor(fetchone_result=(5,)) + mock_conn = _make_conn(mock_cur) + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = False mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test" - mock_conn = MagicMock() - mock_cursor = MagicMock() - mock_cursor.__enter__ = MagicMock(return_value=mock_cursor) - mock_cursor.__exit__ = MagicMock(return_value=False) - mock_cursor.fetchone.return_value = (5,) # 5 tables found - mock_conn.cursor.return_value = mock_cursor - mock_connect.return_value = mock_conn + from apps.api.app.api.routes import health - from apps.api.app.api.routes.health import _check_migrations - - result = await _check_migrations() + with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings): + mock_psycopg.connect.return_value = mock_conn + result = await health._check_migrations() assert result["status"] == "healthy" assert result["message"] == "Database migrations applied" - mock_connect.assert_called_once_with( + mock_psycopg.connect.assert_called_once_with( "postgresql+psycopg://test:test@localhost/test", connect_timeout=3 ) mock_conn.close.assert_called_once() - @patch("apps.api.app.api.routes.health.settings") - @patch("apps.api.app.api.routes.health.psycopg.connect") - async def test_check_migrations_missing_tables(self, mock_connect, mock_settings): - """迁移表不完整时返回 unhealthy。""" + async def test_check_migrations_missing_tables(self): + mock_cur = _make_cursor(fetchone_result=(2,)) + mock_conn = _make_conn(mock_cur) + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = False mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test" - mock_conn = MagicMock() - mock_cursor = MagicMock() - mock_cursor.__enter__ = MagicMock(return_value=mock_cursor) - mock_cursor.__exit__ = MagicMock(return_value=False) - mock_cursor.fetchone.return_value = (2,) # Only 2 of 5 tables - mock_conn.cursor.return_value = mock_cursor - mock_connect.return_value = mock_conn + from apps.api.app.api.routes import health - from apps.api.app.api.routes.health import _check_migrations - - result = await _check_migrations() + with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings): + mock_psycopg.connect.return_value = mock_conn + result = await health._check_migrations() assert result["status"] == "unhealthy" assert "Missing tables" in result["message"] assert "2/5" in result["message"] - @patch("apps.api.app.api.routes.health.settings") - @patch("apps.api.app.api.routes.health.psycopg.connect") - async def test_check_migrations_connection_failure(self, mock_connect, mock_settings): - """数据库连接失败时返回 unhealthy。""" + async def test_check_migrations_connection_failure(self): + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = False mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test" - mock_connect.side_effect = Exception("connection refused") - from apps.api.app.api.routes.health import _check_migrations + from apps.api.app.api.routes import health - result = await _check_migrations() + with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings): + mock_psycopg.connect.side_effect = Exception("connection refused") + result = await health._check_migrations() assert result["status"] == "unhealthy" assert "Migration check failed" in result["message"] - @patch("apps.api.app.api.routes.health.settings") - async def test_check_migrations_in_memory(self, mock_settings): - """使用内存数据库时跳过迁移检查。""" + async def test_check_migrations_in_memory(self): + mock_settings = MagicMock() mock_settings.USE_IN_MEMORY_DB = True - from apps.api.app.api.routes.health import _check_migrations + from apps.api.app.api.routes import health - result = await _check_migrations() + with patch.object(health, "settings", mock_settings): + result = await health._check_migrations() assert result["status"] == "healthy" assert "no migrations needed" in result["message"] @@ -154,33 +146,26 @@ class TestCheckMigrations: @pytest.mark.asyncio class TestStartupCheck: - """Tests for startup_check() endpoint.""" - @patch("apps.api.app.api.routes.health._check_migrations") - @patch("apps.api.app.api.routes.health._check_database") - async def test_startup_all_healthy(self, mock_db, mock_mig): - """所有检查通过时返回 started。""" - mock_db.return_value = {"status": "healthy"} - mock_mig.return_value = {"status": "healthy"} + async def test_startup_all_healthy(self): + from apps.api.app.api.routes import health - from apps.api.app.api.routes.health import startup_check - - result = await startup_check() + with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db: + mock_db.return_value = {"status": "healthy"} + mock_mig.return_value = {"status": "healthy"} + result = await health.startup_check() assert result["status"] == "started" - @patch("apps.api.app.api.routes.health._check_migrations") - @patch("apps.api.app.api.routes.health._check_database") - async def test_startup_db_unhealthy(self, mock_db, mock_mig): - """数据库不健康时返回 starting + 503。""" - mock_db.return_value = {"status": "unhealthy", "message": "fail"} - mock_mig.return_value = {"status": "healthy"} + async def test_startup_db_unhealthy(self): + import json + from apps.api.app.api.routes import health - from apps.api.app.api.routes.health import startup_check - - result = await startup_check() + with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db: + mock_db.return_value = {"status": "unhealthy", "message": "fail"} + mock_mig.return_value = {"status": "healthy"} + result = await health.startup_check() assert result.status_code == 503 - import json body = json.loads(result.body) assert body["status"] == "starting"