feat: 导入 CapacityReport 初始源码
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
"""API package."""
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
"""API router package."""
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
from fastapi import APIRouter, Body
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from app.auth import create_jwt_token, get_auth_config, save_auth_password
|
||||
|
||||
|
||||
router = APIRouter(tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/api/login")
|
||||
async def login(
|
||||
username: str = Body(..., embed=True),
|
||||
password: str = Body(..., embed=True),
|
||||
):
|
||||
auth = get_auth_config()
|
||||
if username != auth["username"] or password != auth["password"]:
|
||||
return JSONResponse(status_code=401, content={"detail": "账号或密码错误"})
|
||||
|
||||
token = create_jwt_token({"user": username})
|
||||
return {"success": True, "token": token}
|
||||
|
||||
|
||||
@router.post("/api/change-password")
|
||||
async def change_password(
|
||||
current_password: str = Body(..., embed=True),
|
||||
new_password: str = Body(..., embed=True),
|
||||
):
|
||||
auth = get_auth_config()
|
||||
if current_password != auth["password"]:
|
||||
return JSONResponse(status_code=400, content={"detail": "当前密码错误"})
|
||||
if len(new_password) < 4:
|
||||
return JSONResponse(status_code=400, content={"detail": "新密码长度不能少于 4 位"})
|
||||
|
||||
save_auth_password(new_password)
|
||||
return {"success": True, "message": "密码修改成功"}
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.config import CACHE_DIR
|
||||
from app.utils.files import format_size, get_dir_size
|
||||
|
||||
|
||||
router = APIRouter(tags=["cache"])
|
||||
|
||||
|
||||
@router.get("/api/cache/size")
|
||||
async def get_cache_size():
|
||||
if not CACHE_DIR.exists():
|
||||
return {
|
||||
"success": True,
|
||||
"size_bytes": 0,
|
||||
"size_formatted": "0 B",
|
||||
"file_count": 0,
|
||||
"dir_count": 0,
|
||||
}
|
||||
|
||||
total_size = 0
|
||||
file_count = 0
|
||||
dir_count = 0
|
||||
|
||||
try:
|
||||
for item in CACHE_DIR.iterdir():
|
||||
if item.name == "history.json":
|
||||
continue
|
||||
if item.is_dir():
|
||||
dir_count += 1
|
||||
elif item.is_file():
|
||||
file_count += 1
|
||||
total_size += get_dir_size(item)
|
||||
except (PermissionError, OSError) as exc:
|
||||
return {"success": False, "error": str(exc), "size_formatted": "计算失败"}
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"size_bytes": total_size,
|
||||
"size_formatted": format_size(total_size),
|
||||
"file_count": file_count,
|
||||
"dir_count": dir_count,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from threading import Thread
|
||||
|
||||
from fastapi import APIRouter, Body, File, HTTPException, UploadFile
|
||||
|
||||
from app import state
|
||||
from app.api.routers.task_runtime import set_task_stage
|
||||
from app.config import AppConfig, CACHE_DIR, CELLDATA_SCRIPT
|
||||
from app.processor import ProcessLogger
|
||||
from app.services.cell_data import CellDataProcessor, execute_celldata_script, refresh_cell_data
|
||||
from app.utils.files import safe_relative_path
|
||||
|
||||
router = APIRouter(tags=["cell-data"])
|
||||
|
||||
|
||||
@router.post("/api/cell-data/process/start")
|
||||
async def start_cell_data_processing():
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
|
||||
|
||||
task_id = "cell_data_" + datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
work_dir = CACHE_DIR / task_id
|
||||
work_dir.mkdir(parents=True, exist_ok=True)
|
||||
logs: list[str] = []
|
||||
current_stage = "locating"
|
||||
|
||||
def log_callback(message: str) -> None:
|
||||
logs.append(message)
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
def stage_callback(stage: str) -> None:
|
||||
nonlocal current_stage
|
||||
current_stage = stage
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
logger = ProcessLogger(
|
||||
log_file=work_dir / "log.txt",
|
||||
callback=log_callback,
|
||||
stage_callback=stage_callback,
|
||||
)
|
||||
app_config = state.current_config()
|
||||
state.processing_tasks[task_id] = {"logs": [], "status": "processing", "stage": current_stage}
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": current_stage,
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
thread = Thread(target=_run_cell_data_processing, args=(task_id, work_dir, logger, app_config), daemon=True)
|
||||
thread.start()
|
||||
return {"success": True, "message": "CellData 处理已启动", "task_id": task_id, "stage": current_stage}
|
||||
|
||||
|
||||
@router.post("/api/cell-data/process/upload")
|
||||
async def upload_and_start_cell_data_processing(files: list[UploadFile] = File(...)):
|
||||
if not files:
|
||||
raise HTTPException(status_code=400, detail="没有上传文件")
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
|
||||
|
||||
task_id = "cell_data_" + datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
work_dir = CACHE_DIR / task_id
|
||||
upload_dir = work_dir / "uploads"
|
||||
upload_dir.mkdir(parents=True, exist_ok=True)
|
||||
saved_count = 0
|
||||
for file in files:
|
||||
if not file.filename:
|
||||
continue
|
||||
target = upload_dir / safe_relative_path(file.filename)
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
target.write_bytes(await file.read())
|
||||
saved_count += 1
|
||||
if saved_count == 0:
|
||||
raise HTTPException(status_code=400, detail="没有有效上传文件")
|
||||
|
||||
logs: list[str] = []
|
||||
current_stage = "parsing"
|
||||
|
||||
def log_callback(message: str) -> None:
|
||||
logs.append(message)
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
def stage_callback(stage: str) -> None:
|
||||
nonlocal current_stage
|
||||
current_stage = stage
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
logger = ProcessLogger(
|
||||
log_file=work_dir / "log.txt",
|
||||
callback=log_callback,
|
||||
stage_callback=stage_callback,
|
||||
)
|
||||
app_config = state.current_config()
|
||||
state.processing_tasks[task_id] = {"logs": [], "status": "processing", "stage": current_stage}
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": current_stage,
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
thread = Thread(target=_run_uploaded_cell_data_processing, args=(task_id, upload_dir, work_dir, logger, app_config), daemon=True)
|
||||
thread.start()
|
||||
return {
|
||||
"success": True,
|
||||
"message": "CellData 处理已启动",
|
||||
"task_id": task_id,
|
||||
"stage": current_stage,
|
||||
"file_count": saved_count,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/api/cell-data/process/status")
|
||||
async def get_cell_data_processing_status(task_id: str = Body(..., embed=True)):
|
||||
if task_id in state.processing_tasks:
|
||||
task = state.processing_tasks[task_id]
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": task.get("status", "processing"),
|
||||
"stage": task.get("stage", "processing"),
|
||||
"logs": task.get("logs", []),
|
||||
"error": task.get("error"),
|
||||
"result": task.get("result"),
|
||||
"elapsed_time": task.get("elapsed_time"),
|
||||
}
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
|
||||
def _run_cell_data_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
|
||||
started = time.time()
|
||||
try:
|
||||
result = refresh_cell_data(app_config, work_dir, logger)
|
||||
execute_celldata_script(CELLDATA_SCRIPT, app_config, logger)
|
||||
elapsed = round(time.time() - started, 2)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": logger.get_logs(),
|
||||
"status": "completed",
|
||||
"stage": "completed",
|
||||
"elapsed_time": elapsed,
|
||||
"result": {
|
||||
"selected_files": result.selected_files,
|
||||
"parsed_rows": result.parsed_rows,
|
||||
"imported_rows": result.imported_rows,
|
||||
"skipped_rows": result.skipped_rows,
|
||||
},
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.error(str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": logger.get_logs(),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"elapsed_time": round(time.time() - started, 2),
|
||||
}
|
||||
finally:
|
||||
state.reset_task_lock()
|
||||
|
||||
|
||||
def _run_uploaded_cell_data_processing(task_id: str, upload_dir: Path, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
|
||||
started = time.time()
|
||||
try:
|
||||
result = CellDataProcessor(app_config, work_dir, logger).run_local(upload_dir)
|
||||
execute_celldata_script(CELLDATA_SCRIPT, app_config, logger)
|
||||
elapsed = round(time.time() - started, 2)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": logger.get_logs(),
|
||||
"status": "completed",
|
||||
"stage": "completed",
|
||||
"elapsed_time": elapsed,
|
||||
"result": {
|
||||
"selected_files": result.selected_files,
|
||||
"parsed_rows": result.parsed_rows,
|
||||
"imported_rows": result.imported_rows,
|
||||
"skipped_rows": result.skipped_rows,
|
||||
},
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.error(str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": logger.get_logs(),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"elapsed_time": round(time.time() - started, 2),
|
||||
}
|
||||
finally:
|
||||
state.reset_task_lock()
|
||||
@@ -0,0 +1,365 @@
|
||||
import json
|
||||
import re
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from urllib.parse import quote
|
||||
|
||||
import pymysql
|
||||
from fastapi import APIRouter, Body, HTTPException, UploadFile, File
|
||||
from fastapi.responses import Response
|
||||
|
||||
from app import state
|
||||
from app.config import (
|
||||
CellDataConfig,
|
||||
DataMappingsConfig,
|
||||
DEFAULT_CELL_DATA_MAPPING,
|
||||
HistoryRetentionConfig,
|
||||
MetrixConfig,
|
||||
MySQLConfig,
|
||||
RemoteDataConfig,
|
||||
SOURCE_TYPES,
|
||||
WAREHOUSE_TYPES,
|
||||
)
|
||||
from app.services.remote_download import RemoteDataDownloader
|
||||
|
||||
|
||||
router = APIRouter(tags=["config"])
|
||||
|
||||
|
||||
@router.get("/api/config")
|
||||
async def get_config():
|
||||
return state.current_config().to_dict()
|
||||
|
||||
|
||||
@router.get("/api/config/full")
|
||||
async def get_config_full():
|
||||
return state.current_config().to_dict_full()
|
||||
|
||||
|
||||
@router.post("/api/config/mysql")
|
||||
async def update_mysql_config(
|
||||
host: str = Body(...),
|
||||
port: int = Body(...),
|
||||
user: str = Body(...),
|
||||
passwd: str = Body(...),
|
||||
dbname: str = Body(...),
|
||||
):
|
||||
state.reload_config()
|
||||
state.config.mysql.host = host
|
||||
state.config.mysql.port = port
|
||||
state.config.mysql.user = user
|
||||
state.config.mysql.passwd = passwd
|
||||
state.config.mysql.dbname = dbname
|
||||
state.config.save()
|
||||
return {"success": True, "message": "数据库配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/remote")
|
||||
async def update_remote_config(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
state.config.remote_data = RemoteDataConfig.from_dict(config)
|
||||
state.config.save()
|
||||
return {"success": True, "message": "远程数据配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/metrix-enabled")
|
||||
async def update_metrix_enabled(enabled: bool = Body(..., embed=True)):
|
||||
state.reload_config()
|
||||
state.config.metrix_enabled = enabled
|
||||
if not enabled:
|
||||
state.config.source_type = "external"
|
||||
state.config.warehouse_type = "mysql"
|
||||
state.config.save()
|
||||
return {"success": True, "message": "已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/backend")
|
||||
async def update_backend(
|
||||
source_type: str = Body(...),
|
||||
warehouse_type: str = Body(...),
|
||||
):
|
||||
if source_type not in SOURCE_TYPES:
|
||||
raise HTTPException(status_code=400, detail="不支持的源类型")
|
||||
if warehouse_type not in WAREHOUSE_TYPES:
|
||||
raise HTTPException(status_code=400, detail="不支持的仓库类型")
|
||||
state.reload_config()
|
||||
state.config.source_type = source_type
|
||||
state.config.warehouse_type = warehouse_type
|
||||
state.config.save()
|
||||
return {"success": True, "message": "后端类型已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/metrix")
|
||||
async def update_metrix_config(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
state.config.metrix = MetrixConfig.from_dict(config)
|
||||
state.config.save()
|
||||
return {"success": True, "message": "Metrix 连接配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/data-mappings")
|
||||
async def update_data_mappings_config(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
current = state.config.data_mappings.normalized()
|
||||
payload = dict(config)
|
||||
if "table_field_mappings" not in payload:
|
||||
payload["table_field_mappings"] = current.table_field_mappings
|
||||
state.config.data_mappings = DataMappingsConfig.from_dict(payload)
|
||||
state.config.save()
|
||||
return {"success": True, "message": "数据目录映射已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/remote")
|
||||
async def update_cell_data_remote_config(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
current = state.config.cell_data.normalized()
|
||||
state.config.cell_data = CellDataConfig(
|
||||
remote_data=RemoteDataConfig.from_dict(config),
|
||||
mysql=current.mysql,
|
||||
scan_paths=current.scan_paths,
|
||||
year_dir_regex=current.year_dir_regex,
|
||||
month_dir_regex=current.month_dir_regex,
|
||||
day_dir_regex=current.day_dir_regex,
|
||||
file_name_regex=current.file_name_regex,
|
||||
file_time_regex=current.file_time_regex,
|
||||
mapping=current.mapping,
|
||||
).normalized()
|
||||
state.config.save()
|
||||
return {"success": True, "message": "CellData 远程数据源配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/mysql")
|
||||
async def update_cell_data_mysql_config(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
current = state.config.cell_data.normalized()
|
||||
state.config.cell_data = CellDataConfig(
|
||||
remote_data=current.remote_data,
|
||||
mysql=MySQLConfig.from_dict(config, default_dbname="celldata"),
|
||||
scan_paths=current.scan_paths,
|
||||
year_dir_regex=current.year_dir_regex,
|
||||
month_dir_regex=current.month_dir_regex,
|
||||
day_dir_regex=current.day_dir_regex,
|
||||
file_name_regex=current.file_name_regex,
|
||||
file_time_regex=current.file_time_regex,
|
||||
mapping=current.mapping,
|
||||
).normalized()
|
||||
state.config.save()
|
||||
return {"success": True, "message": "CellData 数据库配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/remote/test")
|
||||
def test_cell_data_remote_connection(config: dict[str, Any] | None = Body(None)):
|
||||
try:
|
||||
remote_config = RemoteDataConfig.from_dict(config) if config else state.current_config().cell_data.remote_data
|
||||
RemoteDataDownloader(remote_config).test_connection()
|
||||
return {"success": True, "message": "CellData 远程服务器连接成功"}
|
||||
except Exception as exc:
|
||||
return {"success": False, "message": f"连接失败: {exc}"}
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/mysql/test")
|
||||
def test_cell_data_mysql_connection(config: dict[str, Any] | None = Body(None)):
|
||||
mysql_config = MySQLConfig.from_dict(config, default_dbname="celldata") if config else state.current_config().cell_data.mysql
|
||||
return _test_mysql_config(mysql_config)
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/settings")
|
||||
async def update_cell_data_settings(config: dict[str, Any] = Body(...)):
|
||||
validation = _validate_cell_data_settings(config)
|
||||
if not validation["success"]:
|
||||
raise HTTPException(status_code=400, detail=validation["message"])
|
||||
state.reload_config()
|
||||
current = state.config.cell_data.normalized()
|
||||
state.config.cell_data = CellDataConfig(
|
||||
remote_data=current.remote_data,
|
||||
mysql=current.mysql,
|
||||
scan_paths=config.get("scan_paths", current.scan_paths),
|
||||
year_dir_regex=str(config.get("year_dir_regex", current.year_dir_regex)),
|
||||
month_dir_regex=str(config.get("month_dir_regex", current.month_dir_regex)),
|
||||
day_dir_regex=str(config.get("day_dir_regex", current.day_dir_regex)),
|
||||
file_name_regex=str(config.get("file_name_regex", current.file_name_regex)),
|
||||
file_time_regex=str(config.get("file_time_regex", current.file_time_regex)),
|
||||
mapping=config.get("mapping", current.mapping),
|
||||
).normalized()
|
||||
state.config.save()
|
||||
return {"success": True, "message": "CellData 规则已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/cell-data/settings/validate")
|
||||
async def validate_cell_data_settings(config: dict[str, Any] = Body(...)):
|
||||
return _validate_cell_data_settings(config)
|
||||
|
||||
|
||||
@router.get("/api/config/cell-data/mapping/default")
|
||||
async def get_default_cell_data_mapping():
|
||||
return DEFAULT_CELL_DATA_MAPPING
|
||||
|
||||
|
||||
@router.post("/api/config/history-retention")
|
||||
async def update_history_retention(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
state.config.history_retention = HistoryRetentionConfig.from_dict(config)
|
||||
state.config.save()
|
||||
return {"success": True, "message": "处理历史保留配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/sheet-filter")
|
||||
async def update_sheet_filter(filters: list[str] = Body(...)):
|
||||
state.reload_config()
|
||||
state.config.sheet_filter = filters
|
||||
state.config.save()
|
||||
return {"success": True, "message": "Sheet 过滤规则已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.post("/api/config/extract-fields")
|
||||
async def update_extract_fields(fields: list[dict[str, Any]] = Body(...)):
|
||||
state.reload_config()
|
||||
state.config.extract_fields = fields
|
||||
state.config.save()
|
||||
return {"success": True, "message": "字段映射配置已更新", "update": state.config.update}
|
||||
|
||||
|
||||
@router.get("/api/config/download")
|
||||
async def download_config():
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"Configure_{timestamp}.json"
|
||||
config_data = state.current_config().to_file_dict()
|
||||
content = json.dumps(config_data, ensure_ascii=False, indent=2)
|
||||
return Response(
|
||||
content=content,
|
||||
media_type="application/json",
|
||||
headers={"Content-Disposition": f"attachment; filename*=UTF-8''{quote(filename)}"},
|
||||
)
|
||||
|
||||
|
||||
@router.post("/api/config/upload")
|
||||
async def upload_config(file: UploadFile = File(...)):
|
||||
if not file.filename or not file.filename.endswith(".json"):
|
||||
raise HTTPException(status_code=400, detail="只支持 JSON 格式的配置文件")
|
||||
|
||||
try:
|
||||
data = json.loads((await file.read()).decode("utf-8"))
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("配置文件格式错误:必须是 JSON 对象")
|
||||
|
||||
state.reload_config()
|
||||
_apply_config_data(data)
|
||||
state.config.save()
|
||||
return {"success": True, "message": "配置文件上传成功", "update": state.config.update}
|
||||
except json.JSONDecodeError as exc:
|
||||
raise HTTPException(status_code=400, detail="配置文件格式错误:不是有效的 JSON 文件") from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=f"上传失败: {exc}") from exc
|
||||
|
||||
|
||||
def _apply_config_data(data: dict[str, Any]) -> None:
|
||||
if "MetrixEnabled" in data:
|
||||
state.config.metrix_enabled = bool(data["MetrixEnabled"])
|
||||
|
||||
if data.get("SourceType") in SOURCE_TYPES:
|
||||
state.config.source_type = data["SourceType"]
|
||||
if data.get("WarehouseType") in WAREHOUSE_TYPES:
|
||||
state.config.warehouse_type = data["WarehouseType"]
|
||||
|
||||
# 未启用 Metrix 时强制回落为外部储存 + 直连 MySQL,避免导入后状态不一致
|
||||
if not state.config.metrix_enabled:
|
||||
state.config.source_type = "external"
|
||||
state.config.warehouse_type = "mysql"
|
||||
|
||||
metrix_data = data.get("Metrix")
|
||||
if isinstance(metrix_data, dict):
|
||||
state.config.metrix = MetrixConfig.from_dict(metrix_data)
|
||||
|
||||
data_mappings = data.get("DataMappings")
|
||||
if isinstance(data_mappings, dict):
|
||||
state.config.data_mappings = DataMappingsConfig.from_dict(data_mappings)
|
||||
|
||||
mysql_data = data.get("MySQL_DBInfo")
|
||||
if isinstance(mysql_data, dict):
|
||||
state.config.mysql = MySQLConfig.from_dict(mysql_data)
|
||||
|
||||
if "SheetFilter" in data:
|
||||
state.config.sheet_filter = data["SheetFilter"] if isinstance(data["SheetFilter"], list) else []
|
||||
|
||||
if "ExtractField" in data:
|
||||
state.config.extract_fields = data["ExtractField"] if isinstance(data["ExtractField"], list) else []
|
||||
|
||||
remote_data = data.get("RemoteData")
|
||||
if isinstance(remote_data, dict):
|
||||
state.config.remote_data = RemoteDataConfig.from_dict(remote_data)
|
||||
|
||||
cell_data = data.get("CellData")
|
||||
if isinstance(cell_data, dict):
|
||||
state.config.cell_data = CellDataConfig.from_dict(cell_data)
|
||||
|
||||
history_retention = data.get("HistoryRetention")
|
||||
if isinstance(history_retention, dict):
|
||||
state.config.history_retention = HistoryRetentionConfig.from_dict(history_retention)
|
||||
|
||||
|
||||
def _test_mysql_config(config: MySQLConfig) -> dict[str, Any]:
|
||||
normalized = config.normalized()
|
||||
try:
|
||||
conn = pymysql.connect(
|
||||
host=normalized.host,
|
||||
port=normalized.port,
|
||||
user=normalized.user,
|
||||
password=normalized.passwd,
|
||||
database=normalized.dbname,
|
||||
charset="utf8mb4",
|
||||
cursorclass=pymysql.cursors.DictCursor,
|
||||
connect_timeout=10,
|
||||
)
|
||||
try:
|
||||
with conn.cursor() as cursor:
|
||||
cursor.execute("SELECT 1")
|
||||
finally:
|
||||
conn.close()
|
||||
return {"success": True, "message": "连接成功"}
|
||||
except Exception as exc:
|
||||
return {"success": False, "message": str(exc)}
|
||||
|
||||
|
||||
def _validate_cell_data_settings(config: dict[str, Any]) -> dict[str, Any]:
|
||||
scan_paths = config.get("scan_paths", [])
|
||||
if not isinstance(scan_paths, list) or not any(str(path).strip() for path in scan_paths):
|
||||
return {"success": False, "message": "请至少配置一个扫描路径"}
|
||||
|
||||
for key in ("year_dir_regex", "month_dir_regex", "day_dir_regex", "file_name_regex", "file_time_regex"):
|
||||
try:
|
||||
re.compile(str(config.get(key, "")))
|
||||
except re.error as exc:
|
||||
return {"success": False, "message": f"{key} 正则无效: {exc}"}
|
||||
|
||||
mapping = config.get("mapping")
|
||||
if not isinstance(mapping, dict):
|
||||
return {"success": False, "message": "映射规则必须是 JSON 对象"}
|
||||
if not str(mapping.get("target_table", "")).strip():
|
||||
return {"success": False, "message": "映射规则缺少 target_table"}
|
||||
key_config = mapping.get("key")
|
||||
if not isinstance(key_config, dict) or not key_config.get("field") or not key_config.get("expr"):
|
||||
return {"success": False, "message": "映射规则缺少 key.field 或 key.expr"}
|
||||
sources = mapping.get("sources")
|
||||
if not isinstance(sources, list) or not sources:
|
||||
return {"success": False, "message": "映射规则至少需要一个 sources 项"}
|
||||
|
||||
for index, source in enumerate(sources, 1):
|
||||
if not isinstance(source, dict):
|
||||
return {"success": False, "message": f"sources 第 {index} 项必须是对象"}
|
||||
if not source.get("band") or not source.get("file_prefix"):
|
||||
return {"success": False, "message": f"sources 第 {index} 项缺少 band 或 file_prefix"}
|
||||
fields = source.get("fields")
|
||||
if not isinstance(fields, dict) or not fields:
|
||||
return {"success": False, "message": f"sources 第 {index} 项缺少 fields"}
|
||||
for target, rule in fields.items():
|
||||
if not str(target).strip():
|
||||
return {"success": False, "message": f"sources 第 {index} 项存在空目标字段"}
|
||||
if isinstance(rule, str) and rule.strip():
|
||||
continue
|
||||
if isinstance(rule, dict) and "value" in rule:
|
||||
continue
|
||||
return {"success": False, "message": f"{target} 的映射规则无效"}
|
||||
return {"success": True, "message": "映射规则有效"}
|
||||
|
||||
@@ -0,0 +1,339 @@
|
||||
"""容量看板分析接口。
|
||||
|
||||
基于 4G/5G 结果表(含富集 + 高负荷判定 + 优化建议列)做聚合分析,供前端容量看板展示。
|
||||
读经 `make_warehouse`(直连 MySQL 或 Metrix 平台),分析对象为“主仓库”里的结果表。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.background import BackgroundTask
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR
|
||||
from app.utils.files import remove_file_safely
|
||||
from app.warehouse import make_warehouse
|
||||
|
||||
router = APIRouter(tags=["dashboard"])
|
||||
|
||||
# 每个制式(结果表)的关键列映射
|
||||
RAT = {
|
||||
"4g": {
|
||||
"table": "4G_结果表",
|
||||
"id": "CGI",
|
||||
"name": "小区名称",
|
||||
"ul": "上行PUSCH利用率",
|
||||
"dl": "下行PDSCH利用率",
|
||||
"ul_label": "上行PUSCH利用率",
|
||||
"dl_label": "下行PDSCH利用率",
|
||||
"users": "YY-RRC连接建立最大用户数",
|
||||
},
|
||||
"5g": {
|
||||
"table": "5G_结果表",
|
||||
"id": "NCGI",
|
||||
"name": "CU小区配置名称",
|
||||
"ul": "上行PRB平均利用率",
|
||||
"dl": "下行PRB平均利用率",
|
||||
"ul_label": "上行PRB利用率",
|
||||
"dl_label": "下行PRB利用率",
|
||||
"users": "RRC连接平均连接用户数",
|
||||
},
|
||||
}
|
||||
FLOW = "日均流量(GB)"
|
||||
PROBLEMS = ["高负荷", "利用率预警", "高流量预警"]
|
||||
|
||||
|
||||
def _db():
|
||||
return make_warehouse(state.current_config())
|
||||
|
||||
|
||||
def _esc(value: str) -> str:
|
||||
return str(value).replace("\\", "\\\\").replace("'", "''")
|
||||
|
||||
|
||||
def _rows(db, sql: str) -> list[dict[str, Any]]:
|
||||
ok, result = db.execute_sql(sql)
|
||||
if not ok:
|
||||
raise HTTPException(status_code=500, detail=str(result))
|
||||
return result if isinstance(result, list) else []
|
||||
|
||||
|
||||
def _num(value: Any, digits: int | None = None) -> float | int:
|
||||
try:
|
||||
num = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
if digits is None:
|
||||
return int(num)
|
||||
return round(num, digits)
|
||||
|
||||
|
||||
def _rat(rat: str) -> dict[str, str]:
|
||||
cfg = RAT.get((rat or "").lower())
|
||||
if not cfg:
|
||||
raise HTTPException(status_code=400, detail="rat 仅支持 4g / 5g")
|
||||
return cfg
|
||||
|
||||
|
||||
def _existing_tables(db) -> set[str]:
|
||||
try:
|
||||
return {str(t).lower() for t in db.get_tables()}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise HTTPException(status_code=500, detail=f"读取数据表失败: {exc}") from exc
|
||||
|
||||
|
||||
@router.get("/api/dashboard/status")
|
||||
def dashboard_status():
|
||||
tables = _existing_tables(_db())
|
||||
has_4g = RAT["4g"]["table"].lower() in tables
|
||||
has_5g = RAT["5g"]["table"].lower() in tables
|
||||
return {"has_4g": has_4g, "has_5g": has_5g, "ready": has_4g and has_5g}
|
||||
|
||||
|
||||
@router.get("/api/dashboard/overview")
|
||||
def dashboard_overview(rat: str = Query("4g")):
|
||||
cfg = _rat(rat)
|
||||
db = _db()
|
||||
if cfg["table"].lower() not in _existing_tables(db):
|
||||
raise HTTPException(status_code=404, detail=f"结果表 {cfg['table']} 不存在,请先进行数据处理")
|
||||
|
||||
t = f"`{cfg['table']}`"
|
||||
ul, dl, flow, users = f"`{cfg['ul']}`", f"`{cfg['dl']}`", f"`{FLOW}`", f"`{cfg['users']}`"
|
||||
|
||||
summary = _rows(db, (
|
||||
f"SELECT COUNT(*) total,"
|
||||
f" SUM(`高负荷问题`='高负荷') high_load,"
|
||||
f" SUM(`高负荷问题`='利用率预警') util_warn,"
|
||||
f" SUM(`高负荷问题`='高流量预警') flow_warn,"
|
||||
f" SUM(`高负荷问题` IS NULL) normal,"
|
||||
f" ROUND(AVG({ul})*100,1) avg_ul,"
|
||||
f" ROUND(AVG({dl})*100,1) avg_dl,"
|
||||
f" ROUND(MAX({dl})*100,1) max_dl,"
|
||||
f" ROUND(AVG({flow}),2) avg_flow,"
|
||||
f" ROUND(SUM({flow}),1) total_flow"
|
||||
f" FROM {t}"
|
||||
))
|
||||
s = summary[0] if summary else {}
|
||||
summary_out = {
|
||||
"total": _num(s.get("total")),
|
||||
"high_load": _num(s.get("high_load")),
|
||||
"util_warn": _num(s.get("util_warn")),
|
||||
"flow_warn": _num(s.get("flow_warn")),
|
||||
"normal": _num(s.get("normal")),
|
||||
"avg_ul": _num(s.get("avg_ul"), 1),
|
||||
"avg_dl": _num(s.get("avg_dl"), 1),
|
||||
"max_dl": _num(s.get("max_dl"), 1),
|
||||
"avg_flow": _num(s.get("avg_flow"), 2),
|
||||
"total_flow": _num(s.get("total_flow"), 1),
|
||||
}
|
||||
|
||||
problem_pie = [
|
||||
{"name": "高负荷", "value": summary_out["high_load"]},
|
||||
{"name": "利用率预警", "value": summary_out["util_warn"]},
|
||||
{"name": "高流量预警", "value": summary_out["flow_warn"]},
|
||||
{"name": "正常", "value": summary_out["normal"]},
|
||||
]
|
||||
|
||||
def group_by(col: str, limit: int = 0, skip_unknown: bool = False) -> list[dict[str, Any]]:
|
||||
limit_sql = f" LIMIT {int(limit)}" if limit else ""
|
||||
where_sql = f" WHERE `{col}` IS NOT NULL AND `{col}` <> ''" if skip_unknown else ""
|
||||
rows = _rows(db, (
|
||||
f"SELECT IFNULL(`{col}`,'未知') name, COUNT(*) total,"
|
||||
f" SUM(`是否高负荷小区`='是') high,"
|
||||
f" SUM(`高负荷问题` IS NOT NULL) flagged"
|
||||
f" FROM {t}{where_sql} GROUP BY `{col}` ORDER BY flagged DESC, total DESC{limit_sql}"
|
||||
))
|
||||
return [
|
||||
{"name": str(r.get("name")), "total": _num(r.get("total")),
|
||||
"high": _num(r.get("high")), "flagged": _num(r.get("flagged"))}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
def util_hist(col: str) -> list[dict[str, Any]]:
|
||||
rows = _rows(db, (
|
||||
f"SELECT LEAST(FLOOR({col}*10),9) b, COUNT(*) c FROM {t}"
|
||||
f" WHERE {col} IS NOT NULL GROUP BY b ORDER BY b"
|
||||
))
|
||||
bucket_map = {int(_num(r.get("b"))): _num(r.get("c")) for r in rows}
|
||||
return [{"bucket": f"{i*10}-{i*10+10}%", "value": bucket_map.get(i, 0)} for i in range(10)]
|
||||
|
||||
top_rows = _rows(db, (
|
||||
f"SELECT `{cfg['id']}` id, IFNULL(`{cfg['name']}`,'') name, IFNULL(`制式`,'') `system`,"
|
||||
f" IFNULL(`带宽`,'') band, IFNULL(`站型`,'') station,"
|
||||
f" ROUND({ul}*100,1) ul, ROUND({dl}*100,1) dl, ROUND({flow},2) flow, ROUND({users},0) users,"
|
||||
f" IFNULL(`高负荷问题`,'') problem"
|
||||
f" FROM {t} WHERE `高负荷问题`='高负荷' ORDER BY {dl} DESC, {flow} DESC LIMIT 10"
|
||||
))
|
||||
top_cells = [
|
||||
{"id": str(r.get("id")), "name": str(r.get("name")), "system": str(r.get("system")),
|
||||
"band": str(r.get("band")), "station": str(r.get("station")),
|
||||
"ul": _num(r.get("ul"), 1), "dl": _num(r.get("dl"), 1),
|
||||
"flow": _num(r.get("flow"), 2), "users": _num(r.get("users")), "problem": str(r.get("problem"))}
|
||||
for r in top_rows
|
||||
]
|
||||
|
||||
return {
|
||||
"rat": rat.lower(),
|
||||
"labels": {"ul": cfg["ul_label"], "dl": cfg["dl_label"]},
|
||||
"summary": summary_out,
|
||||
"problem_pie": problem_pie,
|
||||
"by_system": group_by("制式", skip_unknown=True),
|
||||
"by_station": group_by("站型", skip_unknown=True),
|
||||
"by_freq": group_by("频段", limit=12, skip_unknown=True),
|
||||
"ul_hist": util_hist(ul),
|
||||
"dl_hist": util_hist(dl),
|
||||
"top_cells": top_cells,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/api/dashboard/cells")
|
||||
def dashboard_cells(
|
||||
rat: str = Query("4g"),
|
||||
problem: str = Query(""),
|
||||
keyword: str = Query(""),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=200),
|
||||
):
|
||||
cfg = _rat(rat)
|
||||
db = _db()
|
||||
if cfg["table"].lower() not in _existing_tables(db):
|
||||
raise HTTPException(status_code=404, detail=f"结果表 {cfg['table']} 不存在,请先进行数据处理")
|
||||
|
||||
t = f"`{cfg['table']}`"
|
||||
ul, dl, flow, users = f"`{cfg['ul']}`", f"`{cfg['dl']}`", f"`{FLOW}`", f"`{cfg['users']}`"
|
||||
|
||||
wheres = ["`高负荷问题` IS NOT NULL"]
|
||||
if problem in PROBLEMS:
|
||||
wheres = [f"`高负荷问题`='{problem}'"]
|
||||
if keyword.strip():
|
||||
kw = _esc(keyword.strip())
|
||||
wheres.append(f"(`{cfg['id']}` LIKE '%{kw}%' OR `{cfg['name']}` LIKE '%{kw}%')")
|
||||
where_sql = " WHERE " + " AND ".join(wheres)
|
||||
|
||||
total = _num(_rows(db, f"SELECT COUNT(*) c FROM {t}{where_sql}")[0].get("c"))
|
||||
offset = (page - 1) * page_size
|
||||
rows = _rows(db, (
|
||||
f"SELECT `{cfg['id']}` id, IFNULL(`{cfg['name']}`,'') name, IFNULL(`制式`,'') `system`,"
|
||||
f" IFNULL(`带宽`,'') band, IFNULL(`站型`,'') station, IFNULL(`频段`,'') freq,"
|
||||
f" ROUND({ul}*100,1) ul, ROUND({dl}*100,1) dl, ROUND({flow},2) flow, ROUND({users},0) users,"
|
||||
f" IFNULL(`高负荷问题`,'') problem, IFNULL(`是否高负荷小区`,'否') is_high"
|
||||
f" FROM {t}{where_sql}"
|
||||
f" ORDER BY FIELD(`高负荷问题`,'高负荷','高流量预警','利用率预警'), {dl} DESC"
|
||||
f" LIMIT {int(page_size)} OFFSET {int(offset)}"
|
||||
))
|
||||
items = [
|
||||
{"id": str(r.get("id")), "name": str(r.get("name")), "system": str(r.get("system")),
|
||||
"band": str(r.get("band")), "station": str(r.get("station")), "freq": str(r.get("freq")),
|
||||
"ul": _num(r.get("ul"), 1), "dl": _num(r.get("dl"), 1), "flow": _num(r.get("flow"), 2),
|
||||
"users": _num(r.get("users")), "problem": str(r.get("problem")), "is_high": str(r.get("is_high"))}
|
||||
for r in rows
|
||||
]
|
||||
return {"items": items, "total": total, "page": page, "page_size": page_size}
|
||||
|
||||
|
||||
@router.get("/api/dashboard/export")
|
||||
def dashboard_export(rat: str = Query("4g"), problem: str = Query(""), keyword: str = Query("")):
|
||||
cfg = _rat(rat)
|
||||
db = _db()
|
||||
if cfg["table"].lower() not in _existing_tables(db):
|
||||
raise HTTPException(status_code=404, detail=f"结果表 {cfg['table']} 不存在,请先进行数据处理")
|
||||
|
||||
t = f"`{cfg['table']}`"
|
||||
ul, dl, flow, users = f"`{cfg['ul']}`", f"`{cfg['dl']}`", f"`{FLOW}`", f"`{cfg['users']}`"
|
||||
wheres = ["`高负荷问题` IS NOT NULL"]
|
||||
if problem in PROBLEMS:
|
||||
wheres = [f"`高负荷问题`='{problem}'"]
|
||||
if keyword.strip():
|
||||
kw = _esc(keyword.strip())
|
||||
wheres.append(f"(`{cfg['id']}` LIKE '%{kw}%' OR `{cfg['name']}` LIKE '%{kw}%')")
|
||||
where_sql = " WHERE " + " AND ".join(wheres)
|
||||
|
||||
rows = _rows(db, (
|
||||
f"SELECT `{cfg['id']}` id, IFNULL(`{cfg['name']}`,'') name, IFNULL(`制式`,'') sys,"
|
||||
f" IFNULL(`带宽`,'') band, IFNULL(`站型`,'') station, IFNULL(`频段`,'') freq,"
|
||||
f" ROUND({ul}*100,1) ul, ROUND({dl}*100,1) dl, ROUND({flow},2) flow, ROUND({users},0) users,"
|
||||
f" IFNULL(`高负荷问题`,'') problem, IFNULL(`优化建议`,'') suggestion"
|
||||
f" FROM {t}{where_sql}"
|
||||
f" ORDER BY FIELD(`高负荷问题`,'高负荷','高流量预警','利用率预警'), {dl} DESC"
|
||||
))
|
||||
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"问题小区清单_{rat.lower()}_{timestamp}.csv"
|
||||
filepath = CACHE_DIR / filename
|
||||
header = [cfg["id"], "小区名称", "制式", "带宽", "站型", "频段",
|
||||
"上行利用率(%)", "下行利用率(%)", "日均流量(GB)", "用户数", "高负荷问题", "优化建议"]
|
||||
try:
|
||||
with filepath.open("w", encoding="utf-8-sig", newline="") as handle:
|
||||
writer = csv.writer(handle)
|
||||
writer.writerow(header)
|
||||
for r in rows:
|
||||
writer.writerow([
|
||||
r.get("id"), r.get("name"), r.get("sys"), r.get("band"), r.get("station"),
|
||||
r.get("freq"), r.get("ul"), r.get("dl"), r.get("flow"), r.get("users"),
|
||||
r.get("problem"), r.get("suggestion"),
|
||||
])
|
||||
except Exception:
|
||||
remove_file_safely(filepath)
|
||||
raise
|
||||
return FileResponse(
|
||||
path=str(filepath), filename=filename, media_type="text/csv",
|
||||
background=BackgroundTask(remove_file_safely, filepath),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/api/dashboard/cell")
|
||||
def dashboard_cell(rat: str = Query("4g"), id: str = Query(...)):
|
||||
cfg = _rat(rat)
|
||||
db = _db()
|
||||
if cfg["table"].lower() not in _existing_tables(db):
|
||||
raise HTTPException(status_code=404, detail=f"结果表 {cfg['table']} 不存在")
|
||||
|
||||
t = f"`{cfg['table']}`"
|
||||
cell_id = _esc(id)
|
||||
rows = _rows(db, f"SELECT * FROM {t} WHERE `{cfg['id']}`='{cell_id}' LIMIT 1")
|
||||
if not rows:
|
||||
raise HTTPException(status_code=404, detail="未找到该小区")
|
||||
row = rows[0]
|
||||
|
||||
# 同扇区同 PLMN 的兄弟小区(用于详情页的均衡上下文)
|
||||
sector = str(row.get("扇区") or "")
|
||||
siblings: list[dict[str, Any]] = []
|
||||
if sector:
|
||||
plmn = _esc("-".join(str(id).split("-")[:2]))
|
||||
sector_e = _esc(sector)
|
||||
dl = f"`{cfg['dl']}`"
|
||||
sib_rows = _rows(db, (
|
||||
f"SELECT `{cfg['id']}` id, IFNULL(`{cfg['name']}`,'') name, IFNULL(`带宽`,'') band,"
|
||||
f" IFNULL(`频段`,'') freq, ROUND(`{cfg['ul']}`*100,1) ul, ROUND({dl}*100,1) dl,"
|
||||
f" ROUND(`{FLOW}`,2) flow, IFNULL(`高负荷问题`,'正常') problem"
|
||||
f" FROM {t} WHERE `扇区`='{sector_e}' AND SUBSTRING_INDEX(`{cfg['id']}`,'-',2)='{plmn}'"
|
||||
f" AND `{cfg['id']}`<>'{cell_id}' ORDER BY {dl} DESC LIMIT 30"
|
||||
))
|
||||
siblings = [
|
||||
{"id": str(r.get("id")), "name": str(r.get("name")), "band": str(r.get("band")),
|
||||
"freq": str(r.get("freq")), "ul": _num(r.get("ul"), 1), "dl": _num(r.get("dl"), 1),
|
||||
"flow": _num(r.get("flow"), 2), "problem": str(r.get("problem"))}
|
||||
for r in sib_rows
|
||||
]
|
||||
|
||||
# 原始行转为字符串友好的 dict(数值保留,None→空)
|
||||
detail = {str(k): (v if v is not None else "") for k, v in row.items()}
|
||||
return {
|
||||
"rat": rat.lower(),
|
||||
"id": str(row.get(cfg["id"]) or id),
|
||||
"name": str(row.get(cfg["name"]) or ""),
|
||||
"labels": {"ul": cfg["ul_label"], "dl": cfg["dl_label"]},
|
||||
"id_field": cfg["id"],
|
||||
"name_field": cfg["name"],
|
||||
"ul_field": cfg["ul"],
|
||||
"dl_field": cfg["dl"],
|
||||
"users_field": cfg["users"],
|
||||
"flow_field": FLOW,
|
||||
"detail": detail,
|
||||
"siblings": siblings,
|
||||
}
|
||||
@@ -0,0 +1,430 @@
|
||||
import csv
|
||||
import re
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import pandas as pd
|
||||
from fastapi import APIRouter, Body, File, Form, HTTPException, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.background import BackgroundTask
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR
|
||||
from app.database import DatabaseManager, detect_csv_encoding
|
||||
from app.services.platform import make_client
|
||||
from app.utils.files import remove_file_safely
|
||||
from app.warehouse import make_cell_data_warehouse, make_warehouse
|
||||
|
||||
|
||||
router = APIRouter(tags=["database"])
|
||||
INVALID_SHEET_NAME_CHARS = re.compile(r"[:\\/?*\[\]]")
|
||||
|
||||
|
||||
def _resolve_requested_tables(
|
||||
table_name: Optional[str],
|
||||
table_names: Optional[list[str]],
|
||||
) -> list[str]:
|
||||
names = table_names if table_names is not None else ([table_name] if table_name else [])
|
||||
return [name.strip() for name in names if isinstance(name, str) and name.strip()]
|
||||
|
||||
|
||||
def _make_sheet_name(table_name: str, used_names: set[str]) -> str:
|
||||
base = INVALID_SHEET_NAME_CHARS.sub("_", table_name).strip("'").strip() or "Sheet"
|
||||
base = base[:31]
|
||||
sheet_name = base
|
||||
index = 2
|
||||
|
||||
while sheet_name in used_names:
|
||||
suffix = f"_{index}"
|
||||
sheet_name = f"{base[:31 - len(suffix)]}{suffix}" or f"Sheet_{index}"
|
||||
index += 1
|
||||
|
||||
used_names.add(sheet_name)
|
||||
return sheet_name
|
||||
|
||||
|
||||
DatabaseSource = str
|
||||
|
||||
|
||||
def _dataframe_from_table(db: DatabaseManager, table_name: str) -> pd.DataFrame:
|
||||
result = db.query_table(table_name, page=1, page_size=1000000)
|
||||
table_info = db.get_table_info(table_name)
|
||||
columns = [str(column["Field"]) for column in table_info["columns"]]
|
||||
return pd.DataFrame(result["data"], columns=columns)
|
||||
|
||||
|
||||
def _db(database_source: DatabaseSource = "main"):
|
||||
"""Direct MySQL DatabaseManager, or a Metrix-backed warehouse with the same interface."""
|
||||
config = state.current_config()
|
||||
if database_source == "cell_data":
|
||||
return make_cell_data_warehouse(config)
|
||||
if database_source != "main":
|
||||
raise HTTPException(status_code=400, detail="不支持的数据库来源")
|
||||
return make_warehouse(config)
|
||||
|
||||
|
||||
@router.post("/api/database/test")
|
||||
def test_database(database_source: DatabaseSource = Body("main", embed=True)):
|
||||
db = _db(database_source)
|
||||
success, message = db.test_connection()
|
||||
return {"success": success, "message": message}
|
||||
|
||||
|
||||
@router.get("/api/database/info")
|
||||
def get_database_info(database_source: DatabaseSource = "main"):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
return {"success": True, **db.get_server_info()}
|
||||
except Exception as exc:
|
||||
return {"success": False, "error": str(exc)}
|
||||
|
||||
|
||||
@router.get("/api/database/tables")
|
||||
@router.post("/api/database/tables")
|
||||
def get_tables(database_source: DatabaseSource = Body("main", embed=True)):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
return {"tables": db.get_tables()}
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/info")
|
||||
def get_table_info(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
return db.get_table_info(table_name)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/data")
|
||||
def query_table_data(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
page: int = Body(1),
|
||||
page_size: int = Body(50),
|
||||
order_by: Optional[str] = Body(None),
|
||||
order_dir: str = Body("ASC"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
return db.query_table(table_name, page, page_size, order_by=order_by, order_dir=order_dir)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/query")
|
||||
def query_table_with_filter(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
page: int = Body(1),
|
||||
page_size: int = Body(50),
|
||||
filters: Optional[dict[str, str]] = Body(None),
|
||||
order_by: Optional[str] = Body(None),
|
||||
order_dir: str = Body("ASC"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
return db.query_table(
|
||||
table_name,
|
||||
page,
|
||||
page_size,
|
||||
filters=filters or {},
|
||||
order_by=order_by,
|
||||
order_dir=order_dir,
|
||||
)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/truncate")
|
||||
def truncate_table(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
db.truncate_table(table_name)
|
||||
return {"success": True, "message": f"表 {table_name} 已清空"}
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/drop")
|
||||
def drop_table(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
db.drop_table(table_name)
|
||||
return {"success": True, "message": f"表 {table_name} 已删除"}
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/drop-all")
|
||||
def drop_all_tables(database_source: DatabaseSource = Body("main", embed=True)):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
result = db.drop_all_tables()
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"已删除 {result['dropped_count']} 个表",
|
||||
"dropped_count": result["dropped_count"],
|
||||
"tables": result["tables"],
|
||||
}
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/database/table/row/update")
|
||||
def update_table_row(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
identifier: dict = Body(...),
|
||||
values: dict = Body(...),
|
||||
):
|
||||
if not identifier:
|
||||
raise HTTPException(status_code=400, detail="缺少行定位信息")
|
||||
db = _db(database_source)
|
||||
try:
|
||||
affected = db.update_row(table_name, identifier, values)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
if affected == 0:
|
||||
raise HTTPException(status_code=404, detail="未找到匹配的数据行,可能已被修改或删除")
|
||||
return {"success": True, "message": "已更新该行", "affected_rows": affected}
|
||||
|
||||
|
||||
@router.post("/api/database/table/row/delete")
|
||||
def delete_table_row(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
identifier: dict = Body(...),
|
||||
):
|
||||
if not identifier:
|
||||
raise HTTPException(status_code=400, detail="缺少行定位信息")
|
||||
db = _db(database_source)
|
||||
try:
|
||||
affected = db.delete_row(table_name, identifier)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
if affected == 0:
|
||||
raise HTTPException(status_code=404, detail="未找到匹配的数据行,可能已被删除")
|
||||
return {"success": True, "message": "已删除该行", "affected_rows": affected}
|
||||
|
||||
|
||||
def _table_columns(db, table_name: str) -> list[str]:
|
||||
info = db.get_table_info(table_name)
|
||||
return [str(column["Field"]) for column in info.get("columns", [])]
|
||||
|
||||
|
||||
@router.post("/api/database/table/template")
|
||||
def download_table_template(
|
||||
table_name: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
columns = _table_columns(db, table_name)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
if not columns:
|
||||
raise HTTPException(status_code=400, detail="无法获取表字段,无法生成模板")
|
||||
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filepath = CACHE_DIR / f"{table_name}_template_{timestamp}.csv"
|
||||
try:
|
||||
with filepath.open("w", encoding="utf-8-sig", newline="") as handle:
|
||||
csv.writer(handle).writerow(columns)
|
||||
except Exception:
|
||||
remove_file_safely(filepath)
|
||||
raise
|
||||
return FileResponse(
|
||||
path=str(filepath),
|
||||
filename=f"{table_name}_模板.csv",
|
||||
media_type="text/csv",
|
||||
background=BackgroundTask(remove_file_safely, filepath),
|
||||
)
|
||||
|
||||
|
||||
def _read_csv_header(path: Path, encoding: str) -> list[str]:
|
||||
with path.open("r", encoding=encoding, newline="") as handle:
|
||||
for row in csv.reader(handle):
|
||||
return [str(cell).strip() for cell in row]
|
||||
return []
|
||||
|
||||
|
||||
# Sync def so FastAPI runs it in a threadpool: file read + DB insert are blocking.
|
||||
@router.post("/api/database/table/import")
|
||||
def import_table_csv(
|
||||
file: UploadFile = File(...),
|
||||
table_name: str = Form(...),
|
||||
database_source: str = Form("main"),
|
||||
):
|
||||
if not file.filename or not file.filename.lower().endswith(".csv"):
|
||||
raise HTTPException(status_code=400, detail="仅支持 CSV 格式文件")
|
||||
db = _db(database_source)
|
||||
try:
|
||||
columns = _table_columns(db, table_name)
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
if not columns:
|
||||
raise HTTPException(status_code=400, detail="无法获取表字段")
|
||||
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
tmp_path = CACHE_DIR / f"import_{timestamp}.csv"
|
||||
tmp_path.write_bytes(file.file.read())
|
||||
try:
|
||||
encoding = detect_csv_encoding(str(tmp_path))
|
||||
header = _read_csv_header(tmp_path, encoding)
|
||||
if not header:
|
||||
raise HTTPException(status_code=400, detail="CSV 文件为空或缺少表头")
|
||||
missing = [name for name in columns if name not in header]
|
||||
extra = [name for name in header if name not in columns]
|
||||
if missing or extra:
|
||||
parts = []
|
||||
if missing:
|
||||
parts.append("缺少字段: " + ", ".join(missing))
|
||||
if extra:
|
||||
parts.append("多余字段: " + ", ".join(extra))
|
||||
raise HTTPException(status_code=400, detail="CSV 字段与模板不一致,导入失败。" + ";".join(parts))
|
||||
if database_source == "cell_data" and table_name.lower() == "sector":
|
||||
stats = db.upsert_csv(str(tmp_path), table_name, "CGI", encoding=encoding)
|
||||
return {
|
||||
"success": True,
|
||||
"message": (
|
||||
f"导入成功,共 {stats['imported_rows']} 行:"
|
||||
f"新增 {stats['inserted_rows']} 行,更新 {stats['updated_rows']} 行,"
|
||||
f"清理历史重复 {stats['removed_duplicate_rows']} 行"
|
||||
),
|
||||
**stats,
|
||||
}
|
||||
imported = db.import_csv(str(tmp_path), table_name, encoding=encoding)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
finally:
|
||||
remove_file_safely(tmp_path)
|
||||
return {"success": True, "message": f"导入成功,共 {imported} 行", "imported_rows": imported}
|
||||
|
||||
|
||||
@router.post("/api/database/execute")
|
||||
def execute_sql(
|
||||
sql: str = Body(..., embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
):
|
||||
db = _db(database_source)
|
||||
try:
|
||||
success, result = db.execute_sql(sql)
|
||||
if success:
|
||||
return {"success": True, "result": result}
|
||||
raise HTTPException(status_code=400, detail=result)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
# Sync def so FastAPI runs it in a threadpool: exporting large tables (and the Metrix
|
||||
# export-job polling) is blocking and would otherwise freeze the single-worker event loop.
|
||||
@router.post("/api/download")
|
||||
def download_table(
|
||||
table_name: Optional[str] = Body(None, embed=True),
|
||||
table_names: Optional[list[str]] = Body(None, embed=True),
|
||||
database_source: DatabaseSource = Body("main"),
|
||||
file_format: str = Body("csv", alias="format"),
|
||||
):
|
||||
if file_format not in {"csv", "xlsx"}:
|
||||
raise HTTPException(status_code=400, detail="不支持的导出格式")
|
||||
|
||||
requested_tables = _resolve_requested_tables(table_name, table_names)
|
||||
if not requested_tables:
|
||||
raise HTTPException(status_code=400, detail="请选择要导出的数据表")
|
||||
if file_format == "csv" and len(requested_tables) != 1:
|
||||
raise HTTPException(status_code=400, detail="CSV 每次只能导出一张表")
|
||||
|
||||
config = state.current_config()
|
||||
if database_source == "main" and config.warehouse_type == "metrix":
|
||||
return _download_via_metrix(config, requested_tables, file_format)
|
||||
|
||||
db = _db(database_source)
|
||||
try:
|
||||
available_tables = set(db.get_tables())
|
||||
missing_tables = [name for name in requested_tables if name not in available_tables]
|
||||
if missing_tables:
|
||||
raise HTTPException(status_code=400, detail=f"数据表不存在: {', '.join(missing_tables)}")
|
||||
|
||||
table_frames = {
|
||||
name: _dataframe_from_table(db, name)
|
||||
for name in requested_tables
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename_prefix = requested_tables[0] if len(requested_tables) == 1 else "tables"
|
||||
filename = f"{filename_prefix}_{timestamp}.{file_format}"
|
||||
filepath = CACHE_DIR / filename
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
if file_format == "csv":
|
||||
table_frames[requested_tables[0]].to_csv(filepath, index=False, encoding="utf-8-sig")
|
||||
media_type = "text/csv"
|
||||
else:
|
||||
used_sheet_names: set[str] = set()
|
||||
with pd.ExcelWriter(filepath) as writer:
|
||||
for name, df in table_frames.items():
|
||||
df.to_excel(writer, sheet_name=_make_sheet_name(name, used_sheet_names), index=False)
|
||||
media_type = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
|
||||
except Exception:
|
||||
remove_file_safely(filepath)
|
||||
raise
|
||||
|
||||
return FileResponse(
|
||||
path=str(filepath),
|
||||
filename=filename,
|
||||
media_type=media_type,
|
||||
background=BackgroundTask(remove_file_safely, filepath),
|
||||
)
|
||||
|
||||
|
||||
def _download_via_metrix(config, requested_tables: list[str], file_format: str) -> FileResponse:
|
||||
"""Metrix 仓库模式:用平台导出任务生成文件后流式返回(避免分页上限丢行)。"""
|
||||
metrix = config.metrix.normalized()
|
||||
client = make_client(metrix)
|
||||
try:
|
||||
job_id = client.submit_export(metrix.database_conn_id, requested_tables, file_format, metrix.target_database)
|
||||
job = client.wait_job(job_id)
|
||||
if job.get("status") != "success":
|
||||
raise HTTPException(status_code=500, detail=f"导出失败: {job.get('error_code') or job.get('status')}")
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename_prefix = requested_tables[0] if len(requested_tables) == 1 else "tables"
|
||||
filename = f"{filename_prefix}_{timestamp}.{file_format}"
|
||||
filepath = CACHE_DIR / filename
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
client.download_job_file(job_id, filepath)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
media_type = "text/csv" if file_format == "csv" else "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
|
||||
return FileResponse(
|
||||
path=str(filepath),
|
||||
filename=filename,
|
||||
media_type=media_type,
|
||||
background=BackgroundTask(remove_file_safely, filepath),
|
||||
)
|
||||
@@ -0,0 +1,38 @@
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app import state
|
||||
from app.database import DatabaseManager
|
||||
|
||||
|
||||
router = APIRouter(tags=["health"])
|
||||
|
||||
|
||||
@router.get("/health")
|
||||
async def health_check():
|
||||
checks = {
|
||||
"app": {"status": "ok"},
|
||||
"database": {"status": "unknown"},
|
||||
}
|
||||
|
||||
try:
|
||||
db_manager = DatabaseManager(state.current_config())
|
||||
server_info = db_manager.get_server_info()
|
||||
checks["database"] = {
|
||||
"status": "ok",
|
||||
"version": server_info.get("version", "unknown"),
|
||||
"load_data_infile": server_info.get("load_data_infile", False),
|
||||
}
|
||||
except Exception as exc:
|
||||
checks["database"] = {"status": "error", "message": str(exc)}
|
||||
|
||||
is_healthy = all(check.get("status") == "ok" for check in checks.values())
|
||||
return {
|
||||
"status": "healthy" if is_healthy else "unhealthy",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"version": "3.0.0",
|
||||
"uptime_pid": os.getpid(),
|
||||
"checks": checks,
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
import zipfile
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.background import BackgroundTask
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR
|
||||
from app.utils.files import format_size, get_dir_size, remove_file_safely
|
||||
|
||||
|
||||
router = APIRouter(tags=["history"])
|
||||
|
||||
|
||||
def _safe_filename_part(value: str) -> str:
|
||||
safe = "".join(char if char.isalnum() or char in {"-", "_"} else "_" for char in value)
|
||||
return safe.strip("_") or "history"
|
||||
|
||||
|
||||
def _get_safe_work_dir(record_id: str, *, require_finished: bool = True) -> tuple[Path, str]:
|
||||
record = state.history_manager.get(record_id)
|
||||
if not record:
|
||||
raise HTTPException(status_code=404, detail="记录不存在")
|
||||
if require_finished and record.status in {"pending", "processing"}:
|
||||
raise HTTPException(status_code=409, detail="任务尚未完成,暂不能下载历史数据")
|
||||
|
||||
work_dir = Path(record.work_dir).resolve()
|
||||
cache_dir = CACHE_DIR.resolve()
|
||||
try:
|
||||
work_dir.relative_to(cache_dir)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail="历史目录不合法") from exc
|
||||
|
||||
if not work_dir.exists() or not work_dir.is_dir():
|
||||
raise HTTPException(status_code=404, detail="历史数据目录不存在")
|
||||
|
||||
return work_dir, record.id
|
||||
|
||||
|
||||
def _get_safe_history_path(record_id: str, relative_path: str | None, *, require_finished: bool = False) -> tuple[Path, Path, str]:
|
||||
work_dir, safe_record_id = _get_safe_work_dir(record_id, require_finished=require_finished)
|
||||
clean_path = _normalize_relative_path(relative_path)
|
||||
target = work_dir.joinpath(*clean_path.split("/")).resolve() if clean_path else work_dir
|
||||
|
||||
try:
|
||||
target.relative_to(work_dir)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail="历史文件路径不合法") from exc
|
||||
|
||||
if not target.exists():
|
||||
raise HTTPException(status_code=404, detail="历史文件不存在")
|
||||
|
||||
return work_dir, target, safe_record_id
|
||||
|
||||
|
||||
def _normalize_relative_path(value: str | None) -> str:
|
||||
normalized = (value or "").replace("\\", "/").strip("/")
|
||||
if not normalized:
|
||||
return ""
|
||||
|
||||
parts = [part for part in normalized.split("/") if part and part != "."]
|
||||
if any(part == ".." for part in parts):
|
||||
raise HTTPException(status_code=400, detail="历史文件路径不合法")
|
||||
return "/".join(parts)
|
||||
|
||||
|
||||
def _zip_directory(source_dir: Path, archive_path: Path) -> None:
|
||||
with zipfile.ZipFile(archive_path, "w", compression=zipfile.ZIP_DEFLATED, allowZip64=True) as archive:
|
||||
for item in source_dir.rglob("*"):
|
||||
arcname = item.relative_to(source_dir).as_posix()
|
||||
if item.is_dir():
|
||||
archive.writestr(f"{arcname}/", "")
|
||||
elif item.is_file():
|
||||
archive.write(item, arcname)
|
||||
|
||||
|
||||
def _zip_history_item(source_path: Path, archive_path: Path) -> None:
|
||||
root_name = source_path.name or "history"
|
||||
with zipfile.ZipFile(archive_path, "w", compression=zipfile.ZIP_DEFLATED, allowZip64=True) as archive:
|
||||
if source_path.is_file():
|
||||
archive.write(source_path, root_name)
|
||||
return
|
||||
|
||||
has_content = False
|
||||
for item in source_path.rglob("*"):
|
||||
has_content = True
|
||||
arcname = Path(root_name, item.relative_to(source_path)).as_posix()
|
||||
if item.is_dir():
|
||||
archive.writestr(f"{arcname}/", "")
|
||||
elif item.is_file():
|
||||
archive.write(item, arcname)
|
||||
|
||||
if not has_content:
|
||||
archive.writestr(f"{root_name}/", "")
|
||||
|
||||
|
||||
def _history_entry(item: Path, work_dir: Path) -> dict[str, Any]:
|
||||
stat = item.stat()
|
||||
is_dir = item.is_dir()
|
||||
size = get_dir_size(item) if is_dir else stat.st_size
|
||||
modified = datetime.fromtimestamp(stat.st_mtime)
|
||||
return {
|
||||
"name": item.name,
|
||||
"path": item.relative_to(work_dir).as_posix(),
|
||||
"type": "dir" if is_dir else "file",
|
||||
"size": size,
|
||||
"size_formatted": format_size(size),
|
||||
"modified": modified.isoformat(),
|
||||
"modified_formatted": modified.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/api/history")
|
||||
async def get_history(limit: int = Body(50, embed=True)):
|
||||
return {"records": state.history_manager.list(limit)}
|
||||
|
||||
|
||||
@router.post("/api/history/delete")
|
||||
async def delete_history(record_id: str = Body(..., embed=True)):
|
||||
if state.history_manager.delete(record_id):
|
||||
return {"success": True, "message": "删除成功"}
|
||||
raise HTTPException(status_code=404, detail="记录不存在")
|
||||
|
||||
|
||||
@router.post("/api/history/clear")
|
||||
async def clear_history():
|
||||
count = state.history_manager.clear()
|
||||
return {"success": True, "deleted": count}
|
||||
|
||||
|
||||
@router.post("/api/history/detail")
|
||||
async def get_history_detail(record_id: str = Body(..., embed=True)):
|
||||
record = state.history_manager.get(record_id)
|
||||
if not record:
|
||||
raise HTTPException(status_code=404, detail="记录不存在")
|
||||
|
||||
result = record.to_dict()
|
||||
result["logs"] = state.history_manager.get_logs(record_id)
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/api/history/size")
|
||||
async def get_history_size(record_id: str = Body(..., embed=True)):
|
||||
record = state.history_manager.get(record_id)
|
||||
if not record:
|
||||
raise HTTPException(status_code=404, detail="记录不存在")
|
||||
|
||||
work_dir = Path(record.work_dir)
|
||||
if not work_dir.exists():
|
||||
return {"success": True, "size": 0, "size_formatted": "0 B"}
|
||||
|
||||
size = get_dir_size(work_dir)
|
||||
return {"success": True, "size": size, "size_formatted": format_size(size)}
|
||||
|
||||
|
||||
@router.post("/api/history/files")
|
||||
async def list_history_files(
|
||||
record_id: str = Body(..., embed=True),
|
||||
path: str | None = Body("", embed=True),
|
||||
):
|
||||
work_dir, current_path, _safe_record_id = _get_safe_history_path(record_id, path)
|
||||
if not current_path.is_dir():
|
||||
raise HTTPException(status_code=400, detail="请选择目录")
|
||||
|
||||
current_relative = current_path.relative_to(work_dir).as_posix()
|
||||
current_relative = "" if current_relative == "." else current_relative
|
||||
parent_relative = None
|
||||
if current_path != work_dir:
|
||||
parent_relative = current_path.parent.relative_to(work_dir).as_posix()
|
||||
parent_relative = "" if parent_relative == "." else parent_relative
|
||||
|
||||
entries = [_history_entry(item, work_dir) for item in current_path.iterdir()]
|
||||
entries.sort(key=lambda item: (item["type"] != "dir", item["name"].lower()))
|
||||
return {
|
||||
"success": True,
|
||||
"record_id": record_id,
|
||||
"path": current_relative,
|
||||
"parent_path": parent_relative,
|
||||
"entries": entries,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/api/history/download")
|
||||
async def download_history(record_id: str = Body(..., embed=True)):
|
||||
work_dir, safe_record_id = _get_safe_work_dir(record_id)
|
||||
export_dir = CACHE_DIR / ".downloads"
|
||||
export_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"{_safe_filename_part(safe_record_id)}_{timestamp}.zip"
|
||||
archive_path = export_dir / filename
|
||||
|
||||
try:
|
||||
_zip_directory(work_dir, archive_path)
|
||||
except Exception:
|
||||
remove_file_safely(archive_path)
|
||||
raise
|
||||
|
||||
return FileResponse(
|
||||
path=str(archive_path),
|
||||
filename=filename,
|
||||
media_type="application/zip",
|
||||
background=BackgroundTask(remove_file_safely, archive_path),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/api/history/file/download")
|
||||
async def download_history_file(
|
||||
record_id: str = Body(..., embed=True),
|
||||
path: str = Body(..., embed=True),
|
||||
):
|
||||
_work_dir, target_path, safe_record_id = _get_safe_history_path(record_id, path, require_finished=True)
|
||||
if target_path.is_file():
|
||||
return FileResponse(
|
||||
path=str(target_path),
|
||||
filename=target_path.name,
|
||||
media_type="application/octet-stream",
|
||||
)
|
||||
|
||||
if not target_path.is_dir():
|
||||
raise HTTPException(status_code=400, detail="历史文件类型不支持下载")
|
||||
|
||||
export_dir = CACHE_DIR / ".downloads"
|
||||
export_dir.mkdir(parents=True, exist_ok=True)
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"{_safe_filename_part(safe_record_id)}_{_safe_filename_part(target_path.name)}_{timestamp}.zip"
|
||||
archive_path = export_dir / filename
|
||||
|
||||
try:
|
||||
_zip_history_item(target_path, archive_path)
|
||||
except Exception:
|
||||
remove_file_safely(archive_path)
|
||||
raise
|
||||
|
||||
return FileResponse(
|
||||
path=str(archive_path),
|
||||
filename=filename,
|
||||
media_type="application/zip",
|
||||
background=BackgroundTask(remove_file_safely, archive_path),
|
||||
)
|
||||
@@ -0,0 +1,26 @@
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
|
||||
from app.services.license import InvalidActivationCodeError, activate, get_license_info
|
||||
|
||||
|
||||
router = APIRouter(tags=["license"])
|
||||
|
||||
|
||||
@router.get("/api/license/status")
|
||||
async def get_license_status():
|
||||
info = get_license_info()
|
||||
return {"success": True, **info.to_dict()}
|
||||
|
||||
|
||||
@router.post("/api/license/activate")
|
||||
async def activate_license(code: str = Body(..., embed=True)):
|
||||
try:
|
||||
info = activate(code)
|
||||
except InvalidActivationCodeError as exc:
|
||||
raise HTTPException(status_code=400, detail=exc.to_detail()) from exc
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": "激活成功,到期日期已延长 30 天",
|
||||
**info.to_dict(),
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
import time
|
||||
from datetime import date, datetime
|
||||
from pathlib import Path
|
||||
from threading import Thread
|
||||
from typing import Any, Callable, Iterable
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
|
||||
from app import state
|
||||
from app.api.routers.task_runtime import (
|
||||
apply_history_retention_safely,
|
||||
log_license_check,
|
||||
set_task_stage,
|
||||
)
|
||||
from app.config import AppConfig, CACHE_DIR, RemoteDataConfig
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
from app.config import CELLDATA_SCRIPT
|
||||
from app.services.cell_data import copy_celldata_tables_to_capacity, execute_celldata_script, refresh_cell_data
|
||||
from app.services.license import LicenseError, check_processing_allowed
|
||||
from app.services.platform import PlatformStorageDownloader, make_source_downloader
|
||||
from app.services.pipeline import RESULT_TABLES, run_import_and_report
|
||||
from app.services.remote_download import RemoteDataDownloader
|
||||
|
||||
|
||||
router = APIRouter(tags=["remote"])
|
||||
|
||||
|
||||
@router.post("/api/remote/test")
|
||||
def test_remote_connection(config: dict[str, Any] | None = Body(None)):
|
||||
app_config = state.current_config()
|
||||
try:
|
||||
if app_config.source_type == "metrix":
|
||||
PlatformStorageDownloader(app_config).test_connection()
|
||||
return {"success": True, "message": "平台储存连接成功"}
|
||||
# FTP/SFTP: test the posted form config if provided, else the saved one.
|
||||
remote_config = RemoteDataConfig.from_dict(config) if config else app_config.remote_data
|
||||
RemoteDataDownloader(remote_config).test_connection()
|
||||
return {"success": True, "message": "远程服务器连接成功"}
|
||||
except Exception as exc:
|
||||
return {"success": False, "message": f"连接失败: {exc}"}
|
||||
|
||||
|
||||
@router.post("/api/remote/start")
|
||||
async def start_remote_processing():
|
||||
return start_remote_processing_job(source="manual")
|
||||
|
||||
|
||||
@router.get("/api/remote/scheduler/status")
|
||||
def get_scheduler_status():
|
||||
if state.auto_scheduler is None:
|
||||
return {"enabled": False, "running": False, "message": "自动调度器未启动"}
|
||||
return state.auto_scheduler.get_status()
|
||||
|
||||
|
||||
@router.post("/api/remote/scheduler/trigger")
|
||||
def trigger_scheduler_check():
|
||||
if state.auto_scheduler is None:
|
||||
raise HTTPException(status_code=503, detail="自动调度器未启动")
|
||||
return state.auto_scheduler.check_and_run(manual=True)
|
||||
|
||||
|
||||
def start_remote_processing_job(
|
||||
*,
|
||||
source: str = "manual",
|
||||
on_finish: Callable[[str, str], None] | None = None,
|
||||
target_dates: Iterable[date] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
|
||||
|
||||
app_config = state.current_config()
|
||||
remote_config = app_config.remote_data.normalized()
|
||||
if not remote_config.enabled:
|
||||
raise HTTPException(status_code=400, detail="请先启用远程数据配置")
|
||||
|
||||
task_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
work_dir = CACHE_DIR / task_id
|
||||
work_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
logs: list[str] = []
|
||||
current_stage = "downloading"
|
||||
|
||||
def log_callback(message: str) -> None:
|
||||
logs.append(message)
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
def stage_callback(stage: str) -> None:
|
||||
nonlocal current_stage
|
||||
current_stage = stage
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
logger = ProcessLogger(
|
||||
log_file=work_dir / "log.txt",
|
||||
callback=log_callback,
|
||||
stage_callback=stage_callback,
|
||||
)
|
||||
state.history_manager.create(work_dir, 0, record_id=task_id)
|
||||
state.history_manager.update(task_id, status="processing")
|
||||
state.processing_tasks[task_id] = {"logs": [], "status": "processing", "stage": current_stage}
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": "downloading",
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
thread = Thread(
|
||||
target=_run_remote_processing,
|
||||
args=(task_id, work_dir, app_config, remote_config, logger, source, on_finish, target_dates),
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": "自动调度远程下载处理任务已启动" if source == "scheduler" else "远程下载处理任务已启动",
|
||||
"task_id": task_id,
|
||||
"stage": "downloading",
|
||||
}
|
||||
|
||||
|
||||
def _run_remote_processing(
|
||||
task_id: str,
|
||||
work_dir: Path,
|
||||
app_config: AppConfig,
|
||||
remote_config: RemoteDataConfig,
|
||||
logger: ProcessLogger,
|
||||
source: str = "manual",
|
||||
on_finish: Callable[[str, str], None] | None = None,
|
||||
target_dates: Iterable[date] | None = None,
|
||||
) -> None:
|
||||
final_status = "failed"
|
||||
try:
|
||||
if app_config.source_type == "metrix":
|
||||
logger.info(f"开始从平台储存下载,目录: {remote_config.remote_dir}")
|
||||
else:
|
||||
logger.info(
|
||||
f"开始远程下载,协议: {remote_config.protocol.upper()},"
|
||||
f"服务器: {remote_config.host}:{remote_config.port},目录: {remote_config.remote_dir}"
|
||||
)
|
||||
downloader = make_source_downloader(app_config, logger.info)
|
||||
download_result = downloader.download_to(work_dir, target_dates=target_dates)
|
||||
logger.success(
|
||||
f"下载完成,共 {download_result.file_count} 个文件,"
|
||||
f"{_format_bytes(download_result.total_bytes)}"
|
||||
)
|
||||
|
||||
if download_result.file_count == 0:
|
||||
raise RuntimeError("源目录中未下载到任何文件")
|
||||
|
||||
state.history_manager.update(task_id, file_count=download_result.file_count)
|
||||
|
||||
_try_refresh_cell_data(app_config, work_dir, logger)
|
||||
|
||||
if app_config.warehouse_type == "metrix":
|
||||
started = time.time()
|
||||
run_import_and_report(work_dir, app_config, logger)
|
||||
status, error, elapsed = "completed", None, round(time.time() - started, 2)
|
||||
else:
|
||||
logger.set_stage("license")
|
||||
log_license_check(logger, check_processing_allowed(work_dir))
|
||||
processor = DataProcessor(app_config, work_dir, logger)
|
||||
result = processor.process()
|
||||
status = "completed" if result.get("success") else "failed"
|
||||
error = result.get("error")
|
||||
elapsed = result.get("elapsed_time", 0)
|
||||
if status == "completed" and remote_config.auto_delete_source:
|
||||
try:
|
||||
deleted_count = downloader.delete_source_files(download_result.remote_files)
|
||||
logger.success(f"远程源文件清理完成,共删除 {deleted_count} 个文件,目录已保留")
|
||||
final_status = status
|
||||
except Exception as exc:
|
||||
logger.warning(f"远程源文件清理失败,数据处理结果已保留: {exc}")
|
||||
final_status = "source_cleanup_failed"
|
||||
else:
|
||||
final_status = status
|
||||
|
||||
state.history_manager.update(
|
||||
task_id,
|
||||
status=status,
|
||||
elapsed_time=elapsed,
|
||||
error=error,
|
||||
result_tables=RESULT_TABLES,
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": status,
|
||||
"stage": status,
|
||||
"error": error,
|
||||
}
|
||||
except Exception as exc:
|
||||
error_detail = exc.to_detail() if isinstance(exc, LicenseError) else None
|
||||
if source == "scheduler":
|
||||
logger.error(f"自动调度任务失败: {exc}")
|
||||
else:
|
||||
logger.error(f"远程自动化任务失败: {exc}")
|
||||
state.history_manager.update(task_id, status="failed", error=str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"error_detail": error_detail,
|
||||
}
|
||||
finally:
|
||||
apply_history_retention_safely()
|
||||
state.reset_task_lock()
|
||||
if on_finish:
|
||||
on_finish(task_id, final_status)
|
||||
|
||||
|
||||
def _try_refresh_cell_data(app_config: AppConfig, work_dir: Path, logger: ProcessLogger) -> None:
|
||||
# 1) 远程拉取 CellData(仅在配置了 FTP/SFTP 时);失败不阻断容量处理
|
||||
if app_config.cell_data.remote_data.enabled:
|
||||
logger.set_stage("cell_data")
|
||||
logger.info("── CellData 更新 ──")
|
||||
try:
|
||||
result = refresh_cell_data(app_config, work_dir, logger)
|
||||
logger.success(
|
||||
f"CellData 更新完成:{result.imported_rows} 行"
|
||||
f"(解析 {result.parsed_rows},跳过 {result.skipped_rows})"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(f"CellData 远程更新失败,继续容量处理: {exc}")
|
||||
|
||||
# 2) 无论是否远程拉取,都执行 CellData 脚本并把表同步到容量库(覆盖手动上传场景);
|
||||
# celldata 无数据 / 连不上则内部跳过,best-effort 不阻断容量处理
|
||||
logger.set_stage("cell_data")
|
||||
try:
|
||||
execute_celldata_script(CELLDATA_SCRIPT, app_config, logger)
|
||||
copy_celldata_tables_to_capacity(app_config, logger)
|
||||
except Exception as exc:
|
||||
logger.warning(f"CellData 同步到容量库失败,继续容量处理: {exc}")
|
||||
|
||||
|
||||
def _format_bytes(size: int) -> str:
|
||||
value = float(size)
|
||||
for unit in ("B", "KB", "MB", "GB"):
|
||||
if value < 1024:
|
||||
return f"{value:.1f} {unit}"
|
||||
value /= 1024
|
||||
return f"{value:.1f} TB"
|
||||
@@ -0,0 +1,141 @@
|
||||
import shutil
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from threading import Thread
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
|
||||
from app import state
|
||||
from app.api.routers.task_runtime import set_task_stage
|
||||
from app.config import AppConfig, CACHE_DIR, CELLDATA_SCRIPT, SQL_SCRIPT
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
|
||||
|
||||
router = APIRouter(tags=["script"])
|
||||
|
||||
|
||||
def _resolve_script_path(script_type: str) -> tuple[Path, str]:
|
||||
paths = {"report": SQL_SCRIPT, "celldata": CELLDATA_SCRIPT}
|
||||
labels = {"report": "容量报表脚本", "celldata": "CellData 脚本"}
|
||||
path = paths.get(script_type)
|
||||
if path is None:
|
||||
raise HTTPException(status_code=400, detail=f"不支持的脚本类型: {script_type}")
|
||||
return path, labels.get(script_type, script_type)
|
||||
|
||||
|
||||
@router.post("/api/script/content")
|
||||
def get_script_content(script_type: str = Body("report", embed=True)):
|
||||
script_path, label = _resolve_script_path(script_type)
|
||||
try:
|
||||
if not script_path.exists():
|
||||
return {
|
||||
"success": True,
|
||||
"content": f"# {label}文件不存在,请在此编写脚本\n",
|
||||
"modified": None,
|
||||
"path": str(script_path),
|
||||
}
|
||||
|
||||
modified = datetime.fromtimestamp(script_path.stat().st_mtime).strftime("%Y-%m-%d %H:%M:%S")
|
||||
return {
|
||||
"success": True,
|
||||
"content": script_path.read_text(encoding="utf-8"),
|
||||
"modified": modified,
|
||||
"path": str(script_path),
|
||||
}
|
||||
except Exception as exc:
|
||||
return {"success": False, "error": str(exc)}
|
||||
|
||||
|
||||
@router.post("/api/script/execute")
|
||||
async def execute_script(script_type: str = Body("report", embed=True)):
|
||||
script_path, _ = _resolve_script_path(script_type)
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行,请等待完成")
|
||||
|
||||
task_id = f"script_{uuid.uuid4().hex[:8]}"
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": "processing",
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
logs: list[str] = []
|
||||
|
||||
def log_callback(message: str) -> None:
|
||||
logs.append(message)
|
||||
set_task_stage(task_id, "processing", logs)
|
||||
|
||||
logger = ProcessLogger(log_file=None, callback=log_callback)
|
||||
set_task_stage(task_id, "processing", logs)
|
||||
app_config = state.current_config()
|
||||
|
||||
thread = Thread(
|
||||
target=_run_script,
|
||||
args=(task_id, logger, logs, app_config, script_path, script_type),
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
_, label = _resolve_script_path(script_type)
|
||||
return {"success": True, "message": f"{label}执行任务已启动", "task_id": task_id}
|
||||
|
||||
|
||||
@router.post("/api/script/save")
|
||||
async def save_script_content(
|
||||
content: str = Body(...),
|
||||
script_type: str = Body("report"),
|
||||
):
|
||||
script_path, _ = _resolve_script_path(script_type)
|
||||
try:
|
||||
if script_path.exists():
|
||||
shutil.copy(script_path, script_path.with_suffix(".sql.bak"))
|
||||
|
||||
script_path.write_text(content, encoding="utf-8")
|
||||
modified = datetime.fromtimestamp(script_path.stat().st_mtime).strftime("%Y-%m-%d %H:%M:%S")
|
||||
return {"success": True, "message": "脚本保存成功", "modified": modified}
|
||||
except Exception as exc:
|
||||
return {"success": False, "error": str(exc)}
|
||||
|
||||
|
||||
def _run_script(
|
||||
task_id: str,
|
||||
logger: ProcessLogger,
|
||||
logs: list[str],
|
||||
app_config: AppConfig,
|
||||
script_path: Path,
|
||||
script_type: str,
|
||||
) -> None:
|
||||
temp_work_dir: Path | None = None
|
||||
labels = {"report": "容量报表脚本", "celldata": "CellData 脚本"}
|
||||
label = labels.get(script_type, script_type)
|
||||
try:
|
||||
logger.info(f"开始执行{label}...")
|
||||
if script_type == "celldata":
|
||||
_run_celldata_script(script_path, app_config, logger)
|
||||
elif app_config.warehouse_type == "metrix":
|
||||
from app.services.pipeline import run_report_sql
|
||||
|
||||
run_report_sql(app_config, logger)
|
||||
else:
|
||||
temp_work_dir = CACHE_DIR / task_id
|
||||
temp_work_dir.mkdir(parents=True, exist_ok=True)
|
||||
processor = DataProcessor(app_config, temp_work_dir, logger)
|
||||
processor._execute_sql_script()
|
||||
logger.success(f"{label}执行完成")
|
||||
set_task_stage(task_id, "completed", logs, status="completed")
|
||||
except Exception as exc:
|
||||
logger.error(f"{label}执行失败: {exc}")
|
||||
set_task_stage(task_id, "failed", logs, status="failed")
|
||||
finally:
|
||||
if temp_work_dir and temp_work_dir.exists():
|
||||
shutil.rmtree(temp_work_dir, ignore_errors=True)
|
||||
state.reset_task_lock()
|
||||
|
||||
|
||||
def _run_celldata_script(script_path: Path, app_config: AppConfig, logger: ProcessLogger) -> None:
|
||||
from app.services.cell_data import execute_celldata_script
|
||||
|
||||
execute_celldata_script(script_path, app_config, logger)
|
||||
@@ -0,0 +1,36 @@
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app import state
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def set_task_stage(task_id: str, stage: str, logs: list[str], status: str = "processing") -> None:
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": logs.copy(),
|
||||
"status": status,
|
||||
"stage": stage,
|
||||
}
|
||||
if state.global_task_lock["task_id"] == task_id:
|
||||
state.global_task_lock["stage"] = stage
|
||||
|
||||
|
||||
def log_license_check(logger: Any, info: Any) -> None:
|
||||
if info.current_date:
|
||||
logger.info(
|
||||
f"授权校验通过,数据日期: {info.current_date.isoformat()},"
|
||||
f"到期日期: {info.expires_on.isoformat()}"
|
||||
)
|
||||
elif info.zip_count:
|
||||
logger.warning("未从 ZIP 文件名识别到日期,已跳过授权日期比对")
|
||||
else:
|
||||
logger.warning("未找到 ZIP 文件,已跳过授权日期比对")
|
||||
|
||||
|
||||
def apply_history_retention_safely() -> None:
|
||||
try:
|
||||
state.apply_history_retention()
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to apply history retention: %s", exc)
|
||||
@@ -0,0 +1,241 @@
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from threading import Thread
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
|
||||
from app import state
|
||||
from app.api.routers.task_runtime import (
|
||||
apply_history_retention_safely,
|
||||
log_license_check,
|
||||
set_task_stage,
|
||||
)
|
||||
from app.config import AppConfig
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
from app.config import CELLDATA_SCRIPT
|
||||
from app.services.cell_data import copy_celldata_tables_to_capacity, execute_celldata_script, refresh_cell_data
|
||||
from app.services.license import LicenseError, check_processing_allowed
|
||||
|
||||
|
||||
router = APIRouter(tags=["tasks"])
|
||||
|
||||
|
||||
@router.get("/api/task/status")
|
||||
async def get_global_task_status():
|
||||
if state.global_task_lock["locked"]:
|
||||
task_id = state.global_task_lock["task_id"]
|
||||
if task_id and _task_finished(task_id):
|
||||
state.reset_task_lock()
|
||||
return {"has_active": False}
|
||||
|
||||
return {
|
||||
"has_active": True,
|
||||
"task_id": task_id,
|
||||
"stage": state.global_task_lock["stage"],
|
||||
"started_at": state.global_task_lock["started_at"],
|
||||
"logs": [],
|
||||
}
|
||||
|
||||
active_tasks = {
|
||||
task_id: task
|
||||
for task_id, task in state.processing_tasks.items()
|
||||
if task.get("status") == "processing"
|
||||
}
|
||||
if active_tasks:
|
||||
task_id = next(iter(active_tasks))
|
||||
return {
|
||||
"has_active": True,
|
||||
"task_id": task_id,
|
||||
"stage": active_tasks[task_id].get("stage", "processing"),
|
||||
"logs": active_tasks[task_id].get("logs", []),
|
||||
}
|
||||
|
||||
return {"has_active": False}
|
||||
|
||||
|
||||
@router.post("/api/task/lock")
|
||||
async def lock_task(task_id: str = Body(..., embed=True)):
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行")
|
||||
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": "uploading",
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
return {"success": True, "message": "任务已锁定"}
|
||||
|
||||
|
||||
@router.post("/api/task/unlock")
|
||||
async def unlock_task(task_id: str | None = Body(None, embed=True)):
|
||||
if task_id and state.global_task_lock["task_id"] != task_id:
|
||||
raise HTTPException(status_code=403, detail="无权解锁此任务")
|
||||
|
||||
state.reset_task_lock()
|
||||
return {"success": True, "message": "任务已解锁"}
|
||||
|
||||
|
||||
@router.get("/api/process/active")
|
||||
async def get_active_task():
|
||||
return await get_global_task_status()
|
||||
|
||||
|
||||
@router.post("/api/process/start")
|
||||
async def start_processing(task_id: str = Body(..., embed=True)):
|
||||
record = state.history_manager.get(task_id)
|
||||
if not record:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if record.status == "processing":
|
||||
raise HTTPException(status_code=400, detail="任务正在处理中")
|
||||
|
||||
work_dir = Path(record.work_dir)
|
||||
if not work_dir.exists():
|
||||
raise HTTPException(status_code=400, detail="工作目录不存在")
|
||||
|
||||
logs: list[str] = []
|
||||
current_stage = "processing"
|
||||
|
||||
def log_callback(message: str) -> None:
|
||||
logs.append(message)
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
def stage_callback(stage: str) -> None:
|
||||
nonlocal current_stage
|
||||
current_stage = stage
|
||||
set_task_stage(task_id, current_stage, logs)
|
||||
|
||||
logger = ProcessLogger(
|
||||
log_file=work_dir / "log.txt",
|
||||
callback=log_callback,
|
||||
stage_callback=stage_callback,
|
||||
)
|
||||
app_config = state.current_config()
|
||||
state.history_manager.update(task_id, status="processing")
|
||||
state.processing_tasks[task_id] = {"logs": [], "status": "processing", "stage": current_stage}
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": task_id,
|
||||
"stage": "processing",
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
thread = Thread(target=_run_processing, args=(task_id, work_dir, logger, app_config), daemon=True)
|
||||
thread.start()
|
||||
|
||||
return {"success": True, "message": "处理任务已启动", "task_id": task_id}
|
||||
|
||||
|
||||
@router.post("/api/process/status")
|
||||
async def get_processing_status(task_id: str = Body(..., embed=True)):
|
||||
if task_id in state.processing_tasks:
|
||||
task_info = state.processing_tasks[task_id]
|
||||
logs = task_info.get("logs") or state.history_manager.get_logs(task_id)
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": task_info["status"],
|
||||
"stage": task_info.get("stage"),
|
||||
"logs": logs,
|
||||
"error": task_info.get("error"),
|
||||
"error_detail": task_info.get("error_detail"),
|
||||
}
|
||||
|
||||
record = state.history_manager.get(task_id)
|
||||
if not record:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": record.status,
|
||||
"stage": record.status,
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"elapsed_time": record.elapsed_time,
|
||||
"error": record.error,
|
||||
}
|
||||
|
||||
|
||||
def _task_finished(task_id: str) -> bool:
|
||||
record = state.history_manager.get(task_id)
|
||||
if record and record.status in {"completed", "failed"}:
|
||||
return True
|
||||
|
||||
task_info: dict[str, Any] | None = state.processing_tasks.get(task_id)
|
||||
return bool(task_info and task_info.get("status") in {"completed", "failed"})
|
||||
|
||||
|
||||
def _run_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
|
||||
try:
|
||||
_try_refresh_cell_data(app_config, work_dir, logger)
|
||||
if app_config.warehouse_type == "metrix":
|
||||
import time
|
||||
|
||||
from app.services.pipeline import RESULT_TABLES, run_import_and_report
|
||||
|
||||
started = time.time()
|
||||
run_import_and_report(work_dir, app_config, logger)
|
||||
status, error, elapsed, result_tables = "completed", None, round(time.time() - started, 2), RESULT_TABLES
|
||||
else:
|
||||
logger.set_stage("license")
|
||||
log_license_check(logger, check_processing_allowed(work_dir))
|
||||
processor = DataProcessor(app_config, work_dir, logger)
|
||||
result = processor.process()
|
||||
status = "completed" if result.get("success") else "failed"
|
||||
error = result.get("error")
|
||||
elapsed = result.get("elapsed_time", 0)
|
||||
result_tables = ["4G_结果表", "5G_结果表"]
|
||||
state.history_manager.update(
|
||||
task_id,
|
||||
status=status,
|
||||
elapsed_time=elapsed,
|
||||
error=error,
|
||||
result_tables=result_tables,
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": status,
|
||||
"stage": status,
|
||||
"error": error,
|
||||
}
|
||||
except Exception as exc:
|
||||
error_detail = exc.to_detail() if isinstance(exc, LicenseError) else None
|
||||
logger.error(str(exc))
|
||||
state.history_manager.update(task_id, status="failed", error=str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"error_detail": error_detail,
|
||||
}
|
||||
finally:
|
||||
apply_history_retention_safely()
|
||||
state.reset_task_lock()
|
||||
|
||||
|
||||
def _try_refresh_cell_data(app_config: AppConfig, work_dir: Path, logger: ProcessLogger) -> None:
|
||||
# 1) 远程拉取 CellData(仅在配置了 FTP/SFTP 时);失败不阻断容量处理
|
||||
if app_config.cell_data.remote_data.enabled:
|
||||
logger.set_stage("cell_data")
|
||||
logger.info("── CellData 更新 ──")
|
||||
try:
|
||||
result = refresh_cell_data(app_config, work_dir, logger)
|
||||
logger.success(
|
||||
f"CellData 更新完成:{result.imported_rows} 行"
|
||||
f"(解析 {result.parsed_rows},跳过 {result.skipped_rows})"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(f"CellData 远程更新失败,继续容量处理: {exc}")
|
||||
|
||||
# 2) 无论是否远程拉取,都执行 CellData 脚本并把表同步到容量库(覆盖手动上传场景);
|
||||
# celldata 无数据 / 连不上则内部跳过,best-effort 不阻断容量处理
|
||||
logger.set_stage("cell_data")
|
||||
try:
|
||||
execute_celldata_script(CELLDATA_SCRIPT, app_config, logger)
|
||||
copy_celldata_tables_to_capacity(app_config, logger)
|
||||
except Exception as exc:
|
||||
logger.warning(f"CellData 同步到容量库失败,继续容量处理: {exc}")
|
||||
@@ -0,0 +1,112 @@
|
||||
from datetime import datetime
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import APIRouter, File, HTTPException, UploadFile
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR
|
||||
from app.utils.files import safe_relative_path
|
||||
|
||||
|
||||
router = APIRouter(tags=["upload"])
|
||||
|
||||
|
||||
@router.post("/api/upload/create")
|
||||
async def create_upload_session():
|
||||
session_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
work_dir = CACHE_DIR / session_id
|
||||
work_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
state.upload_sessions[session_id] = {
|
||||
"work_dir": work_dir,
|
||||
"files": [],
|
||||
"created_at": datetime.now().isoformat(),
|
||||
}
|
||||
|
||||
return {"success": True, "session_id": session_id, "work_dir": str(work_dir)}
|
||||
|
||||
|
||||
@router.post("/api/upload")
|
||||
async def upload_files(
|
||||
files: list[UploadFile] = File(...),
|
||||
session_id: Optional[str] = None,
|
||||
):
|
||||
if not files:
|
||||
raise HTTPException(status_code=400, detail="没有上传文件")
|
||||
|
||||
is_new_session = False
|
||||
if not session_id or session_id not in state.upload_sessions:
|
||||
if state.global_task_lock["locked"]:
|
||||
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
|
||||
|
||||
session_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
work_dir = CACHE_DIR / session_id
|
||||
work_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
state.global_task_lock.update(
|
||||
{
|
||||
"locked": True,
|
||||
"task_id": session_id,
|
||||
"stage": "uploading",
|
||||
"started_at": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
is_new_session = True
|
||||
|
||||
state.upload_sessions[session_id] = {
|
||||
"work_dir": work_dir,
|
||||
"files": [],
|
||||
"created_at": datetime.now().isoformat(),
|
||||
}
|
||||
else:
|
||||
if state.global_task_lock["locked"] and state.global_task_lock["task_id"] != session_id:
|
||||
raise HTTPException(status_code=409, detail="已有其他任务在运行")
|
||||
|
||||
session: dict[str, Any] = state.upload_sessions[session_id]
|
||||
work_dir = session["work_dir"]
|
||||
|
||||
try:
|
||||
saved_files: list[str] = []
|
||||
for file in files:
|
||||
if not file.filename:
|
||||
continue
|
||||
|
||||
relative_path = safe_relative_path(file.filename)
|
||||
file_path = work_dir / relative_path
|
||||
file_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
file_path.write_bytes(await file.read())
|
||||
|
||||
saved_name = str(relative_path).replace("\\", "/")
|
||||
saved_files.append(saved_name)
|
||||
session["files"].append(saved_name)
|
||||
|
||||
record = state.history_manager.get(session_id)
|
||||
if record:
|
||||
state.history_manager.update(session_id, file_count=len(session["files"]))
|
||||
else:
|
||||
state.history_manager.create(work_dir, len(session["files"]), record_id=session_id)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"task_id": session_id,
|
||||
"session_id": session_id,
|
||||
"work_dir": str(work_dir),
|
||||
"file_count": len(saved_files),
|
||||
"total_files": len(session["files"]),
|
||||
"files": saved_files,
|
||||
}
|
||||
except Exception as exc:
|
||||
if is_new_session:
|
||||
state.reset_task_lock()
|
||||
raise HTTPException(status_code=500, detail=f"上传失败: {exc}") from exc
|
||||
|
||||
|
||||
@router.post("/api/upload/complete/{session_id}")
|
||||
async def complete_upload_session(session_id: str):
|
||||
if session_id not in state.upload_sessions:
|
||||
raise HTTPException(status_code=404, detail="上传会话不存在")
|
||||
|
||||
session = state.upload_sessions[session_id]
|
||||
state.history_manager.update(session_id, file_count=len(session["files"]))
|
||||
|
||||
return {"success": True, "session_id": session_id, "total_files": len(session["files"])}
|
||||
Reference in New Issue
Block a user