"""共享测试夹具 — 所有测试文件共用同一文件级 SQLite 数据库。 使用同步引擎在模块级创建表,避免跨事件循环导致表缺失。 """ import json import os from unittest.mock import AsyncMock, patch import pytest from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.pool import StaticPool from app.core.database import Base, get_db from app.main import app TEST_DB_PATH = "./test.db" TEST_DATABASE_URL = f"sqlite+aiosqlite:///{TEST_DB_PATH}" SYNC_DATABASE_URL = f"sqlite:///{TEST_DB_PATH}" # --- 同步引擎:模块级创建表 --- _sync_engine = create_engine(SYNC_DATABASE_URL, echo=False) def _create_tables() -> None: """用同步引擎创建全部表。""" Base.metadata.create_all(_sync_engine) def _drop_tables() -> None: """用同步引擎销毁全部表。""" Base.metadata.drop_all(_sync_engine) # 模块导入时立即创建表 if os.path.exists(TEST_DB_PATH): os.remove(TEST_DB_PATH) _create_tables() # --- 异步引擎:测试运行时使用 --- # StaticPool 保持单连接复用,connect_args 允许跨线程访问 test_engine = create_async_engine( TEST_DATABASE_URL, echo=False, poolclass=StaticPool, connect_args={"check_same_thread": False}, ) test_session_factory = async_sessionmaker(test_engine, class_=AsyncSession, expire_on_commit=False) async def override_get_db(): """覆盖数据库依赖,使用测试 SQLite。""" async with test_session_factory() as session: try: yield session await session.commit() except Exception: await session.rollback() raise app.dependency_overrides[get_db] = override_get_db @pytest.fixture(scope="session", autouse=True) async def setup_db(): """确保表已创建,测试结束后清理。""" # 表已在模块级同步创建,此处仅做清理 yield await test_engine.dispose() _drop_tables() _sync_engine.dispose() if os.path.exists(TEST_DB_PATH): os.remove(TEST_DB_PATH) @pytest.fixture(autouse=True) async def mock_llm(): """Mock LLM 客户端 — 所有 AI 相关测试不依赖真实 API。 chat() 尝试解析 JSON 返回 Python 对象(dict/list), 使 service 层 isinstance(result, dict/list) 检查通过。 无法解析时返回原始字符串。 """ mock_dict = {"summary": "AI 分析结果(测试 mock)", "score": 75, "confidence": 0.8} mock_list = [{"type": "tech", "title": "测试协同", "match_reason": "mock"}] async def mock_chat_stream(self, messages, temperature=0.3, max_tokens=2000): yield json.dumps(mock_dict) async def mock_chat(self, messages, temperature=0.3, max_tokens=2000): # service 层传的可能是 str 或 list[dict] # 根据 prompt 内容判断期望返回类型 if isinstance(messages, str): prompt_text = messages elif isinstance(messages, list): prompt_text = " ".join(str(m.get("content", "")) for m in messages) else: prompt_text = str(messages) # 如果 prompt 中包含 "JSON 数组" 或 "数组格式",返回 list if "数组" in prompt_text or "JSON 数组" in prompt_text: return mock_list # 默认返回 dict return mock_dict with patch("app.services.llm_client.LLMClient.chat", mock_chat), \ patch("app.services.llm_client.LLMClient.chat_stream", mock_chat_stream): yield @pytest.fixture def client(): """创建测试客户端。""" return TestClient(app) @pytest.fixture def auth_headers(client: TestClient): """注册并登录投资人用户,返回认证头。""" client.post( "/api/v1/auth/register", json={ "email": "smoke_test@example.com", "password": "password123", "name": "冒烟测试用户", "tenant_name": "冒烟测试机构", "role": "investor", }, ) resp = client.post( "/api/v1/auth/login", json={"email": "smoke_test@example.com", "password": "password123"}, ) token = resp.json()["data"]["access_token"] return {"Authorization": f"Bearer {token}"} @pytest.fixture def company_id(client: TestClient, auth_headers: dict): """创建测试企业,返回企业 ID。""" resp = client.post( "/api/v1/companies", json={"name": "冒烟测试公司", "industry": "AI"}, headers=auth_headers, ) return resp.json()["data"]["id"]