""" 管理会计OS 测试配置 使用 SQLite 内存数据库进行测试,避免依赖外部 MySQL。 测试前自动建表,测试后自动清理。 重要:此文件在pytest收集测试时最先加载,确保环境变量在app模块导入前注入。 """ import os # 必须在任何app模块导入之前设置环境变量(通过 pytest.ini 的 python_files 保证加载顺序) os.environ.setdefault("CMA_DB_USER", "test") os.environ.setdefault("CMA_DB_PASS", "test") os.environ.setdefault("CMA_DB_HOST", "localhost") os.environ.setdefault("CMA_DB_PORT", "3306") os.environ.setdefault("CMA_DB_NAME", "test") import pytest from typing import Generator from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker, Session from sqlalchemy.pool import StaticPool # 在 import app 模块之前就先覆盖掉 database.py 的 engine # 方式:直接 monkey-patch database 模块 from app import database as db_module from app.database import Base # SQLite 内存引擎 TEST_ENGINE = create_engine( "sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) TEST_SESSION_LOCAL = sessionmaker(autocommit=False, autoflush=False, bind=TEST_ENGINE) # MySQL-only 的 date_format() 在 SQLite 下注册等价实现(仅测试库) # 生产用 MySQL 原生函数;此处仅为让测试能跑通 expenses 月度累计校验/stats 统计 def _sqlite_date_format(dt_val, fmt): if dt_val is None: return None import datetime as _dt if isinstance(dt_val, str): for f in ("%Y-%m-%d %H:%M:%S", "%Y-%m-%d", "%Y-%m"): try: dt_val = _dt.datetime.strptime(str(dt_val)[:19], f) break except ValueError: continue else: return None if isinstance(dt_val, _dt.datetime): d = dt_val elif isinstance(dt_val, _dt.date): d = _dt.datetime(dt_val.year, dt_val.month, dt_val.day) else: return None return { "%Y": f"{d.year:04d}", "%Y-%m": f"{d.year:04d}-{d.month:02d}", "%Y-%m-%d": f"{d.year:04d}-{d.month:02d}-{d.day:02d}", }.get(fmt) from sqlalchemy import event # noqa: E402 event.listen(TEST_ENGINE, "connect", lambda dbapi_conn, rec: dbapi_conn.create_function("date_format", 2, _sqlite_date_format)) # 替换 database 模块的全局引擎 db_module._engine = TEST_ENGINE db_module._SessionLocal = TEST_SESSION_LOCAL # 然后导入models(它们依赖Base) from app.models import ( User, StrategicMap, KPIDefinition, KPIValue, KPIAlert, ActionPlan, OrgNode, StrategicMapVersion, ) import hashlib @pytest.fixture(autouse=True) def setup_db(): """每个测试函数自动初始化和清理数据库""" Base.metadata.create_all(bind=TEST_ENGINE) yield Base.metadata.drop_all(bind=TEST_ENGINE) @pytest.fixture def db() -> Generator[Session, None, None]: """提供数据库 session""" session = TEST_SESSION_LOCAL() try: yield session finally: session.close() @pytest.fixture def client(db) -> Generator[TestClient, None, None]: """提供测试 HTTP 客户端""" from app.main import app from app.models import Entity # 账套模式:确保测试库存在 entity_id=1 的active实体 ent = db.query(Entity).filter(Entity.id == 1).first() if not ent: db.add(Entity(id=1, name="测试企业", short_name="测试", status="active")) db.commit() # 重写依赖,使用测试数据库 app.dependency_overrides[db_module.get_db] = lambda: db with TestClient(app) as c: yield c app.dependency_overrides.clear() # ── 测试数据工厂 ── def create_test_user(db: Session, **kwargs) -> User: """创建测试用户""" defaults = { "username": "testadmin", "password_hash": hashlib.sha256("admin123".encode()).hexdigest(), "name": "测试管理员", "role": "ceo", } defaults.update(kwargs) user = User(**defaults) db.add(user) db.commit() db.refresh(user) return user def get_token_for_user(client: TestClient, username: str = "testadmin", password: str = "admin123") -> str: """获取测试用户的token(账套模式:需entity_id)""" resp = client.post("/api/cma/auth/login", json={ "username": username, "password": password, "entity_id": 1, }) if resp.status_code != 200: raise RuntimeError(f"登录失败: {resp.status_code} {resp.text[:300]}") data = resp.json() return data.get("token") or data.get("access_token") def auth_header(token: str) -> dict: return {"Authorization": f"Bearer {token}"} def create_test_kpi(db: Session, **kwargs) -> KPIDefinition: """创建测试KPI""" defaults = { "kpi_code": "TEST_001", "kpi_name": "测试KPI", "dimension": "finance", "target_value": 100.0, "unit": "%", "status": "active", } defaults.update(kwargs) kpi = KPIDefinition(**defaults) db.add(kpi) db.commit() db.refresh(kpi) return kpi def create_test_map(db: Session, **kwargs) -> StrategicMap: """创建测试战略地图""" defaults = { "title": "测试地图", "status": "draft", "dimensions": [ {"key": "finance", "name": "财务维度", "icon": "💰", "color": "#409eff", "objectives": []}, {"key": "customer", "name": "客户维度", "icon": "🤝", "color": "#67c23a", "objectives": []}, ], "canvas_data": {"connections": []}, } defaults.update(kwargs) m = StrategicMap(**defaults) db.add(m) db.commit() db.refresh(m) return m