- 新增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
191 lines
5.6 KiB
Python
191 lines
5.6 KiB
Python
"""
|
||
管理会计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
|