Files
cma-management/backend/tests/conftest.py
T
Hermes CI Fix a13a080381 test: CMA自动化测试补覆盖 226→452用例, 覆盖率36%→60%
- 新增8个测试文件(bot_bridge/kpi_causality/cash/predict/reports/tax_compliance/expenses/probe_cost)
- 增强 budget/auth/users + conftest账套模式适配
- 测试驱动修复: bot_bridge导入batch_id→source_batch; cash_forecast extra空dict
- 全量: 451 passed, 1 xfailed; 报告 docs/cma-test-coverage-report.md
2026-08-20 06:57:24 +08:00

191 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
管理会计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