"""数据对接 API""" import pandas as pd import io, json, hashlib, re from datetime import datetime from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File from sqlalchemy.orm import Session from sqlalchemy import func from typing import Optional from app.database import get_db from app.auth_middleware import require_auth, require_role from app.models import KPIValue, DataSourceConfig, OperationLog, KPIDefinition router = APIRouter(prefix="/api/cma/data", tags=["数据对接"], dependencies=[Depends(require_role("ceo", "finance", "it"))], ) @router.post("/import-excel") async def import_excel(file: UploadFile = File(...), kpi_col: str = Query("kpi_code", description="Excel中KPI编码列名"), period_col: str = Query("period", description="Excel中期间列名"), value_col: str = Query("actual_value", description="Excel中实际值列名"), default_period: str = Query(None, description="如文件无期间列,统一使用此值"), db: Session = Depends(get_db)): content = await file.read() df = pd.read_excel(io.BytesIO(content)) required = [kpi_col, value_col] if not default_period: required.append(period_col) missing = [c for c in required if c not in df.columns] if missing: raise HTTPException(400, f"Excel缺少列: {missing}。当前文件列: {list(df.columns)}") if len(df) == 0: raise HTTPException(400, "Excel文件为空,没有数据行") from app.models import KPIDefinition kpi_map = {k.kpi_code: k.id for k in db.query(KPIDefinition).all()} batch = hashlib.md5(str(datetime.now().timestamp()).encode()).hexdigest()[:12] count = 0 skipped = [] for idx, row in df.iterrows(): kpi_code = str(row.get(kpi_col, "")).strip() period = str(row.get(period_col, default_period or "")).strip() if period_col in df.columns else (default_period or "").strip() value = row.get(value_col) if not kpi_code or not period or pd.isna(value): skipped.append(f"第{idx+2}行: 缺少必填字段") continue kid = kpi_map.get(kpi_code) if not kid: skipped.append(f"第{idx+2}行: KPI编码「{kpi_code}」不存在") continue db.add(KPIValue( kpi_id=kid, period=period, actual_value=float(value), source_type="excel", source_batch=batch, data_status="verified", )) count += 1 db.commit() msg = f"✅ 导入成功 {count} 条数据" if skipped: msg += f",{len(skipped)}条跳过:\n" + "\n".join(skipped[:10]) if len(skipped) > 10: msg += f"\n...还有{len(skipped)-10}条" return {"message": msg, "batch": batch, "total": count, "skipped": len(skipped)} # ── 智能导入(BOT自动识别,无需手动映射) ── _SMART_MAP = { # KPI编码列匹配模式 → 标准kpi_code "kpi_code_patterns": [ re.compile(r'^(kpi_?code|指标编码|编码)$', re.I), re.compile(r'^(科目|项目|账户|报表项目|项目名称)$'), re.compile(r'^(指标名称?|kpi名称?|name)$', re.I), ], # 期间列匹配 "period_patterns": [ re.compile(r'^(period|期间|月份?|年月|日期|会计期间)$', re.I), re.compile(r'^(报表期[间]?|所属期)$'), ], # 数值列匹配 "value_patterns": [ re.compile(r'^(actual_?value|数值|实际值|实际金额)$', re.I), re.compile(r'^(本期金额|本月数|本期|期末余额|期末数)$'), re.compile(r'^(金额|数据|value)$', re.I), ], # 文件名→期间提取 "period_in_filename": re.compile(r'[-_]?(\d{4})[-_]?(\d{1,2})'), # 文件名→报表类型 "statement_types": { "利润表": "PL", "利润": "PL", "income": "PL", "现金流量表": "CF", "现金流": "CF", "cashflow": "CF", "cash_flow": "CF", "资产负债表": "BS", "资产负": "BS", "balance": "BS", }, } def _smart_detect_kpi_col(cols: list[str]) -> str | None: for pat in _SMART_MAP["kpi_code_patterns"]: for c in cols: if pat.match(c.strip()): return c return None def _smart_detect_period_col(cols: list[str]) -> str | None: for pat in _SMART_MAP["period_patterns"]: for c in cols: if pat.match(c.strip()): return c return None def _smart_detect_value_col(cols: list[str]) -> str | None: for pat in _SMART_MAP["value_patterns"]: for c in cols: if pat.match(c.strip()): return c return None def _smart_extract_period_from_filename(filename: str) -> str | None: m = _SMART_MAP["period_in_filename"].search(filename) if m: return f"{m.group(1)}-{int(m.group(2)):02d}" return None def _smart_detect_statement_type(filename: str) -> str | None: for kw, tp in _SMART_MAP["statement_types"].items(): if kw in filename: return tp return None @router.post("/import-excel-smart") async def import_excel_smart(file: UploadFile = File(...), db: Session = Depends(get_db)): """智能导入 — BOT自动识别列名/期间/报表类型,无需手动映射""" content = await file.read() fname = file.filename or "未知文件" try: df = pd.read_excel(io.BytesIO(content)) except Exception as e: raise HTTPException(400, f"无法读取Excel文件: {e}") if len(df) == 0: raise HTTPException(400, "Excel文件为空") cols = list(df.columns) if len(cols) < 2: raise HTTPException(400, f"Excel列数过少: {cols}") # 4. 智能检测列 kpi_col = _smart_detect_kpi_col(cols) or cols[0] value_col = _smart_detect_value_col(cols) or cols[-1] period_col = _smart_detect_period_col(cols) # 5. 从文件名提取期间 period = _smart_extract_period_from_filename(fname) if not period_col else None # 6. 检测报表类型(用于自动生成KPI编码前缀) stype = _smart_detect_statement_type(fname) # 7. 预加载KPI字典 from app.models import KPIDefinition kpis = {k.kpi_code: k for k in db.query(KPIDefinition).all()} known_codes = set(kpis.keys()) # 构建别名映射(去掉空格/大小写/特殊字符) alias_map: dict[str, str] = {} for code in known_codes: clean = re.sub(r'[\s\-_()()]', '', code).lower() alias_map[clean] = code # 8. 遍历导入 batch = hashlib.md5(str(datetime.now().timestamp()).encode()).hexdigest()[:12] imported = 0 skipped_rows = [] for idx, row in df.iterrows(): raw_kpi = str(row.get(kpi_col, "")).strip() raw_val = row.get(value_col) raw_period = str(row.get(period_col, period or "")).strip() if period_col else (period or "") if not raw_kpi or pd.isna(raw_val): skipped_rows.append(f"第{idx+2}行: 缺数据") continue if not raw_period: skipped_rows.append(f"第{idx+2}行: 无法确定期间") continue # 智能匹配KPI编码 kpi_code = None if raw_kpi in known_codes: kpi_code = raw_kpi else: # 别名匹配 clean_key = re.sub(r'[\s\-_()()]', '', raw_kpi).lower() kpi_code = alias_map.get(clean_key) # 模糊匹配(中文科目名→KPI编码) if not kpi_code: for code, kpi_obj in kpis.items(): if raw_kpi in kpi_obj.kpi_name or kpi_obj.kpi_name in raw_kpi: kpi_code = code break if not kpi_code: skipped_rows.append(f"第{idx+2}行: 「{raw_kpi}」未匹配到KPI") continue try: val = float(raw_val) except: skipped_rows.append(f"第{idx+2}行: 数值格式错误「{raw_val}」") continue db.add(KPIValue( kpi_id=kpis[kpi_code].id, period=raw_period, actual_value=val, source_type="excel", source_batch=batch, data_status="verified", )) imported += 1 db.commit() # 9. 返回汇总 stype_label = {"PL": "利润表", "CF": "现金流量表", "BS": "资产负债表"}.get(stype or "", "数据表") msg = f"✅ {stype_label}识别成功,导入{imported}条" if skipped_rows: msg += f",{len(skipped_rows)}条跳过:\n" + "\n".join(skipped_rows[:8]) if len(skipped_rows) > 8: msg += f"\n...还有{len(skipped_rows) - 8}条" return {"message": msg, "batch": batch, "total": imported, "skipped": len(skipped_rows)} @router.get("/sources") def list_sources(db: Session = Depends(get_db)): sources = db.query(DataSourceConfig).all() return {"data": [{c.name: getattr(s, c.name) for c in DataSourceConfig.__table__.columns} for s in sources]} @router.post("/sources") def create_source(data: dict, db: Session = Depends(get_db)): source = DataSourceConfig( name=data.get("name", ""), source_type=data.get("source_type", "manual"), api_endpoint=data.get("api_endpoint"), api_key=data.get("api_key"), query_sql=data.get("query_sql"), sync_type=data.get("sync_type", "manual"), status="active", ) db.add(source) db.commit() db.refresh(source) # 操作日志 db.add(OperationLog(action="create_source", target_type="source", detail=source.name)) db.commit() return {"data": {c.name: getattr(source, c.name) for c in DataSourceConfig.__table__.columns}} @router.put("/sources/{source_id}") def update_source(source_id: int, data: dict, db: Session = Depends(get_db)): source = db.query(DataSourceConfig).filter(DataSourceConfig.id == source_id).first() if not source: raise HTTPException(404, "数据源不存在") for key in ["name", "source_type", "api_endpoint", "api_key", "query_sql", "sync_type", "status"]: if key in data: setattr(source, key, data[key]) db.commit() db.refresh(source) db.add(OperationLog(action="update_source", target_type="source", detail=source.name)) db.commit() return {"data": {c.name: getattr(source, c.name) for c in DataSourceConfig.__table__.columns}} @router.delete("/sources/{source_id}") def delete_source(source_id: int, db: Session = Depends(get_db)): source = db.query(DataSourceConfig).filter(DataSourceConfig.id == source_id).first() if not source: raise HTTPException(404, "数据源不存在") db.add(OperationLog(action="delete_source", target_type="source", detail=source.name)) db.delete(source) db.commit() return {"message": "删除成功"} @router.get("/sync-kpis") def sync_kpis_from_erp(db: Session = Depends(get_db)): """从ERP数据源同步KPI值(调用erp_sync模块)""" from scripts.erp_sync import run_sync import traceback from datetime import datetime as dt try: # 获取所有标记为erp数据源的KPI erp_kpis = db.query(KPIDefinition).filter( KPIDefinition.status == "active", KPIDefinition.data_source_type == "erp", ).all() kpi_count = len(erp_kpis) kpi_codes = [k.kpi_code for k in erp_kpis] # 执行同步 (dry_run=False, use_api=False 使用本地fallback) run_sync(dry_run=False, kpi_codes=kpi_codes, use_api=False) # 记录操作日志 log = OperationLog( action="sync_kpis", target_type="kpi", detail=f"ERP同步: {kpi_count}个KPI, 编码: {', '.join(kpi_codes[:10])}{'...' if kpi_count > 10 else ''}", ) db.add(log) db.commit() return { "message": f"ERP数据同步完成", "total_kpis": kpi_count, "kpi_codes": kpi_codes, "synced_at": dt.now().isoformat(), } except Exception as e: log = OperationLog( action="sync_kpis_error", target_type="kpi", detail=f"ERP同步失败: {str(e)[:500]}", ) db.add(log) db.commit() raise HTTPException(500, f"ERP同步失败: {str(e)}")