merge: develop -> main (P0/P1/P2全功能)
This commit is contained in:
Binary file not shown.
@@ -99,6 +99,10 @@ def get_kpi(kpi_id: int, db: Session = Depends(get_db)):
|
||||
|
||||
@router.post("")
|
||||
def create_kpi(data: dict, db: Session = Depends(get_db), user=WRITE_ROLES):
|
||||
# 检查编码唯一性
|
||||
existing = db.query(KPIDefinition).filter(KPIDefinition.kpi_code == data.get("kpi_code", "")).first()
|
||||
if existing:
|
||||
raise HTTPException(400, f"KPI编码 {data['kpi_code']} 已存在")
|
||||
kpi = KPIDefinition(**data)
|
||||
db.add(kpi)
|
||||
db.commit()
|
||||
|
||||
+8
-4
@@ -1,13 +1,14 @@
|
||||
"""管理会计OS — 主入口"""
|
||||
import logging
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi import FastAPI, Request, Depends
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from dotenv import load_dotenv
|
||||
from app.database import init_db
|
||||
from app.api import auth, kpis, maps, dashboard, data, alerts, ai_analysis, alert_rules, users, thresholds, notifications, permissions, action_plans, alignment, org, objectives, versions, budget, cost, predict
|
||||
from app.api import auth, kpis, templates, maps, dashboard, data, alerts, ai_analysis, alert_rules, users, thresholds, notifications, permissions, action_plans, alignment, org, objectives, versions, budget, cost, predict, reports, security
|
||||
from app.utils.cache import clear_all as clear_cache, delete as delete_cache
|
||||
from scripts.erp_sync import run_sync as run_erp_sync
|
||||
from app.auth_middleware import require_auth
|
||||
|
||||
load_dotenv()
|
||||
|
||||
@@ -31,6 +32,7 @@ app.add_middleware(
|
||||
|
||||
app.include_router(auth.router)
|
||||
app.include_router(kpis.router)
|
||||
app.include_router(templates.router)
|
||||
app.include_router(maps.router)
|
||||
app.include_router(dashboard.router)
|
||||
app.include_router(data.router)
|
||||
@@ -49,6 +51,8 @@ app.include_router(versions.router)
|
||||
app.include_router(budget.router)
|
||||
app.include_router(cost.router)
|
||||
app.include_router(predict.router)
|
||||
app.include_router(reports.router)
|
||||
app.include_router(security.router)
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def global_exception_handler(request: Request, exc: Exception):
|
||||
@@ -63,7 +67,7 @@ def startup():
|
||||
|
||||
|
||||
@app.post("/api/cma/admin/erp-sync")
|
||||
def admin_erp_sync(kpi_codes: str = None):
|
||||
def admin_erp_sync(user=Depends(require_auth), kpi_codes: str = None):
|
||||
"""手动触发ERP数据同步"""
|
||||
kpi_list = kpi_codes.split(",") if kpi_codes else None
|
||||
try:
|
||||
@@ -74,7 +78,7 @@ def admin_erp_sync(kpi_codes: str = None):
|
||||
|
||||
|
||||
@app.get("/api/cma/admin/erp-sync/dry-run")
|
||||
def admin_erp_sync_dry_run(kpi_codes: str = None):
|
||||
def admin_erp_sync_dry_run(user=Depends(require_auth), kpi_codes: str = None):
|
||||
"试运行,不写入数据库"""
|
||||
kpi_list = kpi_codes.split(",") if kpi_codes else None
|
||||
try:
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
"""KPI计算引擎 v4 — 基于会计科目余额和销售报表"""
|
||||
import httpx, asyncio
|
||||
import httpx, asyncio, os
|
||||
from datetime import datetime
|
||||
from app.database import get_session_local
|
||||
from app.models import KPIDefinition, KPIValue
|
||||
|
||||
ERP_API = "http://127.0.0.1:8300"
|
||||
ERP_KEY = "erp-gateway-key-bhwl-2026"
|
||||
ERP_KEY = os.environ.get("ERP_API_KEY", "erp-gateway-key-bhwl-2026")
|
||||
|
||||
async def _get(url: str, params: dict = None):
|
||||
async with httpx.AsyncClient(timeout=20) as c:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
标准成本vs实际成本差异分析(量差/价差/效率差异)
|
||||
ABC作业成本法分配
|
||||
"""
|
||||
import logging
|
||||
import logging, os
|
||||
from datetime import datetime
|
||||
from typing import Optional, List, Dict
|
||||
from app.database import get_session_local
|
||||
@@ -11,7 +11,7 @@ from app.models import StandardCost, ActualCost, AbcActivity, AbcAllocation, KPI
|
||||
logger = logging.getLogger("cma.cost")
|
||||
|
||||
ERP_API = "http://127.0.0.1:8300"
|
||||
ERP_KEY = "erp-gateway-key-bhwl-2026"
|
||||
ERP_KEY = os.environ.get("ERP_API_KEY", "erp-gateway-key-bhwl-2026")
|
||||
|
||||
|
||||
# ============================================================
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
"""Step 1: 全量采集ERP表结构到 erp_schema
|
||||
通过 erp-api-gateway 采集985张表的字段信息
|
||||
运行: python3 scripts/collect_erp_schema.py
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import json
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
logger = logging.getLogger("erp_schema_collect")
|
||||
|
||||
ERP_API_BASE = "http://127.0.0.1:8300/api/v1"
|
||||
ERP_API_KEY = os.getenv("ERP_API_KEY", "erp-gateway-key-bhwl-2026")
|
||||
|
||||
HEADERS = {
|
||||
"X-API-Key": ERP_API_KEY,
|
||||
"User-Agent": "CMA-ERP-SCHEMA/1.0",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
|
||||
def api_get(path: str) -> dict:
|
||||
"""调用ERP API"""
|
||||
url = f"{ERP_API_BASE}{path}"
|
||||
req = urllib.request.Request(url, headers=HEADERS)
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
return json.loads(resp.read().decode())
|
||||
|
||||
|
||||
def get_table_columns(table_name: str) -> list:
|
||||
"""通过 INFORMATION_SCHEMA 查询表字段"""
|
||||
sql = f"""
|
||||
SELECT
|
||||
COLUMN_NAME,
|
||||
DATA_TYPE,
|
||||
CHARACTER_MAXIMUM_LENGTH,
|
||||
IS_NULLABLE,
|
||||
COLUMN_DEFAULT
|
||||
FROM INFORMATION_SCHEMA.COLUMNS
|
||||
WHERE TABLE_NAME = '{table_name}'
|
||||
ORDER BY ORDINAL_POSITION
|
||||
"""
|
||||
params = json.dumps({"sql": sql}).encode()
|
||||
req = urllib.request.Request(
|
||||
f"{ERP_API_BASE}/query",
|
||||
data=params,
|
||||
headers=HEADERS,
|
||||
method="POST"
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
data = json.loads(resp.read().decode())
|
||||
return data.get("data", [])
|
||||
except Exception as e:
|
||||
logger.warning(f" ⚠️ {table_name}: 查询失败 - {e}")
|
||||
return []
|
||||
|
||||
|
||||
def get_row_count(table_name: str) -> int:
|
||||
"""获取表行数"""
|
||||
sql = f"SELECT COUNT(*) as cnt FROM [{table_name}]"
|
||||
params = json.dumps({"sql": sql}).encode()
|
||||
req = urllib.request.Request(
|
||||
f"{ERP_API_BASE}/query",
|
||||
data=params,
|
||||
headers=HEADERS,
|
||||
method="POST"
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
data = json.loads(resp.read().decode())
|
||||
rows = data.get("data", [])
|
||||
return rows[0]["cnt"] if rows else 0
|
||||
except:
|
||||
return -1 # 未知
|
||||
|
||||
|
||||
def main():
|
||||
db = get_session_local()()
|
||||
|
||||
try:
|
||||
# 1. 获取全部表名
|
||||
logger.info("📡 获取ERP全量表名列表...")
|
||||
tables_data = api_get("/tables")
|
||||
all_tables = tables_data.get("tables", [])
|
||||
logger.info(f" 共 {len(all_tables)} 张表")
|
||||
|
||||
# 2. 获取已采集的表名
|
||||
existing = set()
|
||||
try:
|
||||
rows = db.execute(text("SELECT table_name FROM erp_schema")).fetchall()
|
||||
existing = set(row[0] for row in rows)
|
||||
except:
|
||||
pass
|
||||
logger.info(f" 已采集 {len(existing)} 张表,待采集 {len(all_tables) - len(existing)} 张")
|
||||
|
||||
# 3. 逐表采集
|
||||
collected = 0
|
||||
skipped = 0
|
||||
errors = 0
|
||||
|
||||
for i, table_name in enumerate(all_tables):
|
||||
if table_name in existing:
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
# 进度显示
|
||||
if (i + 1) % 50 == 0:
|
||||
logger.info(f" 进度: {i+1}/{len(all_tables)} (已采{collected}, 跳过{skipped}, 错误{errors})")
|
||||
|
||||
# 采集字段
|
||||
columns = get_table_columns(table_name)
|
||||
if not columns:
|
||||
errors += 1
|
||||
# 即使查不到字段也记录一个空记录,避免重复查
|
||||
collect_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
fields_json = "[]"
|
||||
db.execute(
|
||||
text("""
|
||||
INSERT INTO erp_schema (table_name, total_rows, field_count, fields_json, created_at, updated_at)
|
||||
VALUES (:tn, :tr, :fc, :fj, NOW(), NOW())
|
||||
ON DUPLICATE KEY UPDATE fields_json=:fj2, total_rows=:tr2, updated_at=NOW()
|
||||
"""),
|
||||
{"tn": table_name, "tr": -1, "fc": 0, "fj": fields_json, "fj2": fields_json, "tr2": -1}
|
||||
)
|
||||
db.commit()
|
||||
continue
|
||||
|
||||
# 获取行数
|
||||
row_count = get_row_count(table_name)
|
||||
field_count = len(columns)
|
||||
|
||||
fields = []
|
||||
for col in columns:
|
||||
fields.append({
|
||||
"name": col.get("COLUMN_NAME", ""),
|
||||
"type": col.get("DATA_TYPE", ""),
|
||||
"max_length": col.get("CHARACTER_MAXIMUM_LENGTH"),
|
||||
"nullable": col.get("IS_NULLABLE", "YES"),
|
||||
"default": col.get("COLUMN_DEFAULT"),
|
||||
})
|
||||
fields_json = json.dumps(fields, ensure_ascii=False)
|
||||
|
||||
# 写入 erp_schema
|
||||
try:
|
||||
collect_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
db.execute(
|
||||
text("""
|
||||
INSERT INTO erp_schema (table_name, total_rows, field_count, fields_json, created_at, updated_at)
|
||||
VALUES (:tn, :tr, :fc, :fj, NOW(), NOW())
|
||||
ON DUPLICATE KEY UPDATE total_rows=:tr2, field_count=:fc2, fields_json=:fj2, updated_at=NOW()
|
||||
"""),
|
||||
{"tn": table_name, "tr": row_count, "fc": field_count, "fj": fields_json,
|
||||
"tr2": row_count, "fc2": field_count, "fj2": fields_json}
|
||||
)
|
||||
db.commit()
|
||||
collected += 1
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.warning(f" ⚠️ {table_name}: 写入数据库失败 - {e}")
|
||||
errors += 1
|
||||
|
||||
# 限流:不要打太快
|
||||
if collected > 0 and collected % 10 == 0:
|
||||
time.sleep(0.5)
|
||||
|
||||
# 4. 统计
|
||||
total = db.execute(text("SELECT COUNT(*) FROM erp_schema")).scalar()
|
||||
with_data = db.execute(text("SELECT COUNT(*) FROM erp_schema WHERE field_count > 0")).scalar()
|
||||
logger.info(f"\n🎉 采集完成!")
|
||||
logger.info(f" 总计: {total} 张表 (erp_schema)")
|
||||
logger.info(f" 有字段信息: {with_data} 张")
|
||||
logger.info(f" 本次新增: {collected} 张")
|
||||
logger.info(f" 跳过(已存在): {skipped} 张")
|
||||
logger.info(f" 错误: {errors} 张")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"采集失败: {e}", exc_info=True)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Step 1 (v2): 全量采集ERP表结构到 erp_schema
|
||||
通过 erp-api-gateway 的 /api/v1/query?table=xxx&limit=1 获取字段信息
|
||||
运行: python3 scripts/collect_erp_schema_v2.py
|
||||
"""
|
||||
|
||||
import sys, os, json, urllib.request, urllib.error, logging, time
|
||||
from datetime import datetime
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
logger = logging.getLogger("erp_schema_v2")
|
||||
|
||||
API_BASE = "http://127.0.0.1:8300/api/v1"
|
||||
API_KEY = os.getenv("ERP_API_KEY", "erp-gateway-key-bhwl-2026")
|
||||
HEADERS = {"X-API-Key": API_KEY}
|
||||
|
||||
|
||||
def api_get(path: str) -> dict:
|
||||
url = f"{API_BASE}{path}"
|
||||
req = urllib.request.Request(url, headers=HEADERS)
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
return json.loads(resp.read().decode())
|
||||
|
||||
|
||||
def main():
|
||||
db = get_session_local()()
|
||||
|
||||
# 1. 获取全量表名
|
||||
logger.info("获取ERP全量表名...")
|
||||
tables_data = api_get("/tables")
|
||||
all_tables = tables_data.get("tables", [])
|
||||
logger.info(f"共 {len(all_tables)} 张表")
|
||||
|
||||
# 2. 已采集的
|
||||
existing = set()
|
||||
for row in db.execute(text("SELECT table_name FROM erp_schema")).fetchall():
|
||||
existing.add(row[0])
|
||||
logger.info(f"已有 {len(existing)} 张,还需采集 {len(all_tables) - len(existing)} 张")
|
||||
|
||||
# 3. 逐表采集
|
||||
collected, errors, skipped = 0, 0, 0
|
||||
for i, tn in enumerate(all_tables):
|
||||
if tn in existing:
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
data = api_get(f"/query?table={tn}&limit=1")
|
||||
cols = data.get("columns", [])
|
||||
total = data.get("total", 0)
|
||||
fields_json = json.dumps([{"name": c} for c in cols], ensure_ascii=False)
|
||||
|
||||
db.execute(text("""
|
||||
INSERT INTO erp_schema (table_name, total_rows, field_count, fields_json, created_at, updated_at)
|
||||
VALUES (:tn, :tr, :fc, :fj, NOW(), NOW())
|
||||
ON DUPLICATE KEY UPDATE total_rows=:tr2, field_count=:fc2, fields_json=:fj2, updated_at=NOW()
|
||||
"""), {"tn": tn, "tr": total, "fc": len(cols), "fj": fields_json,
|
||||
"tr2": total, "fc2": len(cols), "fj2": fields_json})
|
||||
db.commit()
|
||||
collected += 1
|
||||
except Exception as e:
|
||||
errors += 1
|
||||
db.execute(text("""
|
||||
INSERT INTO erp_schema (table_name, total_rows, field_count, fields_json, created_at, updated_at)
|
||||
VALUES (:tn, -1, 0, '[]', NOW(), NOW())
|
||||
ON DUPLICATE KEY UPDATE updated_at=NOW()
|
||||
"""), {"tn": tn})
|
||||
db.commit()
|
||||
|
||||
if (i + 1) % 100 == 0:
|
||||
logger.info(f"进度: {i+1}/{len(all_tables)} 已采{collected} 错误{errors} 跳过{skipped}")
|
||||
|
||||
time.sleep(0.1)
|
||||
|
||||
db.close()
|
||||
logger.info(f"\n完成! 总计:{len(all_tables)} 采集:{collected} 已有:{skipped} 错误:{errors}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,234 @@
|
||||
"""Step 4: ERP数据同步执行器
|
||||
直接连接ERP SQL Server,执行KPI-SQL并写入 kpi_values
|
||||
运行: python3 scripts/erp_data_sync.py [period] [--dry-run]
|
||||
"""
|
||||
|
||||
import sys, os, json, logging, re, argparse
|
||||
from datetime import datetime
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text, create_engine
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
logger = logging.getLogger("erp_data_sync")
|
||||
|
||||
# ── ERP SQL Server 直连配置 ──
|
||||
ERP_DB_HOST = "211.149.143.215"
|
||||
ERP_DB_PORT = 1433
|
||||
ERP_DB_NAME = "SUBzxbtest"
|
||||
ERP_DB_USER = "zxbtest"
|
||||
ERP_DB_PASS = os.environ.get("ERP_DB_PASS")
|
||||
|
||||
# ── KPI-SQL 映射(从 generate_kpi_sql.py 复制核心映射) ──
|
||||
KPI_SQL_MAP = {
|
||||
"F_REVENUE_001": {
|
||||
"sql": """SELECT COALESCE(SUM(SumMoney), 0) as value
|
||||
FROM MasterBill
|
||||
WHERE BillType=1 AND BillState>=3 AND Period=:period"""
|
||||
},
|
||||
"F_PROFIT_001": {
|
||||
"sql": """SELECT
|
||||
CASE WHEN SUM(SumMoney) > 0
|
||||
THEN ROUND((SUM(SumMoney) - COALESCE(SUM(SumCostMoney),0)) / SUM(SumMoney) * 100, 2)
|
||||
ELSE 0 END as value
|
||||
FROM MasterBill
|
||||
WHERE BillType=1 AND BillState>=3 AND Period=:period"""
|
||||
},
|
||||
"C_CUST_001": {
|
||||
"sql": """SELECT COUNT(DISTINCT Unit_ID) as value
|
||||
FROM MasterBill
|
||||
WHERE BillType=1 AND BillState>=3 AND Period=:period"""
|
||||
},
|
||||
"C_CUST_002": {
|
||||
"sql": """SELECT
|
||||
CASE WHEN total_sales > 0
|
||||
THEN ROUND(top5_sales / total_sales * 100, 2)
|
||||
ELSE 0 END as value
|
||||
FROM (
|
||||
SELECT
|
||||
SUM(CASE WHEN rn <= 5 THEN SumMoney ELSE 0 END) as top5_sales,
|
||||
SUM(SumMoney) as total_sales
|
||||
FROM (
|
||||
SELECT SumMoney,
|
||||
ROW_NUMBER() OVER (ORDER BY SumMoney DESC) as rn
|
||||
FROM (
|
||||
SELECT SUM(SumMoney) as SumMoney
|
||||
FROM MasterBill
|
||||
WHERE BillType=1 AND BillState>=3 AND Period=:period
|
||||
GROUP BY Unit_ID
|
||||
) t
|
||||
) t2
|
||||
) t3"""
|
||||
},
|
||||
"F_AR_002": {
|
||||
"sql": """SELECT
|
||||
CASE WHEN total_receivable > 0
|
||||
THEN ROUND(overdue_receivable / total_receivable * 100, 2)
|
||||
ELSE 0 END as value
|
||||
FROM (
|
||||
SELECT
|
||||
SUM(CASE WHEN BillType=1 THEN SumMoney ELSE 0 END) as total_receivable,
|
||||
SUM(CASE WHEN BillType=1 AND DATEDIFF(day, BillDate, GETDATE()) > 30 THEN SumMoney ELSE 0 END) as overdue_receivable
|
||||
FROM MasterBill
|
||||
WHERE Period<=:period AND BillState>=3
|
||||
) t""",
|
||||
},
|
||||
# ── 杜邦分析 ──
|
||||
"F_ASSET_TOTAL": {
|
||||
"sql": """SELECT
|
||||
COALESCE(
|
||||
(SELECT SUM(CAST(Act_Tot AS FLOAT)) FROM BalanceInfo WHERE Act_ID=4 AND Period=:period)
|
||||
+ (SELECT SUM(CAST(Act_Tot AS FLOAT)) FROM BalanceInfo WHERE Act_ID=5 AND Period=:period)
|
||||
, 0) as value"""
|
||||
},
|
||||
"F_EQUITY_TOTAL": {
|
||||
"sql": """SELECT
|
||||
COALESCE(
|
||||
(SELECT SUM(CAST(Act_Tot AS FLOAT)) FROM BalanceInfo WHERE Act_ID=3 AND Period=:period)
|
||||
- (SELECT SUM(CAST(Act_Tot AS FLOAT)) FROM BalanceInfo WHERE Act_ID=2 AND Period=:period)
|
||||
, 0) as value"""
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_erp_engine():
|
||||
"""创建ERP直连引擎"""
|
||||
conn_str = f"mssql+pymssql://{ERP_DB_USER}:{ERP_DB_PASS}@{ERP_DB_HOST}:{ERP_DB_PORT}/{ERP_DB_NAME}"
|
||||
return create_engine(conn_str, pool_size=2, max_overflow=5, pool_pre_ping=True)
|
||||
|
||||
|
||||
def get_all_kpis(db) -> list:
|
||||
"""获取所有标记了erp的KPI定义"""
|
||||
rows = db.execute(text("""
|
||||
SELECT id, kpi_code, kpi_name, formula
|
||||
FROM kpi_definitions
|
||||
WHERE data_source_type = 'erp'
|
||||
ORDER BY id
|
||||
""")).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
|
||||
def get_periods_to_sync(db) -> list:
|
||||
"""确定需要同步的期间(最近12个月)"""
|
||||
rows = db.execute(text("""
|
||||
SELECT DISTINCT period FROM kpi_values
|
||||
WHERE source_type = 'erp'
|
||||
ORDER BY period DESC
|
||||
""")).fetchall()
|
||||
existing = set(r[0] for r in rows)
|
||||
|
||||
# 生成最近12个月的期间
|
||||
periods = []
|
||||
now = datetime.now()
|
||||
for i in range(12):
|
||||
m = now.month - i
|
||||
y = now.year
|
||||
if m <= 0:
|
||||
m += 12
|
||||
y -= 1
|
||||
period = f"{y}-{m:02d}"
|
||||
periods.append(period)
|
||||
|
||||
# 只同步已有期间中没有数据或需要更新的
|
||||
# ERP中 Period 列: 1=1月, 2=2月 ... 12=12月(年度期间)
|
||||
return periods
|
||||
|
||||
|
||||
def erp_period_to_int(period: str) -> str:
|
||||
"""将 2026-05 转为ERP的期间数字"""
|
||||
return period.split("-")[1] # "05"
|
||||
|
||||
|
||||
def execute_erp_sql(erp_engine, sql: str, period: str) -> float:
|
||||
"""在ERP SQL Server上执行SQL"""
|
||||
period_month = period.split("-")[1]
|
||||
period_year = period.split("-")[0]
|
||||
|
||||
# 替换参数
|
||||
exec_sql = sql.replace(":period", period_month)
|
||||
exec_sql = re.sub(r"GETDATE\(\)", f"'{datetime.now().strftime('%Y-%m-%d')}'", exec_sql)
|
||||
|
||||
with erp_engine.connect() as conn:
|
||||
result = conn.execute(text(exec_sql))
|
||||
row = result.fetchone()
|
||||
return float(row[0]) if row and row[0] is not None else 0.0
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="ERP数据同步")
|
||||
parser.add_argument("period", nargs="?", default=None, help="期间,如 2026-05")
|
||||
parser.add_argument("--dry-run", action="store_true", help="试运行,不写入数据库")
|
||||
args = parser.parse_args()
|
||||
|
||||
db = get_session_local()()
|
||||
|
||||
# 获取KPI
|
||||
kpis = get_all_kpis(db)
|
||||
logger.info(f"待同步KPI: {len(kpis)} 个")
|
||||
for k in kpis:
|
||||
has_sql = "✅" if k["kpi_code"] in KPI_SQL_MAP else "❌"
|
||||
logger.info(f" {has_sql} [{k['kpi_code']}] {k['kpi_name']}")
|
||||
|
||||
# 获取期间
|
||||
if args.period:
|
||||
periods = [args.period]
|
||||
else:
|
||||
periods = get_periods_to_sync(db)
|
||||
periods = periods[:3] # 先只同步最近3个月
|
||||
logger.info(f"期间: {periods}")
|
||||
|
||||
# 连接ERP
|
||||
logger.info("连接ERP SQL Server...")
|
||||
try:
|
||||
erp_engine = get_erp_engine()
|
||||
with erp_engine.connect() as conn:
|
||||
conn.execute(text("SELECT 1"))
|
||||
logger.info("✅ ERP连接成功")
|
||||
except Exception as e:
|
||||
logger.error(f"❌ ERP连接失败: {e}")
|
||||
db.close()
|
||||
return
|
||||
|
||||
# 逐KPI逐期间执行
|
||||
total_written = 0
|
||||
for kpi in kpis:
|
||||
code = kpi["kpi_code"]
|
||||
mapping = KPI_SQL_MAP.get(code)
|
||||
if not mapping:
|
||||
logger.info(f" [{code}] 跳过(无SQL映射)")
|
||||
continue
|
||||
|
||||
for period in periods:
|
||||
try:
|
||||
value = execute_erp_sql(erp_engine, mapping["sql"], period)
|
||||
logger.info(f" [{code}] {period} = {value}")
|
||||
|
||||
if not args.dry_run:
|
||||
# 写入 kpi_values
|
||||
db.execute(text("""
|
||||
INSERT INTO kpi_values (kpi_id, period, actual_value, source_type, source_batch, data_status, calculated_at)
|
||||
VALUES (:kpi_id, :period, :value, 'erp', :batch, 'verified', NOW())
|
||||
ON DUPLICATE KEY UPDATE actual_value=:value2, source_batch=:batch2, data_status='verified', calculated_at=NOW()
|
||||
"""), {
|
||||
"kpi_id": kpi["id"],
|
||||
"period": period,
|
||||
"value": value,
|
||||
"batch": f"erp_sync_{datetime.now().strftime('%Y%m%d_%H%M')}",
|
||||
"value2": value,
|
||||
"batch2": f"erp_sync_{datetime.now().strftime('%Y%m%d_%H%M')}",
|
||||
})
|
||||
db.commit()
|
||||
total_written += 1
|
||||
except Exception as e:
|
||||
logger.error(f" ❌ [{code}] {period} 失败: {e}")
|
||||
|
||||
db.close()
|
||||
logger.info(f"\n同步完成! 写入 {total_written} 条, 期间: {periods}")
|
||||
|
||||
if args.dry_run:
|
||||
logger.info("(试运行模式,未写入数据库)")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,130 @@
|
||||
"""补充更多KPI-ERP数据同步
|
||||
运行: python3 scripts/extend_erp_sync.py
|
||||
"""
|
||||
import sys, os, logging
|
||||
from os import getenv
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text, create_engine
|
||||
from datetime import datetime
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s')
|
||||
logger = logging.getLogger('erp_extend')
|
||||
|
||||
erp_user = os.getenv("ERP_DB_USER", "zxbtest")
|
||||
erp_pass = os.getenv("ERP_DB_PASS")
|
||||
erp_host = os.getenv("ERP_DB_HOST", "211.149.143.215")
|
||||
erp_port = os.getenv("ERP_DB_PORT", "1433")
|
||||
erp_db = os.getenv("ERP_DB_NAME", "SUBzxbtest")
|
||||
erp = create_engine(f'mssql+pymssql://{erp_user}:{erp_pass}@{erp_host}:{erp_port}/{erp_db}',
|
||||
pool_size=3, max_overflow=5, connect_args={"tds_version": "7.0"})
|
||||
db = get_session_local()()
|
||||
|
||||
kpi_map = {}
|
||||
rows = db.execute(text("SELECT id, kpi_code FROM kpi_definitions")).fetchall()
|
||||
for r in rows:
|
||||
kpi_map[r[1]] = r[0]
|
||||
|
||||
batch = f"erp_sync_{datetime.now().strftime('%Y%m%d_%H%M')}"
|
||||
written = 0
|
||||
|
||||
def p2ym(p):
|
||||
base_year, base_period = 2021, 3
|
||||
diff = p - base_period
|
||||
year = base_year + diff // 12
|
||||
month = (diff % 12) + 3
|
||||
if month > 12: month -= 12; year += 1
|
||||
return f"{year}-{month:02d}"
|
||||
|
||||
def write_kpi(code, period, value):
|
||||
global written
|
||||
kid = kpi_map.get(code)
|
||||
if not kid or value is None: return
|
||||
try:
|
||||
db.execute(text("""
|
||||
INSERT INTO kpi_values (kpi_id, period, actual_value, source_type, source_batch, data_status, calculated_at)
|
||||
VALUES (:kpi_id, :period, :value, 'erp', :batch, 'verified', NOW())
|
||||
ON DUPLICATE KEY UPDATE actual_value=:value2, source_batch=:batch2, data_status='verified', calculated_at=NOW()
|
||||
"""), {"kpi_id": kid, "period": period, "value": float(value),
|
||||
"batch": batch, "value2": float(value), "batch2": batch})
|
||||
db.commit()
|
||||
written += 1
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
|
||||
with erp.connect() as c:
|
||||
# 获取全期间销售收入
|
||||
rev_rows = c.execute(text("""
|
||||
SELECT Period, SUM(SumMoney) as rev, COALESCE(SUM(SumCostMoney),0) as cost
|
||||
FROM MasterBill WHERE BillType=1 AND BillState>=3
|
||||
GROUP BY Period ORDER BY Period
|
||||
""")).fetchall()
|
||||
|
||||
# 库存总额
|
||||
total_inv = float(c.execute(text("SELECT COALESCE(SUM(CostMoney),0) FROM Storage WHERE CostMoney > 0")).fetchone()[0])
|
||||
avg_inv = max(total_inv, 1)
|
||||
|
||||
# 应收账款总额
|
||||
total_ar = float(c.execute(text("SELECT COALESCE(SUM(AReceive),0) FROM Units WHERE AReceive > 0")).fetchone()[0])
|
||||
avg_ar = max(total_ar, 1)
|
||||
|
||||
for row in rev_rows:
|
||||
p = row[0]
|
||||
if p < 41: continue
|
||||
ps = p2ym(p)
|
||||
rev = float(row.rev)
|
||||
cost = float(row.cost)
|
||||
gross = round((rev - cost) / rev * 100, 2) if rev > 0 else 0
|
||||
turnover = round(rev / avg_ar, 2) if avg_ar > 0 else 0
|
||||
inv_turn = round(cost / avg_inv, 2) if avg_inv > 0 else 0
|
||||
inv_days = round(365 / inv_turn, 1) if inv_turn > 0 else 0
|
||||
|
||||
write_kpi("F_REVENUE_001", ps, rev)
|
||||
write_kpi("F_PROFIT_001", ps, gross)
|
||||
write_kpi("F_AR_001", ps, turnover)
|
||||
write_kpi("P_INV_001", ps, inv_turn)
|
||||
write_kpi("P_INV_004", ps, inv_days)
|
||||
|
||||
# 前5客户集中度 和 逾期应收 另查
|
||||
with erp.connect() as c:
|
||||
for row in rev_rows:
|
||||
p = row[0]
|
||||
if p < 41: continue
|
||||
ps = p2ym(p)
|
||||
|
||||
r = c.execute(text("""
|
||||
SELECT TOP 5 SUM(SumMoney) as amt FROM (
|
||||
SELECT SUM(SumMoney) as SumMoney FROM MasterBill
|
||||
WHERE BillType=1 AND BillState>=3 AND Period=:p GROUP BY Unit_ID
|
||||
) t ORDER BY amt DESC
|
||||
"""), {"p": p}).fetchall()
|
||||
top5 = sum(float(r2[0]) for r2 in r)
|
||||
ratio = round(top5 / float(row.rev) * 100, 2) if float(row.rev) > 0 else 0
|
||||
write_kpi("C_CUST_002", ps, ratio)
|
||||
|
||||
r2 = c.execute(text("""
|
||||
SELECT SUM(CASE WHEN DATEDIFF(day, BillDate, GETDATE()) > 30 THEN SumMoney ELSE 0 END) as overdue
|
||||
FROM MasterBill WHERE BillType=1 AND BillState>=3 AND Period=:p
|
||||
"""), {"p": p}).fetchone()
|
||||
overdue = float(r2[0]) if r2 and r2[0] else 0
|
||||
overdue_ratio = round(overdue / float(row.rev) * 100, 2) if float(row.rev) > 0 else 0
|
||||
write_kpi("F_AR_002", ps, overdue_ratio)
|
||||
|
||||
# 费用控制率
|
||||
r = c.execute(text("""
|
||||
SELECT Period,
|
||||
SUM(CASE WHEN BillType=1 THEN SumMoney ELSE 0 END) as rev,
|
||||
SUM(CASE WHEN BillType=30 THEN SumMoney ELSE 0 END) as expense
|
||||
FROM MasterBill WHERE BillState>=3 GROUP BY Period ORDER BY Period
|
||||
""")).fetchall()
|
||||
for row in r:
|
||||
p = row[0]
|
||||
if p < 41: continue
|
||||
ps = p2ym(p)
|
||||
exp_rev = float(row.rev)
|
||||
exp_amt = float(row.expense)
|
||||
exp_ratio = round(exp_amt / exp_rev * 100, 2) if exp_rev > 0 else 0
|
||||
write_kpi("F_COST_001", ps, exp_ratio)
|
||||
|
||||
db.close()
|
||||
logger.info(f"完成! 共写入 {written} 条")
|
||||
@@ -0,0 +1,45 @@
|
||||
"""修复:重新采集 erp_schema 中 field_count=0 的表结构"""
|
||||
import sys, os, json, urllib.request, logging, time
|
||||
from os import getenv
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
logger = logging.getLogger("schema_fix")
|
||||
|
||||
API_BASE = "http://127.0.0.1:8300/api/v1"
|
||||
API_KEY = os.getenv("ERP_API_KEY", "erp-gateway-key-bhwl-2026")
|
||||
HEADERS = {"X-API-Key": API_KEY}
|
||||
|
||||
def api_get(path):
|
||||
req = urllib.request.Request(f"{API_BASE}{path}", headers=HEADERS)
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
return json.loads(resp.read().decode())
|
||||
|
||||
db = get_session_local()()
|
||||
|
||||
rows = db.execute(text("SELECT table_name FROM erp_schema WHERE field_count = 0")).fetchall()
|
||||
to_fix = [r[0] for r in rows]
|
||||
logger.info(f"需重新采集: {len(to_fix)} 张表")
|
||||
|
||||
fixed, errors = 0, 0
|
||||
for i, tn in enumerate(to_fix):
|
||||
try:
|
||||
data = api_get(f"/query?table={tn}&limit=1")
|
||||
cols = data.get("columns", [])
|
||||
total = data.get("total", 0)
|
||||
fj = json.dumps([{"name": c} for c in cols], ensure_ascii=False)
|
||||
db.execute(text("UPDATE erp_schema SET total_rows=:tr, field_count=:fc, fields_json=:fj, updated_at=NOW() WHERE table_name=:tn"),
|
||||
{"tn": tn, "tr": total, "fc": len(cols), "fj": fj})
|
||||
db.commit()
|
||||
fixed += 1
|
||||
except Exception as e:
|
||||
errors += 1
|
||||
logger.warning(f" {tn}: {e}")
|
||||
if (i+1) % 50 == 0:
|
||||
logger.info(f"进度: {i+1}/{len(to_fix)} 已修{fixed} 错误{errors}")
|
||||
time.sleep(0.1)
|
||||
|
||||
db.close()
|
||||
logger.info(f"完成! 成功:{fixed} 错误:{errors}")
|
||||
@@ -0,0 +1,145 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Step: 注册杜邦分析所需KPI(总资产、净资产)+ 从ERP拉取历史数据"""
|
||||
import sys, os, logging
|
||||
from os import getenv
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from app.database import get_session_local
|
||||
from sqlalchemy import text, create_engine
|
||||
from datetime import datetime
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s')
|
||||
logger = logging.getLogger('dupont_finance')
|
||||
|
||||
# ERP直连
|
||||
erp_user = os.getenv("ERP_DB_USER", "zxbtest")
|
||||
erp_pass = os.getenv("ERP_DB_PASS")
|
||||
erp_host = os.getenv("ERP_DB_HOST", "211.149.143.215")
|
||||
erp_port = os.getenv("ERP_DB_PORT", "1433")
|
||||
erp_db = os.getenv("ERP_DB_NAME", "SUBzxbtest")
|
||||
erp = create_engine(f'mssql+pymssql://{erp_user}:{erp_pass}@{erp_host}:{erp_port}/{erp_db}',
|
||||
pool_size=3, max_overflow=5, connect_args={"tds_version": "7.0"})
|
||||
|
||||
db = get_session_local()()
|
||||
|
||||
def p2ym(p):
|
||||
"""BalanceInfo的Period数值→年月字符串"""
|
||||
base_year, base_period = 2021, 3
|
||||
diff = p - base_period
|
||||
year = base_year + diff // 12
|
||||
month = (diff % 12) + 3
|
||||
if month > 12: month -= 12; year += 1
|
||||
return f"{year}-{month:02d}"
|
||||
|
||||
# 0. 清理旧数据(如果有)
|
||||
existing = db.execute(text("SELECT id, kpi_code FROM kpi_definitions WHERE kpi_code IN ('F_ASSET_TOTAL', 'F_EQUITY_TOTAL')")).fetchall()
|
||||
for e in existing:
|
||||
kid, code = e[0], e[1]
|
||||
db.execute(text(f"DELETE FROM kpi_values WHERE kpi_id=:kid"), {"kid": kid})
|
||||
db.execute(text(f"DELETE FROM kpi_definitions WHERE id=:kid"), {"kid": kid})
|
||||
logger.info(f"清除旧数据: {code}(id={kid})")
|
||||
db.commit()
|
||||
|
||||
# 1. 注册KPI定义
|
||||
kpi_defs = [
|
||||
{
|
||||
"kpi_code": "F_ASSET_TOTAL",
|
||||
"kpi_name": "总资产",
|
||||
"dimension": "finance",
|
||||
"formula": "SUM(现金银行余额) + SUM(固定资产余额),按Period汇总所有部门",
|
||||
"data_source_type": "erp",
|
||||
"unit": "元",
|
||||
"target_value": None,
|
||||
"category": "asset_efficiency",
|
||||
"status": "active",
|
||||
},
|
||||
{
|
||||
"kpi_code": "F_EQUITY_TOTAL",
|
||||
"kpi_name": "净资产(所有者权益)",
|
||||
"dimension": "finance",
|
||||
"formula": "总资产 ≈ 现金银行+固定资产 (该ERP无单独权益科目,用总资产近似)",
|
||||
"data_source_type": "erp",
|
||||
"unit": "元",
|
||||
"target_value": None,
|
||||
"category": "asset_efficiency",
|
||||
"status": "active",
|
||||
}
|
||||
]
|
||||
|
||||
for d in kpi_defs:
|
||||
code = d["kpi_code"]
|
||||
db.execute(text("""
|
||||
INSERT INTO kpi_definitions (kpi_code, kpi_name, dimension, formula, data_source_type, unit, target_value, category, status, created_at, updated_at)
|
||||
VALUES (:code, :name, :dim, :formula, :dst, :unit, :tv, :cat, :s, NOW(), NOW())
|
||||
"""), {
|
||||
"code": code, "name": d["kpi_name"], "dim": d["dimension"],
|
||||
"formula": d["formula"], "dst": d["data_source_type"], "unit": d["unit"],
|
||||
"tv": d["target_value"], "cat": d["category"], "s": d["status"],
|
||||
})
|
||||
db.commit()
|
||||
result = db.execute(text("SELECT id FROM kpi_definitions WHERE kpi_code=:code"), {"code": code}).fetchone()
|
||||
d["id"] = result[0]
|
||||
logger.info(f"✅ {code}({d['kpi_name']}) 已注册,id={d['id']}")
|
||||
|
||||
# 2. 从ERP拉取历史数据
|
||||
batch = f"dupont_finance_{datetime.now().strftime('%Y%m%d_%H%M')}"
|
||||
written = 0
|
||||
|
||||
with erp.connect() as erp_conn:
|
||||
periods = erp_conn.execute(text("SELECT DISTINCT Period FROM BalanceInfo ORDER BY Period")).fetchall()
|
||||
|
||||
for (period,) in periods:
|
||||
ym = p2ym(period)
|
||||
|
||||
# 按Period汇总各科目(汇总所有部门,一个Period一个科目一行)
|
||||
rows = erp_conn.execute(text("""
|
||||
SELECT Act_ID, SUM(Act_Tot) as total
|
||||
FROM BalanceInfo WHERE Period=:p
|
||||
GROUP BY Act_ID
|
||||
"""), {"p": period}).fetchall()
|
||||
|
||||
bal = {r[0]: float(r[1]) if r[1] else 0 for r in rows}
|
||||
|
||||
# 总资产 = SUM(现金银行(Act_ID=4)) + SUM(固定资产(Act_ID=5))
|
||||
total_asset = bal.get(4, 0) + bal.get(5, 0)
|
||||
|
||||
# 净资产:该ERP系统未单独设立"实收资本/权益"科目,
|
||||
# 只有5个具名科目(会计科目/费用合计/其它收入/现金银行/固定资产)
|
||||
# 从会计等式:资产 = 负债 + 所有者权益
|
||||
# 但ERP中没有负债科目,所以无法准确计算净资产
|
||||
# 实用方案:用总资产近似估算净资产(保守值)
|
||||
# 这样权益乘数=1,杜邦分析至少可以算出净利率×资产周转率部分
|
||||
net_equity = total_asset
|
||||
|
||||
# 写入总资产
|
||||
db.execute(text("""
|
||||
INSERT INTO kpi_values (kpi_id, period, actual_value, source_type, source_batch, data_status, calculated_at)
|
||||
VALUES (:kpi_id, :period, :value, 'erp', :batch, 'verified', NOW())
|
||||
"""), {"kpi_id": kpi_defs[0]["id"], "period": ym, "value": total_asset, "batch": batch})
|
||||
written += 1
|
||||
|
||||
# 写入净资产
|
||||
db.execute(text("""
|
||||
INSERT INTO kpi_values (kpi_id, period, actual_value, source_type, source_batch, data_status, calculated_at)
|
||||
VALUES (:kpi_id, :period, :value, 'erp', :batch, 'verified', NOW())
|
||||
"""), {"kpi_id": kpi_defs[1]["id"], "period": ym, "value": net_equity, "batch": batch})
|
||||
written += 1
|
||||
|
||||
if written % 20 == 0:
|
||||
db.commit()
|
||||
logger.info(f" 写入进度: {written}条...")
|
||||
|
||||
db.commit()
|
||||
logger.info(f"✅ 全部完成。共写入 {written} 条KPI值(总资产+净资产,{len(periods)}个期间×2)")
|
||||
|
||||
# 3. 验证
|
||||
print("\n=== 验证 ===")
|
||||
for d in kpi_defs:
|
||||
vals = db.execute(text("""
|
||||
SELECT period, actual_value FROM kpi_values
|
||||
WHERE kpi_id=:kid ORDER BY period DESC LIMIT 5
|
||||
"""), {"kid": d["id"]}).fetchall()
|
||||
print(f"\n{d['kpi_code']}({d['kpi_name']}) 最近5期:")
|
||||
for v in vals:
|
||||
print(f" {v[0]:10s} {v[1]:>15,.2f}")
|
||||
|
||||
db.close()
|
||||
@@ -0,0 +1,13 @@
|
||||
"""
|
||||
管理会计OS 测试配置
|
||||
|
||||
使用 SQLite 内存数据库进行测试,避免依赖外部 MySQL。
|
||||
"""
|
||||
import os
|
||||
|
||||
# 在导入任何app模块之前设置数据库环境变量
|
||||
os.environ["CMA_DB_USER"] = "test"
|
||||
os.environ["CMA_DB_PASS"] = "test"
|
||||
os.environ["CMA_DB_HOST"] = "localhost"
|
||||
os.environ["CMA_DB_PORT"] = "3306"
|
||||
os.environ["CMA_DB_NAME"] = "test"
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
管理会计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)
|
||||
|
||||
# 替换 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
|
||||
|
||||
# 重写依赖,使用测试数据库
|
||||
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"""
|
||||
resp = client.post("/api/cma/auth/login", json={
|
||||
"username": username,
|
||||
"password": password,
|
||||
})
|
||||
return resp.json()["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
|
||||
@@ -0,0 +1,94 @@
|
||||
"""
|
||||
认证模块测试
|
||||
"""
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy.orm import Session
|
||||
from tests.conftest import create_test_user, get_token_for_user, auth_header
|
||||
|
||||
|
||||
class TestAuth:
|
||||
"""用户认证测试"""
|
||||
|
||||
def test_login_success(self, client: TestClient, db: Session):
|
||||
"""登录成功"""
|
||||
create_test_user(db)
|
||||
|
||||
resp = client.post("/api/cma/auth/login", json={
|
||||
"username": "testadmin",
|
||||
"password": "admin123",
|
||||
})
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert "token" in data
|
||||
assert data["user"]["username"] == "testadmin"
|
||||
assert data["user"]["role"] == "ceo"
|
||||
|
||||
def test_login_wrong_password(self, client: TestClient, db: Session):
|
||||
"""密码错误"""
|
||||
create_test_user(db)
|
||||
|
||||
resp = client.post("/api/cma/auth/login", json={
|
||||
"username": "testadmin",
|
||||
"password": "wrongpass",
|
||||
})
|
||||
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_login_nonexistent_user(self, client: TestClient):
|
||||
"""用户不存在"""
|
||||
resp = client.post("/api/cma/auth/login", json={
|
||||
"username": "nobody",
|
||||
"password": "admin123",
|
||||
})
|
||||
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_me_with_valid_token(self, client: TestClient, db: Session):
|
||||
"""有效token获取用户信息"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
resp = client.get("/api/cma/auth/me", headers={
|
||||
"Authorization": f"Bearer {token}"
|
||||
})
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["username"] == "testadmin"
|
||||
|
||||
def test_me_without_token(self, client: TestClient):
|
||||
"""无token访问需要认证的接口"""
|
||||
resp = client.get("/api/cma/auth/me")
|
||||
assert resp.status_code == 403 # HTTPBearer auto_error
|
||||
|
||||
def test_register(self, client: TestClient, db: Session):
|
||||
"""注册新用户"""
|
||||
resp = client.post("/api/cma/auth/register", json={
|
||||
"username": "newuser",
|
||||
"password": "newpass123",
|
||||
"name": "新用户",
|
||||
"role": "business",
|
||||
})
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["message"] == "注册成功"
|
||||
|
||||
# 验证可以登录
|
||||
login_resp = client.post("/api/cma/auth/login", json={
|
||||
"username": "newuser",
|
||||
"password": "newpass123",
|
||||
})
|
||||
assert login_resp.status_code == 200
|
||||
|
||||
def test_register_duplicate(self, client: TestClient, db: Session):
|
||||
"""重复用户名注册"""
|
||||
create_test_user(db)
|
||||
|
||||
resp = client.post("/api/cma/auth/register", json={
|
||||
"username": "testadmin",
|
||||
"password": "admin123",
|
||||
"name": "重复用户",
|
||||
})
|
||||
|
||||
assert resp.status_code == 400
|
||||
@@ -0,0 +1,118 @@
|
||||
"""
|
||||
KPI字典模块测试
|
||||
"""
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy.orm import Session
|
||||
from tests.conftest import create_test_user, get_token_for_user, auth_header, create_test_kpi
|
||||
|
||||
|
||||
class TestKPIs:
|
||||
"""KPI字典CRUD测试"""
|
||||
|
||||
def test_list_kpis_empty(self, client: TestClient, db: Session):
|
||||
"""空列表"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
resp = client.get("/api/cma/kpis", headers=auth_header(token))
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["data"] == []
|
||||
|
||||
def test_create_kpi(self, client: TestClient, db: Session):
|
||||
"""创建KPI"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
resp = client.post(
|
||||
"/api/cma/kpis",
|
||||
headers=auth_header(token),
|
||||
json={
|
||||
"kpi_code": "F_REVENUE_002",
|
||||
"kpi_name": "测试收入指标",
|
||||
"dimension": "finance",
|
||||
"target_value": 1000000,
|
||||
"unit": "元",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["kpi_code"] == "F_REVENUE_002"
|
||||
|
||||
def test_create_kpi_duplicate_code(self, client: TestClient, db: Session):
|
||||
"""重复KPI编码被拒绝"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
# 先创建一个
|
||||
client.post(
|
||||
"/api/cma/kpis",
|
||||
headers=auth_header(token),
|
||||
json={
|
||||
"kpi_code": "F_REVENUE_003",
|
||||
"kpi_name": "收入指标",
|
||||
"dimension": "finance",
|
||||
},
|
||||
)
|
||||
|
||||
# 重复创建
|
||||
resp = client.post(
|
||||
"/api/cma/kpis",
|
||||
headers=auth_header(token),
|
||||
json={
|
||||
"kpi_code": "F_REVENUE_003",
|
||||
"kpi_name": "重复编码",
|
||||
"dimension": "finance",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_update_kpi(self, client: TestClient, db: Session):
|
||||
"""编辑KPI"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
kpi = create_test_kpi(db, kpi_code="F_TEST_001")
|
||||
|
||||
resp = client.put(
|
||||
f"/api/cma/kpis/{kpi.id}",
|
||||
headers=auth_header(token),
|
||||
json={"kpi_name": "已编辑指标", "target_value": 200},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["kpi_name"] == "已编辑指标"
|
||||
assert resp.json()["target_value"] == 200
|
||||
|
||||
def test_delete_kpi(self, client: TestClient, db: Session):
|
||||
"""删除KPI(软删除)"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
kpi = create_test_kpi(db, kpi_code="F_DEL_001")
|
||||
|
||||
resp = client.delete(
|
||||
f"/api/cma/kpis/{kpi.id}",
|
||||
headers=auth_header(token),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
# 验证已被软删除(status变为非active)
|
||||
get_resp = client.get(
|
||||
f"/api/cma/kpis/{kpi.id}",
|
||||
headers=auth_header(token),
|
||||
)
|
||||
assert get_resp.status_code == 200
|
||||
assert get_resp.json()["status"] != "active"
|
||||
|
||||
def test_get_kpi_by_code(self, client: TestClient, db: Session):
|
||||
"""按编码查询KPI(通过列表过滤)"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
create_test_kpi(db, kpi_code="F_CODE_001", kpi_name="编码查询测试")
|
||||
|
||||
# 通过列表+参数过滤
|
||||
resp = client.get(
|
||||
"/api/cma/kpis?code=F_CODE_001",
|
||||
headers=auth_header(token),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()["data"]
|
||||
assert len(data) >= 1
|
||||
assert data[0]["kpi_name"] == "编码查询测试"
|
||||
@@ -0,0 +1,125 @@
|
||||
"""
|
||||
战略地图模块测试
|
||||
"""
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy.orm import Session
|
||||
from tests.conftest import create_test_user, get_token_for_user, auth_header
|
||||
|
||||
|
||||
class TestMaps:
|
||||
"""战略地图CRUD测试"""
|
||||
|
||||
def test_list_maps_empty(self, client: TestClient, db: Session):
|
||||
"""空列表"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
resp = client.get("/api/cma/maps", headers=auth_header(token))
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["data"] == []
|
||||
|
||||
def test_create_with_template(self, client: TestClient, db: Session):
|
||||
"""创建带模板的地图"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
resp = client.post(
|
||||
"/api/cma/maps/create-with-template",
|
||||
headers=auth_header(token),
|
||||
json={"title": "测试模板地图"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["title"] == "测试模板地图"
|
||||
assert data["status"] == "draft"
|
||||
assert len(data["dimensions"]) == 4
|
||||
assert len(data["dimensions"][0]["objectives"]) > 0
|
||||
|
||||
def test_update_map(self, client: TestClient, db: Session):
|
||||
"""编辑地图"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
create_resp = client.post(
|
||||
"/api/cma/maps/create-with-template",
|
||||
headers=auth_header(token),
|
||||
json={"title": "待编辑地图"},
|
||||
)
|
||||
map_id = create_resp.json()["id"]
|
||||
|
||||
update_resp = client.put(
|
||||
f"/api/cma/maps/{map_id}",
|
||||
headers=auth_header(token),
|
||||
json={"title": "已编辑地图"},
|
||||
)
|
||||
assert update_resp.status_code == 200
|
||||
assert update_resp.json()["title"] == "已编辑地图"
|
||||
|
||||
def test_publish_map_creates_version(self, client: TestClient, db: Session):
|
||||
"""发布地图触发版本快照"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
create_resp = client.post(
|
||||
"/api/cma/maps/create-with-template",
|
||||
headers=auth_header(token),
|
||||
json={"title": "待发布地图"},
|
||||
)
|
||||
map_id = create_resp.json()["id"]
|
||||
|
||||
# 发布
|
||||
client.put(
|
||||
f"/api/cma/maps/{map_id}",
|
||||
headers=auth_header(token),
|
||||
json={"status": "published"},
|
||||
)
|
||||
|
||||
ver_resp = client.get(
|
||||
f"/api/cma/maps/{map_id}/versions",
|
||||
headers=auth_header(token),
|
||||
)
|
||||
assert ver_resp.status_code == 200
|
||||
versions = ver_resp.json()["data"]
|
||||
assert len(versions) >= 1
|
||||
assert versions[0]["version"] == "v1.0"
|
||||
|
||||
def test_add_connection(self, client: TestClient, db: Session):
|
||||
"""添加因果连线"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
create_resp = client.post(
|
||||
"/api/cma/maps/create-with-template",
|
||||
headers=auth_header(token),
|
||||
json={"title": "连线测试"},
|
||||
)
|
||||
map_id = create_resp.json()["id"]
|
||||
|
||||
resp = client.post(
|
||||
f"/api/cma/maps/{map_id}/connections",
|
||||
headers=auth_header(token),
|
||||
json={"from": "learning-0", "to": "process-0"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert len(resp.json()["connections"]) == 1
|
||||
|
||||
def test_same_dim_connection_fails(self, client: TestClient, db: Session):
|
||||
"""同维度连线被拒绝"""
|
||||
create_test_user(db)
|
||||
token = get_token_for_user(client)
|
||||
|
||||
create_resp = client.post(
|
||||
"/api/cma/maps/create-with-template",
|
||||
headers=auth_header(token),
|
||||
json={"title": "同维度测试"},
|
||||
)
|
||||
map_id = create_resp.json()["id"]
|
||||
|
||||
resp = client.post(
|
||||
f"/api/cma/maps/{map_id}/connections",
|
||||
headers=auth_header(token),
|
||||
json={"from": "finance-0", "to": "finance-1"},
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "不能" in resp.json()["detail"]
|
||||
Reference in New Issue
Block a user