feat: 双模式集成并精简 API 文档/Token、设置页卡片自适应
- 源/仓库各自可在直连(FTP/SFTP、MySQL)与 Metrix 存储/数据库平台间独立选择,两侧互不依赖 - Metrix 模式下源走平台储存 API、仓库走平台导入 + run-script(single_session)、查看导出代理到平台 - 去掉对外 API 文档与 API Token(前后端 + auth/config 解耦),业务接口仅登录态可访问 - 授权默认到期日改为 2026-12-30 - 设置页卡片改横向自适应(宽屏并排、窄屏换行),处理历史保留卡片收窄
This commit is contained in:
@@ -1,135 +0,0 @@
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException, Request
|
||||
|
||||
from app.auth import resolve_login_context
|
||||
from app.services.api_tokens import create_token, delete_token, delete_tokens, list_tokens, regenerate_token, update_token
|
||||
|
||||
|
||||
router = APIRouter(tags=["api-tokens"])
|
||||
|
||||
|
||||
def _require_login(request: Request) -> None:
|
||||
if resolve_login_context(request) is None:
|
||||
raise HTTPException(status_code=401, detail="未登录或登录已过期")
|
||||
|
||||
|
||||
def _bad_expiration_error(exc: ValueError) -> HTTPException:
|
||||
return HTTPException(status_code=400, detail="到期日期格式无效,请使用 YYYY-MM-DD 或 ISO 日期时间")
|
||||
|
||||
|
||||
def _resolve_expires_at(payload: dict[str, Any]) -> str | None:
|
||||
if bool(payload.get("permanent", False)):
|
||||
return None
|
||||
|
||||
expires_at = str(payload.get("expires_at") or "").strip()
|
||||
if not expires_at:
|
||||
raise HTTPException(status_code=400, detail="请选择 Token 到期日期,或设置为永久有效")
|
||||
return expires_at
|
||||
|
||||
|
||||
@router.get("/api/tokens")
|
||||
async def get_tokens(request: Request):
|
||||
_require_login(request)
|
||||
return {"success": True, "tokens": list_tokens()}
|
||||
|
||||
|
||||
@router.post("/api/tokens/create")
|
||||
async def create_api_token(request: Request, payload: dict[str, Any] = Body(...)):
|
||||
_require_login(request)
|
||||
name = str(payload.get("name", "")).strip()
|
||||
enabled = bool(payload.get("enabled", True))
|
||||
raw_expires_at = _resolve_expires_at(payload)
|
||||
|
||||
try:
|
||||
raw_token, record = create_token(
|
||||
name=name,
|
||||
expires_at=raw_expires_at,
|
||||
enabled=enabled,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise _bad_expiration_error(exc) from exc
|
||||
return {
|
||||
"success": True,
|
||||
"message": "API Token 创建成功",
|
||||
"token": raw_token,
|
||||
"record": record,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/api/tokens/update")
|
||||
async def update_api_token(request: Request, payload: dict[str, Any] = Body(...)):
|
||||
_require_login(request)
|
||||
token_id = str(payload.get("id", "")).strip()
|
||||
if not token_id:
|
||||
raise HTTPException(status_code=400, detail="缺少 Token ID")
|
||||
|
||||
changes: dict[str, Any] = {}
|
||||
if "name" in payload:
|
||||
changes["name"] = payload.get("name")
|
||||
if "enabled" in payload:
|
||||
changes["enabled"] = payload.get("enabled")
|
||||
if "permanent" in payload or "expires_at" in payload:
|
||||
changes["expires_at"] = _resolve_expires_at(payload)
|
||||
|
||||
try:
|
||||
record = update_token(token_id, **changes)
|
||||
return {"success": True, "message": "API Token 已更新", "record": record}
|
||||
except ValueError as exc:
|
||||
raise _bad_expiration_error(exc) from exc
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/tokens/regenerate")
|
||||
async def regenerate_api_token(request: Request, payload: dict[str, Any] = Body(...)):
|
||||
_require_login(request)
|
||||
token_id = str(payload.get("id", "")).strip()
|
||||
if not token_id:
|
||||
raise HTTPException(status_code=400, detail="缺少 Token ID")
|
||||
|
||||
try:
|
||||
raw_token, record = regenerate_token(token_id)
|
||||
return {
|
||||
"success": True,
|
||||
"message": "API Token 已重新生成",
|
||||
"token": raw_token,
|
||||
"record": record,
|
||||
}
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/api/tokens/delete")
|
||||
async def delete_api_token(request: Request, payload: dict[str, Any] = Body(...)):
|
||||
_require_login(request)
|
||||
token_id = str(payload.get("id", "")).strip()
|
||||
if not token_id:
|
||||
raise HTTPException(status_code=400, detail="缺少 Token ID")
|
||||
|
||||
delete_token(token_id)
|
||||
return {"success": True, "message": "API Token 已删除"}
|
||||
|
||||
|
||||
@router.post("/api/tokens/batch-delete")
|
||||
async def batch_delete_api_tokens(request: Request, payload: dict[str, Any] = Body(...)):
|
||||
_require_login(request)
|
||||
token_ids = payload.get("ids", [])
|
||||
if not isinstance(token_ids, list) or not token_ids:
|
||||
raise HTTPException(status_code=400, detail="请选择要删除的 Token")
|
||||
|
||||
deleted_count = delete_tokens([str(token_id) for token_id in token_ids])
|
||||
return {"success": True, "message": f"已删除 {deleted_count} 个 API Token", "deleted_count": deleted_count}
|
||||
|
||||
|
||||
@router.get("/api/docs-info")
|
||||
async def docs_info(request: Request):
|
||||
_require_login(request)
|
||||
return {
|
||||
"success": True,
|
||||
"docs_url": "/api/docs-ui",
|
||||
"openapi_url": "/api/openapi.json",
|
||||
"token_header": "Authorization: Bearer <token>",
|
||||
"alt_header": "X-API-Token: <token>",
|
||||
"note": "API 文档仅登录后可访问。",
|
||||
}
|
||||
@@ -7,8 +7,14 @@ from fastapi import APIRouter, Body, HTTPException, UploadFile, File
|
||||
from fastapi.responses import Response
|
||||
|
||||
from app import state
|
||||
from app.config import HistoryRetentionConfig, RJDataConfig, RemoteDataConfig
|
||||
from app.services.api_tokens import export_tokens, import_tokens
|
||||
from app.config import (
|
||||
HistoryRetentionConfig,
|
||||
MetrixConfig,
|
||||
RJDataConfig,
|
||||
RemoteDataConfig,
|
||||
SOURCE_TYPES,
|
||||
WAREHOUSE_TYPES,
|
||||
)
|
||||
|
||||
|
||||
router = APIRouter(tags=["config"])
|
||||
@@ -50,6 +56,30 @@ async def update_remote_config(config: dict[str, Any] = Body(...)):
|
||||
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/history-retention")
|
||||
async def update_history_retention(config: dict[str, Any] = Body(...)):
|
||||
state.reload_config()
|
||||
@@ -79,7 +109,6 @@ 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()
|
||||
config_data["ApiTokens"] = export_tokens()
|
||||
content = json.dumps(config_data, ensure_ascii=False, indent=2)
|
||||
return Response(
|
||||
content=content,
|
||||
@@ -100,8 +129,6 @@ async def upload_config(file: UploadFile = File(...)):
|
||||
|
||||
state.reload_config()
|
||||
_apply_config_data(data)
|
||||
if "ApiTokens" in data:
|
||||
import_tokens(data["ApiTokens"])
|
||||
state.config.save()
|
||||
return {"success": True, "message": "配置文件上传成功", "update": state.config.update}
|
||||
except json.JSONDecodeError as exc:
|
||||
@@ -113,6 +140,15 @@ async def upload_config(file: UploadFile = File(...)):
|
||||
|
||||
|
||||
def _apply_config_data(data: dict[str, Any]) -> None:
|
||||
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_data = data.get("Metrix")
|
||||
if isinstance(metrix_data, dict):
|
||||
state.config.metrix = MetrixConfig.from_dict(metrix_data)
|
||||
|
||||
mysql_data = data.get("MySQL_DBInfo")
|
||||
if isinstance(mysql_data, dict):
|
||||
for key in ("host", "port", "user", "passwd", "dbname"):
|
||||
|
||||
+51
-13
@@ -10,7 +10,9 @@ from starlette.background import BackgroundTask
|
||||
from app import state
|
||||
from app.config import CACHE_DIR
|
||||
from app.database import DatabaseManager
|
||||
from app.services.platform import make_client
|
||||
from app.utils.files import remove_file_safely
|
||||
from app.warehouse import make_warehouse
|
||||
|
||||
|
||||
router = APIRouter(tags=["database"])
|
||||
@@ -47,19 +49,20 @@ def _dataframe_from_table(db: DatabaseManager, table_name: str) -> pd.DataFrame:
|
||||
return pd.DataFrame(result["data"], columns=columns)
|
||||
|
||||
|
||||
def _db() -> DatabaseManager:
|
||||
return DatabaseManager(state.current_config())
|
||||
def _db():
|
||||
"""Direct MySQL DatabaseManager, or a Metrix-backed warehouse with the same interface."""
|
||||
return make_warehouse(state.current_config())
|
||||
|
||||
|
||||
@router.post("/api/database/test")
|
||||
async def test_database():
|
||||
def test_database():
|
||||
db = _db()
|
||||
success, message = db.test_connection()
|
||||
return {"success": success, "message": message}
|
||||
|
||||
|
||||
@router.get("/api/database/info")
|
||||
async def get_database_info():
|
||||
def get_database_info():
|
||||
db = _db()
|
||||
try:
|
||||
return {"success": True, **db.get_server_info()}
|
||||
@@ -69,7 +72,7 @@ async def get_database_info():
|
||||
|
||||
@router.get("/api/database/tables")
|
||||
@router.post("/api/database/tables")
|
||||
async def get_tables():
|
||||
def get_tables():
|
||||
db = _db()
|
||||
try:
|
||||
return {"tables": db.get_tables()}
|
||||
@@ -78,7 +81,7 @@ async def get_tables():
|
||||
|
||||
|
||||
@router.post("/api/database/table/info")
|
||||
async def get_table_info(table_name: str = Body(..., embed=True)):
|
||||
def get_table_info(table_name: str = Body(..., embed=True)):
|
||||
db = _db()
|
||||
try:
|
||||
return db.get_table_info(table_name)
|
||||
@@ -87,7 +90,7 @@ async def get_table_info(table_name: str = Body(..., embed=True)):
|
||||
|
||||
|
||||
@router.post("/api/database/table/data")
|
||||
async def query_table_data(
|
||||
def query_table_data(
|
||||
table_name: str = Body(..., embed=True),
|
||||
page: int = Body(1),
|
||||
page_size: int = Body(50),
|
||||
@@ -102,7 +105,7 @@ async def query_table_data(
|
||||
|
||||
|
||||
@router.post("/api/database/table/query")
|
||||
async def query_table_with_filter(
|
||||
def query_table_with_filter(
|
||||
table_name: str = Body(..., embed=True),
|
||||
page: int = Body(1),
|
||||
page_size: int = Body(50),
|
||||
@@ -125,7 +128,7 @@ async def query_table_with_filter(
|
||||
|
||||
|
||||
@router.post("/api/database/table/truncate")
|
||||
async def truncate_table(table_name: str = Body(..., embed=True)):
|
||||
def truncate_table(table_name: str = Body(..., embed=True)):
|
||||
db = _db()
|
||||
try:
|
||||
db.truncate_table(table_name)
|
||||
@@ -135,7 +138,7 @@ async def truncate_table(table_name: str = Body(..., embed=True)):
|
||||
|
||||
|
||||
@router.post("/api/database/table/drop")
|
||||
async def drop_table(table_name: str = Body(..., embed=True)):
|
||||
def drop_table(table_name: str = Body(..., embed=True)):
|
||||
db = _db()
|
||||
try:
|
||||
db.drop_table(table_name)
|
||||
@@ -145,7 +148,7 @@ async def drop_table(table_name: str = Body(..., embed=True)):
|
||||
|
||||
|
||||
@router.post("/api/database/table/drop-all")
|
||||
async def drop_all_tables():
|
||||
def drop_all_tables():
|
||||
db = _db()
|
||||
try:
|
||||
result = db.drop_all_tables()
|
||||
@@ -160,7 +163,7 @@ async def drop_all_tables():
|
||||
|
||||
|
||||
@router.post("/api/database/execute")
|
||||
async def execute_sql(sql: str = Body(..., embed=True)):
|
||||
def execute_sql(sql: str = Body(..., embed=True)):
|
||||
db = _db()
|
||||
try:
|
||||
success, result = db.execute_sql(sql)
|
||||
@@ -173,8 +176,10 @@ async def execute_sql(sql: str = Body(..., embed=True)):
|
||||
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")
|
||||
async def download_table(
|
||||
def download_table(
|
||||
table_name: Optional[str] = Body(None, embed=True),
|
||||
table_names: Optional[list[str]] = Body(None, embed=True),
|
||||
file_format: str = Body("csv", alias="format"),
|
||||
@@ -188,6 +193,10 @@ async def download_table(
|
||||
if file_format == "csv" and len(requested_tables) != 1:
|
||||
raise HTTPException(status_code=400, detail="CSV 每次只能导出一张表")
|
||||
|
||||
config = state.current_config()
|
||||
if config.warehouse_type == "metrix":
|
||||
return _download_via_metrix(config, requested_tables, file_format)
|
||||
|
||||
db = _db()
|
||||
try:
|
||||
available_tables = set(db.get_tables())
|
||||
@@ -230,3 +239,32 @@ async def download_table(
|
||||
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),
|
||||
)
|
||||
|
||||
+34
-17
@@ -1,3 +1,4 @@
|
||||
import time
|
||||
from datetime import date, datetime
|
||||
from pathlib import Path
|
||||
from threading import Thread
|
||||
@@ -14,6 +15,8 @@ from app.api.routers.task_runtime import (
|
||||
from app.config import AppConfig, CACHE_DIR, RemoteDataConfig
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
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
|
||||
|
||||
|
||||
@@ -22,12 +25,17 @@ router = APIRouter(tags=["remote"])
|
||||
|
||||
@router.post("/api/remote/test")
|
||||
async def test_remote_connection(config: dict[str, Any] | None = Body(None)):
|
||||
remote_config = RemoteDataConfig.from_dict(config) if config else state.current_config().remote_data
|
||||
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}"}
|
||||
return {"success": False, "message": f"连接失败: {exc}"}
|
||||
|
||||
|
||||
@router.post("/api/remote/start")
|
||||
@@ -123,28 +131,37 @@ def _run_remote_processing(
|
||||
) -> None:
|
||||
final_status = "failed"
|
||||
try:
|
||||
logger.info(
|
||||
f"开始远程下载,协议: {remote_config.protocol.upper()},"
|
||||
f"服务器: {remote_config.host}:{remote_config.port},目录: {remote_config.remote_dir}"
|
||||
)
|
||||
downloader = RemoteDataDownloader(remote_config, logger.info)
|
||||
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"下载完成,共 {download_result.file_count} 个文件,"
|
||||
f"{_format_bytes(download_result.total_bytes)}"
|
||||
)
|
||||
|
||||
if download_result.file_count == 0:
|
||||
raise RuntimeError("远程目录中未下载到任何文件")
|
||||
raise RuntimeError("源目录中未下载到任何文件")
|
||||
|
||||
state.history_manager.update(task_id, file_count=download_result.file_count)
|
||||
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")
|
||||
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)
|
||||
@@ -159,9 +176,9 @@ def _run_remote_processing(
|
||||
state.history_manager.update(
|
||||
task_id,
|
||||
status=status,
|
||||
elapsed_time=result.get("elapsed_time", 0),
|
||||
elapsed_time=elapsed,
|
||||
error=error,
|
||||
result_tables=["4G_结果表", "5G_结果表"],
|
||||
result_tables=RESULT_TABLES,
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
|
||||
@@ -84,11 +84,15 @@ def _run_script(task_id: str, logger: ProcessLogger, logs: list[str], app_config
|
||||
temp_work_dir: Path | None = None
|
||||
try:
|
||||
logger.info("开始执行 SQL 脚本...")
|
||||
temp_work_dir = CACHE_DIR / task_id
|
||||
temp_work_dir.mkdir(parents=True, exist_ok=True)
|
||||
if app_config.warehouse_type == "metrix":
|
||||
from app.services.pipeline import run_report_sql
|
||||
|
||||
processor = DataProcessor(app_config, temp_work_dir, logger)
|
||||
processor._execute_sql_script()
|
||||
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("SQL 脚本执行完成")
|
||||
set_task_stage(task_id, "completed", logs, status="completed")
|
||||
except Exception as exc:
|
||||
|
||||
@@ -168,19 +168,29 @@ def _task_finished(task_id: str) -> bool:
|
||||
|
||||
def _run_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
|
||||
try:
|
||||
logger.set_stage("license")
|
||||
log_license_check(logger, check_processing_allowed(work_dir))
|
||||
if app_config.warehouse_type == "metrix":
|
||||
import time
|
||||
|
||||
processor = DataProcessor(app_config, work_dir, logger)
|
||||
result = processor.process()
|
||||
status = "completed" if result.get("success") else "failed"
|
||||
error = result.get("error")
|
||||
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=result.get("elapsed_time", 0),
|
||||
elapsed_time=elapsed,
|
||||
error=error,
|
||||
result_tables=["4G_结果表", "5G_结果表"],
|
||||
result_tables=result_tables,
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
|
||||
@@ -123,12 +123,6 @@ def resolve_access_context(request: Request) -> AuthContext | None:
|
||||
if payload:
|
||||
return AuthContext(kind="jwt", payload=payload)
|
||||
|
||||
from app.services.api_tokens import verify_api_token
|
||||
|
||||
api_payload = verify_api_token(token)
|
||||
if api_payload:
|
||||
return AuthContext(kind="api_token", payload=api_payload)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
|
||||
+100
@@ -34,6 +34,77 @@ class MySQLConfig:
|
||||
dbname: str = "CapacityReport"
|
||||
|
||||
|
||||
# Source backends: direct FTP/SFTP, or Metrix storage platform.
|
||||
SOURCE_TYPES = ("ftp", "sftp", "metrix")
|
||||
# Warehouse backends: direct MySQL, or Metrix database platform.
|
||||
WAREHOUSE_TYPES = ("mysql", "metrix")
|
||||
|
||||
|
||||
@dataclass
|
||||
class MetrixConfig:
|
||||
"""Connection to a Metrix platform, used when source_type/warehouse_type is 'metrix'.
|
||||
|
||||
Metrix appears in the UI as two connection types ("存储平台"/"数据库平台") that share the
|
||||
same base_url + token; storage_id is the file source, database_conn_id + target_database
|
||||
are the warehouse. data_dir_to_table maps data sub-dirs to staging tables (Metrix mode only).
|
||||
"""
|
||||
base_url: str = "http://host.docker.internal:8000"
|
||||
token: str = ""
|
||||
storage_id: str = ""
|
||||
database_conn_id: str = ""
|
||||
target_database: str = ""
|
||||
recent_days: int = 7
|
||||
data_dir_to_table: Dict[str, str] = field(default_factory=lambda: {"4G": "4G_UD", "5G": "5G_UD"})
|
||||
|
||||
def normalized(self) -> "MetrixConfig":
|
||||
try:
|
||||
recent_days = max(int(self.recent_days), 1)
|
||||
except (TypeError, ValueError):
|
||||
recent_days = 7
|
||||
mapping = {
|
||||
str(k).strip(): str(v).strip()
|
||||
for k, v in (self.data_dir_to_table or {}).items()
|
||||
if str(k).strip() and str(v).strip()
|
||||
}
|
||||
return MetrixConfig(
|
||||
base_url=str(self.base_url or "").strip(),
|
||||
token=str(self.token or "").strip(),
|
||||
storage_id=str(self.storage_id or "").strip(),
|
||||
database_conn_id=str(self.database_conn_id or "").strip(),
|
||||
target_database=str(self.target_database or "").strip(),
|
||||
recent_days=recent_days,
|
||||
data_dir_to_table=mapping or {"4G": "4G_UD", "5G": "5G_UD"},
|
||||
)
|
||||
|
||||
def to_dict(self, include_token: bool = False) -> Dict[str, Any]:
|
||||
n = self.normalized()
|
||||
data = {
|
||||
"base_url": n.base_url,
|
||||
"storage_id": n.storage_id,
|
||||
"database_conn_id": n.database_conn_id,
|
||||
"target_database": n.target_database,
|
||||
"recent_days": n.recent_days,
|
||||
"data_dir_to_table": n.data_dir_to_table,
|
||||
}
|
||||
if include_token:
|
||||
data["token"] = n.token
|
||||
return data
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any] | None) -> "MetrixConfig":
|
||||
data = data or {}
|
||||
mapping = data.get("data_dir_to_table")
|
||||
return cls(
|
||||
base_url=str(data.get("base_url", "http://host.docker.internal:8000")),
|
||||
token=str(data.get("token", "")),
|
||||
storage_id=str(data.get("storage_id", "")),
|
||||
database_conn_id=str(data.get("database_conn_id", "")),
|
||||
target_database=str(data.get("target_database", "")),
|
||||
recent_days=data.get("recent_days", 7),
|
||||
data_dir_to_table=mapping if isinstance(mapping, dict) else {"4G": "4G_UD", "5G": "5G_UD"},
|
||||
).normalized()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AutoSchedulerConfig:
|
||||
enabled: bool = False
|
||||
@@ -250,10 +321,26 @@ class HistoryRetentionConfig:
|
||||
).normalized()
|
||||
|
||||
|
||||
def _normalize_source_type(value: Any, protocol: str = "sftp") -> str:
|
||||
text = str(value or "").strip().lower()
|
||||
if text in SOURCE_TYPES:
|
||||
return text
|
||||
# Back-compat: no explicit source type means direct remote, pick its protocol.
|
||||
return "ftp" if str(protocol).strip().lower() == "ftp" else "sftp"
|
||||
|
||||
|
||||
def _normalize_warehouse_type(value: Any) -> str:
|
||||
text = str(value or "").strip().lower()
|
||||
return text if text in WAREHOUSE_TYPES else "mysql"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AppConfig:
|
||||
update: str = ""
|
||||
source_type: str = "sftp"
|
||||
warehouse_type: str = "mysql"
|
||||
mysql: MySQLConfig = field(default_factory=MySQLConfig)
|
||||
metrix: MetrixConfig = field(default_factory=MetrixConfig)
|
||||
remote_data: RemoteDataConfig = field(default_factory=RemoteDataConfig)
|
||||
history_retention: HistoryRetentionConfig = field(default_factory=HistoryRetentionConfig)
|
||||
rj_data: RJDataConfig = field(default_factory=RJDataConfig)
|
||||
@@ -278,12 +365,16 @@ class AppConfig:
|
||||
dbname=mysql_data.get("dbname", "CapacityReport")
|
||||
)
|
||||
remote_config = RemoteDataConfig.from_dict(data.get("RemoteData"))
|
||||
metrix_config = MetrixConfig.from_dict(data.get("Metrix"))
|
||||
history_retention = HistoryRetentionConfig.from_dict(data.get("HistoryRetention"))
|
||||
rj_data = RJDataConfig.from_dict(data.get("RJData"))
|
||||
|
||||
return cls(
|
||||
update=data.get("Update", ""),
|
||||
source_type=_normalize_source_type(data.get("SourceType"), remote_config.protocol),
|
||||
warehouse_type=_normalize_warehouse_type(data.get("WarehouseType")),
|
||||
mysql=mysql_config,
|
||||
metrix=metrix_config,
|
||||
remote_data=remote_config,
|
||||
history_retention=history_retention,
|
||||
rj_data=rj_data,
|
||||
@@ -303,6 +394,8 @@ class AppConfig:
|
||||
"""转换为配置文件结构(包含敏感字段,用于保存和下载)"""
|
||||
return {
|
||||
"Update": self.update,
|
||||
"SourceType": self.source_type,
|
||||
"WarehouseType": self.warehouse_type,
|
||||
"MySQL_DBInfo": {
|
||||
"host": self.mysql.host,
|
||||
"port": self.mysql.port,
|
||||
@@ -310,6 +403,7 @@ class AppConfig:
|
||||
"passwd": self.mysql.passwd,
|
||||
"dbname": self.mysql.dbname
|
||||
},
|
||||
"Metrix": self.metrix.normalized().to_dict(include_token=True),
|
||||
"RemoteData": self.remote_data.normalized().to_dict(include_password=True),
|
||||
"HistoryRetention": self.history_retention.normalized().to_dict(),
|
||||
"RJData": self.rj_data.normalized().to_dict(),
|
||||
@@ -321,12 +415,15 @@ class AppConfig:
|
||||
"""转换为字典(用于返回给前端,隐藏密码)"""
|
||||
return {
|
||||
"update": self.update,
|
||||
"source_type": self.source_type,
|
||||
"warehouse_type": self.warehouse_type,
|
||||
"mysql": {
|
||||
"host": self.mysql.host,
|
||||
"port": self.mysql.port,
|
||||
"user": self.mysql.user,
|
||||
"dbname": self.mysql.dbname
|
||||
},
|
||||
"metrix": self.metrix.normalized().to_dict(),
|
||||
"remote_data": self.remote_data.normalized().to_dict(),
|
||||
"history_retention": self.history_retention.normalized().to_dict(),
|
||||
"rj_data": self.rj_data.normalized().to_dict(),
|
||||
@@ -338,6 +435,8 @@ class AppConfig:
|
||||
"""转换为完整字典(包含密码,用于编辑时回显)"""
|
||||
return {
|
||||
"update": self.update,
|
||||
"source_type": self.source_type,
|
||||
"warehouse_type": self.warehouse_type,
|
||||
"mysql": {
|
||||
"host": self.mysql.host,
|
||||
"port": self.mysql.port,
|
||||
@@ -345,6 +444,7 @@ class AppConfig:
|
||||
"passwd": self.mysql.passwd,
|
||||
"dbname": self.mysql.dbname
|
||||
},
|
||||
"metrix": self.metrix.normalized().to_dict(include_token=True),
|
||||
"remote_data": self.remote_data.normalized().to_dict(include_password=True),
|
||||
"history_retention": self.history_retention.normalized().to_dict(),
|
||||
"rj_data": self.rj_data.normalized().to_dict(),
|
||||
|
||||
+6
-435
@@ -1,16 +1,15 @@
|
||||
import argparse
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.openapi.utils import get_openapi
|
||||
from fastapi.responses import FileResponse, JSONResponse, RedirectResponse
|
||||
from fastapi.responses import FileResponse, JSONResponse
|
||||
|
||||
from app import state
|
||||
from app.api.routers import (
|
||||
api_tokens,
|
||||
auth,
|
||||
cache,
|
||||
config,
|
||||
@@ -23,319 +22,21 @@ from app.api.routers import (
|
||||
tasks,
|
||||
upload,
|
||||
)
|
||||
from app.auth import extract_access_token, resolve_access_context, resolve_login_context
|
||||
from app.auth import resolve_access_context, resolve_login_context
|
||||
from app.config import BASE_DIR
|
||||
from app.services.api_tokens import touch_token_usage
|
||||
from app.services.auto_scheduler import AutoScheduler
|
||||
|
||||
|
||||
APP_VERSION = "3.0.0"
|
||||
APP_HOST = "0.0.0.0"
|
||||
APP_PORT = 9081
|
||||
FRONTEND_DIST_DIR = BASE_DIR / "frontend" / "dist"
|
||||
# Code/frontend live in the image (/app); runtime state lives on the data volume (BASE_DIR=/data).
|
||||
FRONTEND_DIST_DIR = Path(os.environ.get("CAPAREPORT_FRONTEND_DIR") or (BASE_DIR / "frontend" / "dist"))
|
||||
LOGIN_ONLY_API_PREFIXES = (
|
||||
"/api/config",
|
||||
"/api/change-password",
|
||||
"/api/license",
|
||||
"/api/tokens",
|
||||
)
|
||||
LOGIN_ONLY_API_PATHS = {"/api/openapi.json", "/api/docs-ui", "/api/docs-info"}
|
||||
TAG_LABELS = {
|
||||
"auth": "认证",
|
||||
"upload": "数据处理",
|
||||
"remote": "远程数据",
|
||||
"tasks": "任务状态",
|
||||
"history": "处理历史",
|
||||
"database": "数据库",
|
||||
"script": "脚本",
|
||||
"config": "系统配置",
|
||||
"cache": "缓存",
|
||||
"license": "授权",
|
||||
"api-tokens": "API Token",
|
||||
"health": "健康检查",
|
||||
}
|
||||
OPENAPI_TAGS = [
|
||||
{"name": "认证", "description": "登录、修改密码等登录态接口。"},
|
||||
{"name": "数据处理", "description": "上传源数据并启动容量报表处理流程。"},
|
||||
{"name": "远程数据", "description": "测试 FTP/SFTP 连接,并从远程目录下载后自动处理。"},
|
||||
{"name": "任务状态", "description": "查看当前任务、处理进度和日志。"},
|
||||
{"name": "处理历史", "description": "查询、下载和清理历史处理记录及原始数据。"},
|
||||
{"name": "数据库", "description": "列出表、查询表、导出表和执行 SQL。API Token 可调用这些业务接口。"},
|
||||
{"name": "脚本", "description": "查看、保存和手动执行报表 SQL 脚本。"},
|
||||
{"name": "系统配置", "description": "数据库、远程数据源、过滤规则、字段映射等配置。仅登录态可调用。"},
|
||||
{"name": "缓存", "description": "查看服务端缓存占用。"},
|
||||
{"name": "授权", "description": "查看和延长程序授权有效期。仅登录态可调用。"},
|
||||
{"name": "API Token", "description": "生成、复制、启停、编辑、重生成和批量删除 API Token。仅登录态可调用。"},
|
||||
{"name": "健康检查", "description": "服务健康状态检查。"},
|
||||
]
|
||||
OPENAPI_OPERATION_DOCS = {
|
||||
("post", "/api/login"): {
|
||||
"summary": "登录系统",
|
||||
"description": "使用系统账号密码登录,成功后返回登录 JWT。",
|
||||
"example": {"username": "root", "password": "capacity"},
|
||||
},
|
||||
("post", "/api/change-password"): {
|
||||
"summary": "修改登录密码",
|
||||
"description": "修改当前登录用户密码,需要登录 JWT,不支持 API Token 调用。",
|
||||
"example": {"current_password": "capacity", "new_password": "new-password"},
|
||||
},
|
||||
("get", "/health"): {"summary": "健康检查", "description": "返回服务进程是否可用。"},
|
||||
("post", "/api/upload/create"): {
|
||||
"summary": "创建上传会话",
|
||||
"description": "创建一个待上传的数据处理会话,返回 session_id。通常用于分批上传文件。",
|
||||
},
|
||||
("post", "/api/upload"): {
|
||||
"summary": "上传源数据文件",
|
||||
"description": "上传 ZIP、CSV 或 Excel 数据文件。可传 session_id 追加到已有上传会话;不传则自动创建并锁定任务。",
|
||||
"request_description": "multipart/form-data,files 为一个或多个文件,session_id 可选。",
|
||||
},
|
||||
("post", "/api/upload/complete/{session_id}"): {
|
||||
"summary": "完成上传会话",
|
||||
"description": "标记上传会话文件数量,用于前端展示。",
|
||||
"parameters": {"session_id": "上传会话 ID。"},
|
||||
},
|
||||
("post", "/api/process/start"): {
|
||||
"summary": "开始处理已上传数据",
|
||||
"description": "对指定上传会话目录执行解压、入库和报表 SQL 脚本。",
|
||||
"example": {"task_id": "20260519_172457"},
|
||||
},
|
||||
("post", "/api/process/status"): {
|
||||
"summary": "查询处理任务状态",
|
||||
"description": "返回任务阶段、状态、日志和错误详情。",
|
||||
"example": {"task_id": "20260519_172457"},
|
||||
},
|
||||
("get", "/api/process/active"): {
|
||||
"summary": "查询当前活跃任务",
|
||||
"description": "返回当前是否有上传、远程下载、处理或脚本任务正在执行。",
|
||||
},
|
||||
("get", "/api/task/status"): {
|
||||
"summary": "查询全局任务锁",
|
||||
"description": "返回全局任务锁状态、任务 ID、阶段和最近日志。",
|
||||
},
|
||||
("post", "/api/task/lock"): {
|
||||
"summary": "锁定任务",
|
||||
"description": "内部接口:手动占用全局任务锁。",
|
||||
"example": {"task_id": "manual-task"},
|
||||
},
|
||||
("post", "/api/task/unlock"): {
|
||||
"summary": "释放任务锁",
|
||||
"description": "内部接口:释放全局任务锁。传 task_id 时只释放匹配的任务。",
|
||||
"example": {"task_id": "manual-task"},
|
||||
},
|
||||
("post", "/api/remote/test"): {
|
||||
"summary": "测试远程数据源",
|
||||
"description": "测试 FTP/SFTP 连接。请求体为空时使用系统设置中的远程数据源配置。",
|
||||
"example": {
|
||||
"protocol": "sftp",
|
||||
"host": "127.0.0.1",
|
||||
"port": 22,
|
||||
"user": "user",
|
||||
"passwd": "your-password",
|
||||
"remote_dir": "/CapacityReportData",
|
||||
"passive": True,
|
||||
"timeout": 30,
|
||||
"auto_delete_source": False,
|
||||
},
|
||||
},
|
||||
("post", "/api/remote/start"): {
|
||||
"summary": "远程下载并处理",
|
||||
"description": "从已配置的 FTP/SFTP 目录递归下载源数据,然后自动执行完整处理流程。",
|
||||
},
|
||||
("get", "/api/remote/scheduler/status"): {
|
||||
"summary": "查询远程自动调度状态",
|
||||
"description": "返回自动调度启用状态、目标周、就绪标识、下次检查时间和各远程目录的日期覆盖情况。",
|
||||
},
|
||||
("post", "/api/remote/scheduler/trigger"): {
|
||||
"summary": "手动触发自动调度检查",
|
||||
"description": "立即执行一次远程目录就绪检查;如果已存在就绪标识,会直接触发远程下载并处理。",
|
||||
},
|
||||
("post", "/api/history"): {
|
||||
"summary": "查询处理历史",
|
||||
"description": "按最近时间返回处理历史记录。",
|
||||
"example": {"limit": 50},
|
||||
},
|
||||
("post", "/api/history/detail"): {
|
||||
"summary": "查询历史详情",
|
||||
"description": "返回指定历史记录的基础信息和完整日志。",
|
||||
"example": {"record_id": "20260519_172457"},
|
||||
},
|
||||
("post", "/api/history/download"): {
|
||||
"summary": "下载历史原始数据",
|
||||
"description": "将历史记录对应工作目录压缩为 ZIP 后下载,下载响应完成后自动清理临时压缩包。",
|
||||
"example": {"record_id": "20260519_172457"},
|
||||
},
|
||||
("post", "/api/history/files"): {
|
||||
"summary": "浏览历史文件",
|
||||
"description": "列出指定历史记录工作目录下某个相对目录的文件和子目录,返回类型、大小和修改时间。",
|
||||
"example": {"record_id": "20260519_172457", "path": "4G/FDD"},
|
||||
},
|
||||
("post", "/api/history/file/download"): {
|
||||
"summary": "下载历史单个文件或目录",
|
||||
"description": "下载历史工作目录内的单个文件;如果目标是目录,则先压缩该目录后下载,并在响应完成后清理临时压缩包。",
|
||||
"example": {"record_id": "20260519_172457", "path": "4G/FDD/CapacityReportData4G_202605110000_202605120000.zip"},
|
||||
},
|
||||
("post", "/api/history/size"): {
|
||||
"summary": "查询历史目录大小",
|
||||
"description": "统计指定历史记录工作目录的文件数和占用空间。",
|
||||
"example": {"record_id": "20260519_172457"},
|
||||
},
|
||||
("post", "/api/history/delete"): {
|
||||
"summary": "删除历史记录",
|
||||
"description": "删除指定处理历史及其本地缓存数据。",
|
||||
"example": {"record_id": "20260519_172457"},
|
||||
},
|
||||
("post", "/api/history/clear"): {"summary": "清空处理历史", "description": "删除全部处理历史及其缓存数据。"},
|
||||
("get", "/api/database/info"): {
|
||||
"summary": "查询数据库信息",
|
||||
"description": "返回 MySQL 版本、LOAD DATA INFILE 可用性等诊断信息。",
|
||||
},
|
||||
("post", "/api/database/test"): {"summary": "测试数据库连接", "description": "测试当前数据库配置是否可连接。"},
|
||||
("get", "/api/database/tables"): {"summary": "列出所有数据表", "description": "返回当前数据库中的全部表名。"},
|
||||
("post", "/api/database/tables"): {"summary": "列出所有数据表", "description": "返回当前数据库中的全部表名。"},
|
||||
("post", "/api/database/table/info"): {
|
||||
"summary": "查询数据表结构",
|
||||
"description": "返回指定表的字段结构和行数。",
|
||||
"example": {"table_name": "4G_结果表"},
|
||||
},
|
||||
("post", "/api/database/table/data"): {
|
||||
"summary": "分页查询数据表",
|
||||
"description": "按页读取指定表数据,可指定排序字段和排序方向。",
|
||||
"example": {
|
||||
"table_name": "4G_结果表",
|
||||
"page": 1,
|
||||
"page_size": 50,
|
||||
"order_by": "日均流量(GB)",
|
||||
"order_dir": "DESC",
|
||||
},
|
||||
},
|
||||
("post", "/api/database/table/query"): {
|
||||
"summary": "按条件查询数据表",
|
||||
"description": "支持分页、排序和字段模糊查询。filters 的 key 为字段名,value 为模糊匹配值。",
|
||||
"example": {
|
||||
"table_name": "4G_结果表",
|
||||
"page": 1,
|
||||
"page_size": 50,
|
||||
"filters": {"小区名称": "广州"},
|
||||
"order_by": "日均流量(GB)",
|
||||
"order_dir": "DESC",
|
||||
},
|
||||
},
|
||||
("post", "/api/database/table/truncate"): {
|
||||
"summary": "清空数据表",
|
||||
"description": "保留表结构,删除指定表的全部数据。",
|
||||
"example": {"table_name": "4G_UD"},
|
||||
},
|
||||
("post", "/api/database/table/drop"): {
|
||||
"summary": "删除数据表",
|
||||
"description": "删除指定数据表。",
|
||||
"example": {"table_name": "4G_UD"},
|
||||
},
|
||||
("post", "/api/database/table/drop-all"): {"summary": "删除全部数据表", "description": "删除当前数据库中的全部表。"},
|
||||
("post", "/api/database/execute"): {
|
||||
"summary": "执行自定义 SQL",
|
||||
"description": "执行任意 SQL,包括 SELECT、UPDATE、INSERT、DROP 等。请仅在可信内网环境使用。",
|
||||
"example": {"sql": "SELECT * FROM `4G_结果表` LIMIT 10"},
|
||||
},
|
||||
("post", "/api/download"): {
|
||||
"summary": "导出数据表",
|
||||
"description": "导出 CSV 或 XLSX。CSV 每次只能导出一张表,XLSX 可选择多张表并按表名分 sheet。",
|
||||
"example": {"format": "xlsx", "table_names": ["4G_结果表", "5G_结果表"]},
|
||||
},
|
||||
("get", "/api/script/content"): {"summary": "读取 SQL 脚本", "description": "读取当前 ReportScript.sql 内容和修改时间。"},
|
||||
("post", "/api/script/save"): {
|
||||
"summary": "保存 SQL 脚本",
|
||||
"description": "覆盖保存 ReportScript.sql 内容。",
|
||||
"example": {"content": "SELECT 1;"},
|
||||
},
|
||||
("post", "/api/script/execute"): {"summary": "手动执行 SQL 脚本", "description": "直接执行当前 ReportScript.sql,并返回脚本任务 ID。"},
|
||||
("get", "/api/config"): {"summary": "读取基础配置", "description": "读取当前系统基础配置。仅登录态可访问。"},
|
||||
("get", "/api/config/full"): {"summary": "读取完整配置", "description": "读取数据库、远程数据源、历史保留、过滤规则和字段映射配置。"},
|
||||
("post", "/api/config/mysql"): {
|
||||
"summary": "保存数据库配置",
|
||||
"description": "更新 MySQL 连接配置。",
|
||||
"example": {"host": "capacity-mysql", "port": 3306, "user": "root", "passwd": "your-password", "dbname": "CapacityReport"},
|
||||
},
|
||||
("post", "/api/config/remote"): {
|
||||
"summary": "保存远程数据源配置",
|
||||
"description": "更新 FTP/SFTP 自动下载配置。",
|
||||
"example": {
|
||||
"enabled": True,
|
||||
"protocol": "sftp",
|
||||
"host": "127.0.0.1",
|
||||
"port": 22,
|
||||
"user": "user",
|
||||
"passwd": "your-password",
|
||||
"remote_dir": "/CapacityReportData",
|
||||
"passive": True,
|
||||
"timeout": 30,
|
||||
"auto_delete_source": False,
|
||||
"auto_scheduler": {
|
||||
"enabled": False,
|
||||
"check_interval_hours": 1,
|
||||
"expected_directories": ["4G/FDD", "4G/900", "5G/2.6", "5G/700"],
|
||||
"week_offset": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
("post", "/api/config/history-retention"): {
|
||||
"summary": "保存历史保留配置",
|
||||
"description": "设置处理历史是否自动清理,以及保留最近多少次记录。",
|
||||
"example": {"enabled": True, "keep_count": 20},
|
||||
},
|
||||
("post", "/api/config/sheet-filter"): {
|
||||
"summary": "保存 Sheet 过滤规则",
|
||||
"description": "设置需要跳过处理的 Sheet 关键字列表。",
|
||||
"example": ["指标(计数器)", "Template"],
|
||||
},
|
||||
("post", "/api/config/extract-fields"): {
|
||||
"summary": "保存字段映射配置",
|
||||
"description": "设置 Excel/CSV 源字段到数据库字段的映射规则。",
|
||||
"example": [{"Field": "日期时间", "Extract": ["开始时间"], "Type": "datetime"}],
|
||||
},
|
||||
("get", "/api/config/download"): {
|
||||
"summary": "下载配置文件",
|
||||
"description": "下载当前 Configure.json,并附带 ApiTokens,用于迁移或恢复 API Token 配置。",
|
||||
},
|
||||
("post", "/api/config/upload"): {
|
||||
"summary": "上传配置文件",
|
||||
"description": "上传并应用 Configure.json。文件中包含 ApiTokens 时会同步恢复 API Token。仅登录态可访问。",
|
||||
"request_description": "multipart/form-data,file 为 Configure.json 文件。",
|
||||
},
|
||||
("get", "/api/cache/size"): {"summary": "查询缓存大小", "description": "统计当前 cache 目录的大小、文件数和目录数。"},
|
||||
("get", "/api/license/status"): {"summary": "查询授权状态", "description": "返回当前授权到期日期和激活 key 标签。"},
|
||||
("post", "/api/license/activate"): {
|
||||
"summary": "激活授权延期",
|
||||
"description": "提交激活码,将授权到期日期延长 30 天。",
|
||||
"example": {"code": "sha256-value"},
|
||||
},
|
||||
("get", "/api/tokens"): {"summary": "列出 API Token", "description": "返回已创建 Token 列表,包含可复制的完整 Token。仅登录态可访问。"},
|
||||
("post", "/api/tokens/create"): {
|
||||
"summary": "生成 API Token",
|
||||
"description": "创建新的 API Token。完整 Token 会保存到本地,后续可在列表中重复复制。",
|
||||
"example": {"name": "外部系统接入", "permanent": True, "expires_at": None, "enabled": True},
|
||||
},
|
||||
("post", "/api/tokens/update"): {
|
||||
"summary": "编辑 API Token",
|
||||
"description": "修改 Token 名称、启停状态和有效期。",
|
||||
"example": {"id": "token-id", "name": "外部系统接入", "permanent": False, "expires_at": "2026-12-31", "enabled": True},
|
||||
},
|
||||
("post", "/api/tokens/regenerate"): {
|
||||
"summary": "重生成 API Token",
|
||||
"description": "重生成完整 Token,旧 Token 立即失效。",
|
||||
"example": {"id": "token-id"},
|
||||
},
|
||||
("post", "/api/tokens/delete"): {
|
||||
"summary": "删除 API Token",
|
||||
"description": "删除指定 Token。",
|
||||
"example": {"id": "token-id"},
|
||||
},
|
||||
("post", "/api/tokens/batch-delete"): {
|
||||
"summary": "批量删除 API Token",
|
||||
"description": "按 ID 批量删除 Token。",
|
||||
"example": {"ids": ["token-id-1", "token-id-2"]},
|
||||
},
|
||||
("get", "/api/docs-info"): {"summary": "查询 API 文档入口", "description": "返回 API 文档和 OpenAPI JSON 地址。仅登录态可访问。"},
|
||||
}
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -359,7 +60,6 @@ def create_app() -> FastAPI:
|
||||
openapi_url=None,
|
||||
lifespan=app_lifespan,
|
||||
)
|
||||
app.openapi = lambda: custom_openapi(app) # type: ignore[method-assign]
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
@@ -394,18 +94,12 @@ async def auth_middleware(request: Request, call_next):
|
||||
return JSONResponse(status_code=401, content={"detail": "未授权,请提供有效的 Token"})
|
||||
|
||||
request.state.auth_context = access_context
|
||||
if access_context.kind == "api_token":
|
||||
access_token = extract_access_token(request)
|
||||
if access_token:
|
||||
client_host = request.client.host if request.client else None
|
||||
touch_token_usage(access_token, client_host)
|
||||
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def register_routes(app: FastAPI) -> None:
|
||||
routers = [
|
||||
api_tokens.router,
|
||||
auth.router,
|
||||
health.router,
|
||||
upload.router,
|
||||
@@ -423,18 +117,6 @@ def register_routes(app: FastAPI) -> None:
|
||||
|
||||
|
||||
def register_frontend(app: FastAPI) -> None:
|
||||
@app.get("/api/openapi.json", include_in_schema=False)
|
||||
async def serve_openapi(request: Request):
|
||||
if resolve_login_context(request) is None:
|
||||
return JSONResponse(status_code=401, content={"detail": "未登录或登录已过期"})
|
||||
return JSONResponse(app.openapi())
|
||||
|
||||
@app.get("/api/docs-ui", include_in_schema=False)
|
||||
async def serve_docs_ui(request: Request):
|
||||
if resolve_login_context(request) is None:
|
||||
return JSONResponse(status_code=401, content={"detail": "未登录或登录已过期"})
|
||||
return RedirectResponse(url="/api-docs", status_code=302)
|
||||
|
||||
@app.get("/", include_in_schema=False)
|
||||
@app.get("/{path:path}", include_in_schema=False)
|
||||
async def serve_frontend(path: str = ""):
|
||||
@@ -475,118 +157,7 @@ def _safe_file(root: Path, path: str) -> Path | None:
|
||||
|
||||
|
||||
def _is_login_only_api(path: str) -> bool:
|
||||
return path in LOGIN_ONLY_API_PATHS or path.startswith(LOGIN_ONLY_API_PREFIXES)
|
||||
|
||||
|
||||
def custom_openapi(app: FastAPI) -> dict:
|
||||
if app.openapi_schema:
|
||||
return app.openapi_schema
|
||||
|
||||
schema = get_openapi(
|
||||
title="CapacityReport API",
|
||||
version=app.version,
|
||||
description=(
|
||||
"容量报表数据处理系统接口文档。业务接口支持登录 JWT 或 API Token;"
|
||||
"系统配置、授权、Token 管理和文档本身仅支持登录态访问。"
|
||||
),
|
||||
routes=app.routes,
|
||||
)
|
||||
schema["tags"] = OPENAPI_TAGS
|
||||
components = schema.setdefault("components", {})
|
||||
security_schemes = components.setdefault("securitySchemes", {})
|
||||
security_schemes["BearerAuth"] = {
|
||||
"type": "http",
|
||||
"scheme": "bearer",
|
||||
"bearerFormat": "JWT",
|
||||
"description": "登录 JWT 或 API Token,均可通过 Authorization: Bearer <token> 传递;API Token 也支持 X-API-Token: <token>。",
|
||||
}
|
||||
security_schemes["ApiTokenHeader"] = {
|
||||
"type": "apiKey",
|
||||
"in": "header",
|
||||
"name": "X-API-Token",
|
||||
"description": "API Token 也可以通过 X-API-Token 请求头传递。",
|
||||
}
|
||||
|
||||
for path, methods in schema.get("paths", {}).items():
|
||||
for method, operation in methods.items():
|
||||
if not isinstance(operation, dict):
|
||||
continue
|
||||
|
||||
operation["tags"] = [TAG_LABELS.get(tag, tag) for tag in operation.get("tags", [])]
|
||||
operation["operationId"] = _make_operation_id(method, path)
|
||||
|
||||
if path.startswith("/api/") and path not in {"/api/login"}:
|
||||
operation["security"] = (
|
||||
[{"BearerAuth": []}]
|
||||
if _is_login_only_api(path)
|
||||
else [{"BearerAuth": []}, {"ApiTokenHeader": []}]
|
||||
)
|
||||
|
||||
_apply_operation_doc(operation, OPENAPI_OPERATION_DOCS.get((method.lower(), path), {}))
|
||||
|
||||
app.openapi_schema = schema
|
||||
return schema
|
||||
|
||||
|
||||
def _make_operation_id(method: str, path: str) -> str:
|
||||
normalized_path = (
|
||||
path.strip("/")
|
||||
.replace("/", "_")
|
||||
.replace("-", "_")
|
||||
.replace("{", "")
|
||||
.replace("}", "")
|
||||
)
|
||||
return f"{method.lower()}_{normalized_path or 'root'}"
|
||||
|
||||
|
||||
def _apply_operation_doc(operation: dict, doc: dict) -> None:
|
||||
if not doc:
|
||||
return
|
||||
|
||||
for key in ("summary", "description"):
|
||||
value = doc.get(key)
|
||||
if value:
|
||||
operation[key] = value
|
||||
|
||||
request_description = doc.get("request_description")
|
||||
if request_description and isinstance(operation.get("requestBody"), dict):
|
||||
operation["requestBody"]["description"] = request_description
|
||||
|
||||
if "example" in doc:
|
||||
_set_request_example(operation, doc["example"])
|
||||
|
||||
parameter_descriptions = doc.get("parameters")
|
||||
if isinstance(parameter_descriptions, dict):
|
||||
_set_parameter_descriptions(operation, parameter_descriptions)
|
||||
|
||||
|
||||
def _set_request_example(operation: dict, example: object) -> None:
|
||||
request_body = operation.get("requestBody")
|
||||
if not isinstance(request_body, dict):
|
||||
return
|
||||
|
||||
content = request_body.get("content")
|
||||
if not isinstance(content, dict):
|
||||
return
|
||||
|
||||
media = content.get("application/json")
|
||||
if not isinstance(media, dict):
|
||||
media = next((value for value in content.values() if isinstance(value, dict)), None)
|
||||
if isinstance(media, dict):
|
||||
media["example"] = example
|
||||
|
||||
|
||||
def _set_parameter_descriptions(operation: dict, descriptions: dict[str, str]) -> None:
|
||||
parameters = operation.get("parameters")
|
||||
if not isinstance(parameters, list):
|
||||
return
|
||||
|
||||
for parameter in parameters:
|
||||
if not isinstance(parameter, dict):
|
||||
continue
|
||||
name = parameter.get("name")
|
||||
if isinstance(name, str) and name in descriptions:
|
||||
parameter["description"] = descriptions[name]
|
||||
return path.startswith(LOGIN_ONLY_API_PREFIXES)
|
||||
|
||||
|
||||
app = create_app()
|
||||
|
||||
@@ -1,329 +0,0 @@
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import secrets
|
||||
from dataclasses import asdict, dataclass
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
from threading import RLock
|
||||
from typing import Any
|
||||
|
||||
from app.config import BASE_DIR
|
||||
|
||||
|
||||
API_TOKENS_PATH = BASE_DIR / "api_tokens.json"
|
||||
API_TOKEN_PREFIX = "cap_"
|
||||
API_TOKEN_SECRET = "CapaReportApiTokenSecret2026"
|
||||
_STORE_LOCK = RLock()
|
||||
|
||||
|
||||
@dataclass
|
||||
class ApiTokenRecord:
|
||||
id: str
|
||||
name: str
|
||||
token_hash: str
|
||||
prefix: str
|
||||
suffix: str
|
||||
created_at: str
|
||||
expires_at: str | None
|
||||
enabled: bool
|
||||
last_used_at: str | None = None
|
||||
last_used_from: str | None = None
|
||||
token: str | None = None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def ensure_store() -> None:
|
||||
if API_TOKENS_PATH.exists():
|
||||
return
|
||||
with _STORE_LOCK:
|
||||
if API_TOKENS_PATH.exists():
|
||||
return
|
||||
API_TOKENS_PATH.write_text(json.dumps({"tokens": []}, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
|
||||
|
||||
def list_tokens() -> list[dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
return [record_to_public_dict(record) for record in _load_records()]
|
||||
|
||||
|
||||
def export_tokens() -> list[dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
return [record.to_dict() for record in _load_records()]
|
||||
|
||||
|
||||
def import_tokens(items: Any) -> int:
|
||||
if not isinstance(items, list):
|
||||
return 0
|
||||
|
||||
records: list[ApiTokenRecord] = []
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
record = _record_from_dict(item)
|
||||
if record.token_hash:
|
||||
records.append(record)
|
||||
|
||||
with _STORE_LOCK:
|
||||
_save_records(records)
|
||||
return len(records)
|
||||
|
||||
|
||||
def create_token(
|
||||
name: str,
|
||||
expires_in_days: int | None = None,
|
||||
enabled: bool = True,
|
||||
expires_at: str | None = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
raw_token = generate_raw_token()
|
||||
resolved_expires_at = normalize_expiration(expires_at)
|
||||
if resolved_expires_at is None and expires_in_days is not None:
|
||||
resolved_expires_at = expires_at_from_days(expires_in_days)
|
||||
|
||||
record = ApiTokenRecord(
|
||||
id=secrets.token_hex(8),
|
||||
name=name.strip() or "未命名 Token",
|
||||
token_hash=hash_token(raw_token),
|
||||
prefix=raw_token[:12],
|
||||
suffix=raw_token[-12:],
|
||||
created_at=utc_now(),
|
||||
expires_at=resolved_expires_at,
|
||||
enabled=bool(enabled),
|
||||
token=raw_token,
|
||||
)
|
||||
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
records.append(record)
|
||||
_save_records(records)
|
||||
|
||||
return raw_token, record_to_public_dict(record)
|
||||
|
||||
|
||||
def update_token(token_id: str, **changes: Any) -> dict[str, Any]:
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
for index, record in enumerate(records):
|
||||
if record.id != token_id:
|
||||
continue
|
||||
|
||||
if "name" in changes and isinstance(changes["name"], str):
|
||||
record.name = changes["name"].strip() or record.name
|
||||
if "enabled" in changes:
|
||||
record.enabled = bool(changes["enabled"])
|
||||
if "expires_at" in changes:
|
||||
record.expires_at = normalize_expiration(changes["expires_at"])
|
||||
|
||||
records[index] = record
|
||||
_save_records(records)
|
||||
return record_to_public_dict(record)
|
||||
|
||||
raise KeyError(f"Token not found: {token_id}")
|
||||
|
||||
|
||||
def delete_token(token_id: str) -> None:
|
||||
with _STORE_LOCK:
|
||||
records = [record for record in _load_records() if record.id != token_id]
|
||||
_save_records(records)
|
||||
|
||||
|
||||
def delete_tokens(token_ids: list[str]) -> int:
|
||||
token_id_set = {str(token_id).strip() for token_id in token_ids if str(token_id).strip()}
|
||||
if not token_id_set:
|
||||
return 0
|
||||
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
kept_records = [record for record in records if record.id not in token_id_set]
|
||||
_save_records(kept_records)
|
||||
return len(records) - len(kept_records)
|
||||
|
||||
|
||||
def regenerate_token(token_id: str) -> tuple[str, dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
for index, record in enumerate(records):
|
||||
if record.id != token_id:
|
||||
continue
|
||||
|
||||
raw_token = generate_raw_token()
|
||||
record.token_hash = hash_token(raw_token)
|
||||
record.prefix = raw_token[:12]
|
||||
record.suffix = raw_token[-12:]
|
||||
record.token = raw_token
|
||||
record.created_at = utc_now()
|
||||
record.last_used_at = None
|
||||
record.last_used_from = None
|
||||
records[index] = record
|
||||
_save_records(records)
|
||||
return raw_token, record_to_public_dict(record)
|
||||
|
||||
raise KeyError(f"Token not found: {token_id}")
|
||||
|
||||
|
||||
def verify_api_token(raw_token: str) -> dict[str, Any] | None:
|
||||
ensure_store()
|
||||
token_hash = hash_token(raw_token)
|
||||
now = datetime.now(timezone.utc)
|
||||
with _STORE_LOCK:
|
||||
for record in _load_records():
|
||||
if not record.enabled or record.token_hash != token_hash:
|
||||
continue
|
||||
|
||||
expires_at = parse_datetime(record.expires_at)
|
||||
if expires_at and expires_at < now:
|
||||
continue
|
||||
|
||||
return record_to_context(record)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def touch_token_usage(raw_token: str, source: str | None = None) -> None:
|
||||
token_hash = hash_token(raw_token)
|
||||
now = utc_now()
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
updated = False
|
||||
for index, record in enumerate(records):
|
||||
if record.token_hash != token_hash:
|
||||
continue
|
||||
record.last_used_at = now
|
||||
record.last_used_from = source or record.last_used_from
|
||||
records[index] = record
|
||||
updated = True
|
||||
break
|
||||
if updated:
|
||||
_save_records(records)
|
||||
|
||||
|
||||
def generate_raw_token() -> str:
|
||||
return API_TOKEN_PREFIX + secrets.token_urlsafe(36)
|
||||
|
||||
|
||||
def hash_token(raw_token: str) -> str:
|
||||
digest = hmac.new(API_TOKEN_SECRET.encode(), raw_token.encode(), hashlib.sha256).digest()
|
||||
return base64.urlsafe_b64encode(digest).decode().rstrip("=")
|
||||
|
||||
|
||||
def normalize_expiration(value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
trimmed = str(value).strip()
|
||||
if not trimmed:
|
||||
return None
|
||||
parsed = parse_datetime(trimmed)
|
||||
if parsed is None:
|
||||
raise ValueError("Invalid expiration date")
|
||||
return parsed.isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def expires_at_from_days(days: int) -> str:
|
||||
safe_days = max(int(days), 1)
|
||||
expires_at = datetime.now(timezone.utc) + timedelta(days=safe_days)
|
||||
return expires_at.isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def parse_datetime(value: str | None) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value)
|
||||
except ValueError:
|
||||
try:
|
||||
parsed_date = date.fromisoformat(value)
|
||||
except ValueError:
|
||||
return None
|
||||
parsed = datetime.combine(parsed_date, time.max)
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=timezone.utc)
|
||||
return parsed.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def utc_now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def record_to_public_dict(record: ApiTokenRecord, include_hash: bool = False) -> dict[str, Any]:
|
||||
expires_at = parse_datetime(record.expires_at)
|
||||
data = {
|
||||
"id": record.id,
|
||||
"name": record.name,
|
||||
"prefix": record.prefix,
|
||||
"suffix": record.suffix,
|
||||
"created_at": record.created_at,
|
||||
"expires_at": record.expires_at,
|
||||
"enabled": record.enabled,
|
||||
"last_used_at": record.last_used_at,
|
||||
"last_used_from": record.last_used_from,
|
||||
"expired": bool(expires_at and expires_at < datetime.now(timezone.utc)),
|
||||
"token": record.token,
|
||||
"token_available": bool(record.token),
|
||||
}
|
||||
if include_hash:
|
||||
data["token_hash"] = record.token_hash
|
||||
return data
|
||||
|
||||
|
||||
def record_to_context(record: ApiTokenRecord) -> dict[str, Any]:
|
||||
return {
|
||||
"token_id": record.id,
|
||||
"name": record.name,
|
||||
"created_at": record.created_at,
|
||||
"expires_at": record.expires_at,
|
||||
"token_type": "api_token",
|
||||
}
|
||||
|
||||
|
||||
def _load_records() -> list[ApiTokenRecord]:
|
||||
ensure_store()
|
||||
try:
|
||||
raw = json.loads(API_TOKENS_PATH.read_text(encoding="utf-8"))
|
||||
except json.JSONDecodeError:
|
||||
raw = {"tokens": []}
|
||||
|
||||
tokens = raw.get("tokens", []) if isinstance(raw, dict) else []
|
||||
records: list[ApiTokenRecord] = []
|
||||
for item in tokens:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
record = _record_from_dict(item)
|
||||
if record.token_hash:
|
||||
records.append(record)
|
||||
|
||||
records.sort(key=lambda record: record.created_at, reverse=True)
|
||||
return records
|
||||
|
||||
|
||||
def _record_from_dict(item: dict[str, Any]) -> ApiTokenRecord:
|
||||
raw_token = str(item.get("token") or "").strip() or None
|
||||
token_hash = str(item.get("token_hash") or "").strip()
|
||||
if raw_token and not token_hash:
|
||||
token_hash = hash_token(raw_token)
|
||||
|
||||
prefix = str(item.get("prefix") or "")
|
||||
suffix = str(item.get("suffix") or "")
|
||||
if raw_token:
|
||||
prefix = prefix or raw_token[:12]
|
||||
suffix = suffix or raw_token[-12:]
|
||||
|
||||
return ApiTokenRecord(
|
||||
id=str(item.get("id", "")) or secrets.token_hex(8),
|
||||
name=str(item.get("name", "未命名 Token")),
|
||||
token_hash=token_hash,
|
||||
prefix=prefix,
|
||||
suffix=suffix,
|
||||
created_at=str(item.get("created_at", utc_now())),
|
||||
expires_at=item.get("expires_at"),
|
||||
enabled=bool(item.get("enabled", True)),
|
||||
last_used_at=item.get("last_used_at"),
|
||||
last_used_from=item.get("last_used_from"),
|
||||
token=raw_token,
|
||||
)
|
||||
|
||||
|
||||
def _save_records(records: list[ApiTokenRecord]) -> None:
|
||||
payload = {"tokens": [record.to_dict() for record in records]}
|
||||
API_TOKENS_PATH.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
@@ -11,6 +11,7 @@ from fastapi import HTTPException
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR, AutoSchedulerConfig
|
||||
from app.services.platform import make_source_downloader
|
||||
from app.services.remote_download import RemoteDataDownloader, RemoteFileInfo
|
||||
from app.utils.file_dates import parse_file_date_range, required_week_days
|
||||
|
||||
@@ -186,7 +187,7 @@ class AutoScheduler:
|
||||
return self._trigger_processing(target_dates, ready_flag, manual)
|
||||
|
||||
target_days = required_week_days(scheduler.week_offset)
|
||||
downloader = RemoteDataDownloader(remote_config)
|
||||
downloader = make_source_downloader(app_config)
|
||||
rj_config = app_config.rj_data.normalized()
|
||||
rj_directories = set(rj_config.weekly_directories) if rj_config.enabled else set()
|
||||
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
"""纯数据处理流水线(容器版,无数据库代码)。
|
||||
|
||||
输入:包含已下载周 ZIP/CSV/Excel 的工作目录。
|
||||
输出:每张暂存表一个规范化 CSV({表名: csv 路径});由平台数据库导入 API 建表入库,
|
||||
再由平台 run-script 跑报表 SQL 生成结果表。这里不含任何数据库/LOAD DATA 代码。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
import zipfile
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Callable
|
||||
|
||||
import chardet
|
||||
import pandas as pd
|
||||
|
||||
from app.utils.file_dates import select_recent_items_by_directory
|
||||
|
||||
ZERO_TEXTS = {"", "-", "--", "—", "–", "NA", "N/A", "NULL", "NONE", "NAN", "\\N"}
|
||||
DATETIME_FORMATS = [
|
||||
"ISO8601",
|
||||
"%Y-%m-%d %H:%M:%S",
|
||||
"%Y-%m-%d %H:%M",
|
||||
"%Y/%m/%d %H:%M:%S",
|
||||
"%Y/%m/%d %H:%M",
|
||||
"%Y-%m-%d",
|
||||
"%Y/%m/%d",
|
||||
"%Y年%m月%d日 %H:%M:%S",
|
||||
"%Y年%m月%d日",
|
||||
"%Y%m%d%H%M%S",
|
||||
"%Y%m%d",
|
||||
]
|
||||
MAX_WORKERS = 8
|
||||
|
||||
|
||||
class CsvProcessor:
|
||||
def __init__(self, work_dir: Path, config: dict, log: Callable[[str], None]):
|
||||
self.work_dir = Path(work_dir)
|
||||
self.config = config
|
||||
self.log = log
|
||||
self.recent_days = int(config.get("recent_days", 7))
|
||||
self.sheet_filter = set(config.get("sheet_filter", []))
|
||||
self.data_dir_to_table = {k.upper(): v for k, v in (config.get("data_dir_to_table") or {}).items()}
|
||||
self.field_map, self.type_map = _build_global_map(config.get("extract_fields", []))
|
||||
rj = config.get("rj") or {}
|
||||
self.rj_enabled = bool(rj.get("enabled"))
|
||||
self.rj_weekly_dirs = list(rj.get("weekly_directories") or [])
|
||||
self.rj_dir_to_table = dict(rj.get("dir_to_table") or {})
|
||||
self.rj_maps = _build_rj_maps(rj.get("table_field_mappings") or {})
|
||||
self.out_dir = self.work_dir / ".out"
|
||||
|
||||
def process(self) -> dict[str, Path]:
|
||||
self._unzip_files()
|
||||
self._excel_to_csv()
|
||||
return self._build_table_csvs()
|
||||
|
||||
# --- step 1: unzip ---------------------------------------------------
|
||||
def _unzip_files(self) -> None:
|
||||
zips = self._filter_recent(list(self.work_dir.rglob("*.zip")), "ZIP")
|
||||
self.log(f"解压 ZIP: {len(zips)} 个")
|
||||
for zip_file in zips:
|
||||
try:
|
||||
_extract_zip(zip_file, self.log)
|
||||
except Exception as exc: # noqa: BLE001 - keep going on a bad archive
|
||||
self.log(f"[WARN] 解压失败 {zip_file.name}: {exc}")
|
||||
|
||||
# --- step 2: excel -> csv -------------------------------------------
|
||||
def _excel_to_csv(self) -> None:
|
||||
excels = self._filter_recent(list(self._scan(self.work_dir, (".xlsx", ".xls"))), "Excel")
|
||||
if not excels:
|
||||
return
|
||||
self.log(f"Excel 转 CSV: {len(excels)} 个文件")
|
||||
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool:
|
||||
futures = {pool.submit(self._one_excel, f): f for f in excels}
|
||||
for future in as_completed(futures):
|
||||
excel = futures[future]
|
||||
try:
|
||||
future.result()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self.log(f"[WARN] Excel 处理失败 {excel.name}: {exc}")
|
||||
|
||||
def _one_excel(self, excel_file: Path) -> None:
|
||||
xl = pd.ExcelFile(excel_file, engine="openpyxl")
|
||||
try:
|
||||
for sheet in xl.sheet_names:
|
||||
if sheet in self.sheet_filter:
|
||||
continue
|
||||
out = excel_file.parent / f"{excel_file.stem}_{sheet}.csv"
|
||||
xl.parse(sheet).to_csv(out, index=False, encoding="utf-8")
|
||||
finally:
|
||||
xl.close()
|
||||
|
||||
# --- step 3: build one normalized CSV per staging table -------------
|
||||
def _build_table_csvs(self) -> dict[str, Path]:
|
||||
data_dirs = self._find_data_dirs()
|
||||
if not data_dirs:
|
||||
self.log("[WARN] 未找到任何数据目录")
|
||||
return {}
|
||||
self.out_dir.mkdir(parents=True, exist_ok=True)
|
||||
result: dict[str, Path] = {}
|
||||
for table, directory in data_dirs.items():
|
||||
csv_files = self._filter_recent(list(self._scan(directory, (".csv",))), "CSV", root=directory)
|
||||
if not csv_files:
|
||||
continue
|
||||
field_map, type_map = self._maps_for_table(table)
|
||||
out_path = self.out_dir / f"{table}.csv"
|
||||
rows = self._write_table_csv(table, csv_files, field_map, type_map, out_path)
|
||||
if rows > 0:
|
||||
result[table] = out_path
|
||||
self.log(f"暂存表 {table}: {rows} 行 -> {out_path.name}")
|
||||
return result
|
||||
|
||||
def _write_table_csv(self, table, csv_files, field_map, type_map, out_path: Path) -> int:
|
||||
# First pass: union of target columns across this table's CSV files.
|
||||
union: list[str] = []
|
||||
seen: set[str] = set()
|
||||
frames: list[tuple[Path, list[str]]] = []
|
||||
for csv_file in csv_files:
|
||||
headers = _read_headers(csv_file)
|
||||
targets = _ordered_targets(headers, field_map)
|
||||
if not targets:
|
||||
continue
|
||||
frames.append((csv_file, headers))
|
||||
for target in targets:
|
||||
if target not in seen:
|
||||
seen.add(target)
|
||||
union.append(target)
|
||||
if not union:
|
||||
return 0
|
||||
|
||||
total = 0
|
||||
header_written = False
|
||||
for csv_file, _headers in frames:
|
||||
df = self._normalize(csv_file, field_map, type_map, union)
|
||||
if df is None or df.empty:
|
||||
continue
|
||||
df.to_csv(out_path, index=False, header=not header_written, mode="w" if not header_written else "a", encoding="utf-8")
|
||||
header_written = True
|
||||
total += len(df)
|
||||
return total
|
||||
|
||||
def _normalize(self, csv_file: Path, field_map, type_map, union: list[str]):
|
||||
try:
|
||||
df = pd.read_csv(
|
||||
csv_file,
|
||||
encoding=_detect_encoding(csv_file),
|
||||
dtype=str,
|
||||
na_values=[""],
|
||||
keep_default_na=False,
|
||||
low_memory=True,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self.log(f"[WARN] 读取 CSV 失败 {csv_file.name}: {exc}")
|
||||
return None
|
||||
|
||||
col_map: dict[str, str] = {}
|
||||
mapped: set[str] = set()
|
||||
for col in df.columns:
|
||||
target = field_map.get(col)
|
||||
if target and target not in mapped:
|
||||
col_map[col] = target
|
||||
mapped.add(target)
|
||||
if not col_map:
|
||||
return None
|
||||
|
||||
out = df[list(col_map.keys())].copy()
|
||||
out.columns = list(col_map.values())
|
||||
out = out.fillna("")
|
||||
for col in out.columns:
|
||||
col_type = type_map.get(col, "string")
|
||||
if col_type == "datetime":
|
||||
out[col] = _convert_datetime(out[col])
|
||||
elif col_type == "int":
|
||||
out[col] = _convert_int(out[col])
|
||||
elif col_type == "float":
|
||||
out[col] = _convert_float(out[col])
|
||||
else:
|
||||
out[col] = out[col].astype("string").str.replace("%", "", regex=False).str.slice(0, 255)
|
||||
# Reindex to the shared union columns; fill missing per type so numeric
|
||||
# staging columns never carry '' (the report SQL re-types them later).
|
||||
for col in union:
|
||||
if col not in out.columns:
|
||||
out[col] = "0" if type_map.get(col) in ("int", "float") else ""
|
||||
return out[union]
|
||||
|
||||
# --- directory detection --------------------------------------------
|
||||
def _find_data_dirs(self) -> dict[str, Path]:
|
||||
data_dirs: dict[str, Path] = {}
|
||||
target_names = set(self.data_dir_to_table.keys())
|
||||
for sub in self.work_dir.rglob("*"):
|
||||
if sub.is_dir() and sub.name.upper() in target_names:
|
||||
table = self.data_dir_to_table[sub.name.upper()]
|
||||
data_dirs.setdefault(table, sub)
|
||||
if self.rj_enabled:
|
||||
self._find_rj_dirs(data_dirs)
|
||||
return data_dirs
|
||||
|
||||
def _find_rj_dirs(self, data_dirs: dict[str, Path]) -> None:
|
||||
for weekly in self.rj_weekly_dirs:
|
||||
path = self.work_dir / weekly
|
||||
if path.exists() and path.is_dir() and path.name in self.rj_dir_to_table:
|
||||
data_dirs.setdefault(self.rj_dir_to_table[path.name], path)
|
||||
if not any(table in data_dirs for table in self.rj_dir_to_table.values()):
|
||||
for sub in self.work_dir.rglob("*"):
|
||||
if sub.is_dir() and sub.name in self.rj_dir_to_table:
|
||||
data_dirs.setdefault(self.rj_dir_to_table[sub.name], sub)
|
||||
|
||||
def _maps_for_table(self, table: str):
|
||||
if table in self.rj_maps:
|
||||
return self.rj_maps[table]
|
||||
return self.field_map, self.type_map
|
||||
|
||||
# --- helpers ---------------------------------------------------------
|
||||
def _scan(self, directory: Path, extensions: tuple[str, ...]):
|
||||
for ext in extensions:
|
||||
yield from directory.rglob(f"*{ext}")
|
||||
|
||||
def _filter_recent(self, files: list[Path], label: str, root: Path | None = None) -> list[Path]:
|
||||
if not files:
|
||||
return files
|
||||
base = (root or self.work_dir).resolve()
|
||||
|
||||
def parent_key(file_path: Path) -> str:
|
||||
try:
|
||||
parent = file_path.parent.resolve().relative_to(base)
|
||||
except ValueError:
|
||||
parent = file_path.parent
|
||||
text = str(parent).replace("\\", "/")
|
||||
return "" if text == "." else text
|
||||
|
||||
selected, summaries = select_recent_items_by_directory(
|
||||
files,
|
||||
parent_key=parent_key,
|
||||
name_key=lambda f: f.name,
|
||||
days=self.recent_days,
|
||||
)
|
||||
for summary in summaries:
|
||||
if summary.skipped_count and summary.start_date and summary.max_date:
|
||||
self.log(
|
||||
f"{label} {summary.directory or '.'}: 取 {summary.start_date}~{summary.max_date} "
|
||||
f"{summary.selected_count}/{summary.total_count},跳过 {summary.skipped_count} 个旧文件"
|
||||
)
|
||||
return sorted(selected)
|
||||
|
||||
|
||||
def _build_global_map(extract_fields: list[dict]) -> tuple[dict[str, str], dict[str, str]]:
|
||||
field_map: dict[str, str] = {}
|
||||
type_map: dict[str, str] = {}
|
||||
for field in extract_fields:
|
||||
target = field.get("Field")
|
||||
if not target:
|
||||
continue
|
||||
type_map[target] = field.get("Type", "string")
|
||||
for source in field.get("Extract", []):
|
||||
field_map[source] = target
|
||||
return field_map, type_map
|
||||
|
||||
|
||||
def _build_rj_maps(table_field_mappings: dict) -> dict[str, tuple[dict[str, str], dict[str, str]]]:
|
||||
maps: dict[str, tuple[dict[str, str], dict[str, str]]] = {}
|
||||
for table, fields in table_field_mappings.items():
|
||||
field_map: dict[str, str] = {}
|
||||
type_map: dict[str, str] = {}
|
||||
for field in fields:
|
||||
source = field.get("Source")
|
||||
target = field.get("Target")
|
||||
if source and target:
|
||||
field_map[source] = target
|
||||
type_map[target] = field.get("Type", "string")
|
||||
maps[table] = (field_map, type_map)
|
||||
return maps
|
||||
|
||||
|
||||
def _ordered_targets(headers: list[str], field_map: dict[str, str]) -> list[str]:
|
||||
targets: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for col in headers:
|
||||
target = field_map.get(col)
|
||||
if target and target not in seen:
|
||||
seen.add(target)
|
||||
targets.append(target)
|
||||
return targets
|
||||
|
||||
|
||||
def _read_headers(csv_file: Path) -> list[str]:
|
||||
try:
|
||||
df = pd.read_csv(csv_file, encoding=_detect_encoding(csv_file), nrows=0, dtype=str)
|
||||
return list(df.columns)
|
||||
except Exception: # noqa: BLE001
|
||||
return []
|
||||
|
||||
|
||||
def _detect_encoding(file_path: Path) -> str:
|
||||
with open(file_path, "rb") as handle:
|
||||
result = chardet.detect(handle.read(8192))
|
||||
encoding = (result.get("encoding") or "utf-8").lower()
|
||||
if "utf" in encoding:
|
||||
return "utf-8"
|
||||
if "gb" in encoding:
|
||||
return "gbk"
|
||||
return "utf-8"
|
||||
|
||||
|
||||
def _extract_zip(zip_file: Path, log: Callable[[str], None]) -> None:
|
||||
for enc in ("utf-8", "gbk", "cp437"):
|
||||
try:
|
||||
with zipfile.ZipFile(zip_file, "r", metadata_encoding=enc) as zf:
|
||||
_extract_members(zf, zip_file.parent, log)
|
||||
return
|
||||
except (UnicodeDecodeError, zipfile.BadZipFile):
|
||||
continue
|
||||
raise RuntimeError("无法解压(编码检测失败)")
|
||||
|
||||
|
||||
def _extract_members(zf: zipfile.ZipFile, target_dir: Path, log: Callable[[str], None]) -> None:
|
||||
root = target_dir.resolve()
|
||||
for member in zf.infolist():
|
||||
name = member.filename.replace("\\", "/")
|
||||
target = (root / name).resolve()
|
||||
try:
|
||||
target.relative_to(root)
|
||||
except ValueError:
|
||||
log(f"[WARN] 跳过不安全的 ZIP 条目: {member.filename}")
|
||||
continue
|
||||
if member.is_dir():
|
||||
target.mkdir(parents=True, exist_ok=True)
|
||||
continue
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
with zf.open(member) as source, target.open("wb") as out:
|
||||
shutil.copyfileobj(source, out)
|
||||
|
||||
|
||||
def _clean_numeric_text(series: pd.Series) -> pd.Series:
|
||||
return (
|
||||
series.str.strip()
|
||||
.str.replace(",", "", regex=False)
|
||||
.str.replace(",", "", regex=False)
|
||||
.str.replace("%", "", regex=False)
|
||||
.str.replace("%", "", regex=False)
|
||||
.str.replace("\t", "", regex=False)
|
||||
.str.replace(" ", "", regex=False)
|
||||
)
|
||||
|
||||
|
||||
def _numeric_series(series: pd.Series) -> pd.Series:
|
||||
text = series.astype("string")
|
||||
has_percent = text.str.contains(r"[%%]", regex=True, na=False)
|
||||
cleaned = _clean_numeric_text(text)
|
||||
zero_mask = cleaned.isna() | cleaned.str.upper().isin(ZERO_TEXTS)
|
||||
numeric = pd.to_numeric(cleaned.mask(zero_mask, "0"), errors="coerce").fillna(0)
|
||||
numeric[has_percent & numeric.notna()] = numeric[has_percent & numeric.notna()] / 100
|
||||
return numeric
|
||||
|
||||
|
||||
def _convert_int(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
rounded = _numeric_series(series).round()
|
||||
return pd.Series([int(v) for v in rounded], index=series.index, dtype=object)
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
|
||||
|
||||
def _convert_float(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
numeric = _numeric_series(series)
|
||||
return pd.Series([float(v) for v in numeric], index=series.index, dtype=object)
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
|
||||
|
||||
def _convert_datetime(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
valid = series.notna() & (series != "") & (series.astype(str).str.strip() != "")
|
||||
if not valid.any():
|
||||
return pd.Series([None] * len(series), index=series.index)
|
||||
parsed = pd.Series([pd.NaT] * len(series), index=series.index)
|
||||
remaining = valid.copy()
|
||||
for fmt in DATETIME_FORMATS:
|
||||
if not remaining.any():
|
||||
break
|
||||
try:
|
||||
temp = pd.to_datetime(series[remaining], errors="coerce", format=fmt)
|
||||
except Exception: # noqa: BLE001
|
||||
continue
|
||||
ok = temp.notna()
|
||||
if ok.any():
|
||||
idx = remaining[remaining].index[ok]
|
||||
parsed.loc[idx] = temp[ok].values
|
||||
remaining.loc[idx] = False
|
||||
if remaining.any():
|
||||
try:
|
||||
temp = pd.to_datetime(series[remaining], errors="coerce", format="mixed", dayfirst=False)
|
||||
ok = temp.notna()
|
||||
if ok.any():
|
||||
idx = remaining[remaining].index[ok]
|
||||
parsed.loc[idx] = temp[ok].values
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return parsed.dt.strftime("%Y-%m-%d %H:%M:%S")
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
@@ -12,7 +12,7 @@ from typing import Any
|
||||
from app.config import BASE_DIR
|
||||
|
||||
|
||||
DEFAULT_EXPIRES_ON = date(2026, 6, 20)
|
||||
DEFAULT_EXPIRES_ON = date(2026, 12, 30)
|
||||
EXTEND_DAYS = 30
|
||||
LICENSE_FILE = BASE_DIR / "license.dat"
|
||||
_SECRET = b"CapacityReport local license v1"
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
"""Metrix 仓库模式的处理流水线:CSV 处理 → 平台导入暂存表 → run-script(single_session) 跑报表 SQL。
|
||||
|
||||
仅当 warehouse_type == "metrix" 时使用;直连 MySQL 模式走原版 DataProcessor。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from app.config import SQL_SCRIPT, AppConfig, MetrixConfig
|
||||
from app.processor import ProcessLogger
|
||||
from app.services.csv_processor import CsvProcessor
|
||||
from app.services.platform import make_client
|
||||
|
||||
RJ_DIR_TO_TABLE = {
|
||||
"2.6RJGD": "2_6GRJGD",
|
||||
"2.6RJYD": "2_6GRJYD",
|
||||
"700RJGD": "700MRJGD",
|
||||
"700RJYD": "700MRJYD",
|
||||
}
|
||||
RESULT_TABLES = ["4G_结果表", "5G_结果表"]
|
||||
|
||||
|
||||
def validate_metrix(metrix: MetrixConfig) -> None:
|
||||
missing = []
|
||||
if not metrix.base_url:
|
||||
missing.append("平台地址")
|
||||
if not metrix.token:
|
||||
missing.append("API Token")
|
||||
if not metrix.database_conn_id:
|
||||
missing.append("数据库连接 ID")
|
||||
if missing:
|
||||
raise RuntimeError("Metrix 连接配置不完整: " + ", ".join(missing))
|
||||
|
||||
|
||||
def build_processor_config(app_config: AppConfig) -> dict:
|
||||
metrix = app_config.metrix.normalized()
|
||||
rj = app_config.rj_data.normalized()
|
||||
return {
|
||||
"recent_days": metrix.recent_days,
|
||||
"sheet_filter": list(app_config.sheet_filter),
|
||||
"data_dir_to_table": dict(metrix.data_dir_to_table),
|
||||
"extract_fields": app_config.extract_fields,
|
||||
"rj": {
|
||||
"enabled": rj.enabled,
|
||||
"weekly_directories": rj.weekly_directories,
|
||||
"dir_to_table": RJ_DIR_TO_TABLE,
|
||||
"table_field_mappings": rj.table_field_mappings,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def read_report_sql() -> str:
|
||||
if not SQL_SCRIPT.exists():
|
||||
return ""
|
||||
return SQL_SCRIPT.read_text(encoding="utf-8").strip()
|
||||
|
||||
|
||||
def run_report_sql(app_config: AppConfig, logger: ProcessLogger) -> list[dict]:
|
||||
metrix = app_config.metrix.normalized()
|
||||
validate_metrix(metrix)
|
||||
report_sql = read_report_sql()
|
||||
if not report_sql:
|
||||
raise RuntimeError("报表 SQL(ReportScript.sql)为空或不存在")
|
||||
client = make_client(metrix)
|
||||
logger.info("执行报表 SQL(single_session)...")
|
||||
result = client.run_script(
|
||||
metrix.database_conn_id,
|
||||
content=report_sql,
|
||||
database=metrix.target_database,
|
||||
single_session=True,
|
||||
run_timeout=7200,
|
||||
)
|
||||
statements = result.get("results", [])
|
||||
failed = [item for item in statements if not item.get("ok")]
|
||||
if result.get("stopped") or failed:
|
||||
for item in failed[:5]:
|
||||
logger.error(f"[SQL] 第 {item.get('index')} 条失败: {item.get('message')}")
|
||||
raise RuntimeError("报表 SQL 执行失败")
|
||||
logger.success(f"报表 SQL 执行完成,共 {len(statements)} 条语句")
|
||||
return statements
|
||||
|
||||
|
||||
def run_import_and_report(work_dir: Path, app_config: AppConfig, logger: ProcessLogger) -> dict:
|
||||
"""处理工作目录数据 → 平台导入暂存表 → 跑报表 SQL。失败抛 RuntimeError。"""
|
||||
metrix = app_config.metrix.normalized()
|
||||
validate_metrix(metrix)
|
||||
|
||||
logger.set_stage("converting")
|
||||
tables = CsvProcessor(work_dir, build_processor_config(app_config), logger.info).process()
|
||||
if not tables:
|
||||
raise RuntimeError("处理后没有产出任何暂存表数据")
|
||||
|
||||
client = make_client(metrix)
|
||||
conn_id = metrix.database_conn_id
|
||||
target_db = metrix.target_database
|
||||
|
||||
# 导入前 DROP 旧暂存表,让自动建表按当周实际列重建。
|
||||
logger.set_stage("importing")
|
||||
drop_sql = "".join(f"DROP TABLE IF EXISTS `{table}`;\n" for table in tables)
|
||||
drop_result = client.run_script(conn_id, content=drop_sql, database=target_db, run_timeout=600)
|
||||
if drop_result.get("stopped"):
|
||||
raise RuntimeError("清理旧暂存表失败")
|
||||
|
||||
for table, csv_path in tables.items():
|
||||
logger.info(f"导入暂存表 {table} ...")
|
||||
job_id = client.import_csv(conn_id, table, csv_path, mode="overwrite", database=target_db, create_table=True)
|
||||
job = client.wait_job(job_id)
|
||||
if job.get("status") != "success":
|
||||
raise RuntimeError(f"暂存表 {table} 导入失败: {job.get('error_code') or job.get('status')}")
|
||||
logger.success(f"暂存表 {table} 导入完成")
|
||||
|
||||
logger.set_stage("scripting")
|
||||
statements = run_report_sql(app_config, logger)
|
||||
return {"tables": list(tables.keys()), "statements": len(statements)}
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Metrix 平台集成:API 客户端 + 储存下载器。
|
||||
|
||||
当 source_type/warehouse_type 选 "metrix" 时,源数据走平台储存模块、数据仓库走平台数据库模块。
|
||||
连接信息(地址/token/storage_id/database_conn_id/target_database)来自 Configure.json 的 Metrix 段。
|
||||
储存下载器与 RemoteDataDownloader 接口一致,可被源工厂直接替换。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import date
|
||||
from pathlib import Path
|
||||
from typing import Iterable
|
||||
|
||||
import requests
|
||||
|
||||
from app.config import AppConfig, MetrixConfig
|
||||
from app.services.remote_download import RemoteDownloadResult, RemoteFileInfo
|
||||
from app.utils.file_dates import parse_file_date_range, select_recent_items_by_directory
|
||||
|
||||
|
||||
class PlatformClient:
|
||||
"""平台储存 + 数据库模块的最小 API 封装(Bearer Token 鉴权)。"""
|
||||
|
||||
def __init__(self, base_url: str, token: str, timeout: int = 60):
|
||||
if not base_url:
|
||||
raise ValueError("缺少平台地址,请在系统设置的 Metrix 连接中填写")
|
||||
if not token:
|
||||
raise ValueError("缺少平台 API Token,请在系统设置的 Metrix 连接中填写")
|
||||
self.base = base_url.rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.session = requests.Session()
|
||||
self.session.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
# --- 储存模块 --------------------------------------------------------
|
||||
def list_storage_files(self, storage_id: str, path: str = "/", recursive: bool = True) -> list[dict]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/storages/{storage_id}/files",
|
||||
params={"path": path, "recursive": "true" if recursive else "false"},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json().get("entries", [])
|
||||
|
||||
def download_storage_file(self, storage_id: str, path: str, dest: Path) -> None:
|
||||
dest.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self.session.get(
|
||||
f"{self.base}/api/storages/{storage_id}/download",
|
||||
params={"path": path},
|
||||
stream=True,
|
||||
timeout=self.timeout,
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
with dest.open("wb") as handle:
|
||||
for chunk in resp.iter_content(chunk_size=1024 * 64):
|
||||
if chunk:
|
||||
handle.write(chunk)
|
||||
|
||||
def batch_delete_storage(self, storage_id: str, paths: list[str]) -> int:
|
||||
deleted = 0
|
||||
for start in range(0, len(paths), 100):
|
||||
chunk = [p for p in paths[start:start + 100] if p]
|
||||
if not chunk:
|
||||
continue
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/storages/{storage_id}/batch-delete",
|
||||
json={"paths": chunk},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
deleted += len(chunk)
|
||||
return deleted
|
||||
|
||||
# --- 数据库模块 ------------------------------------------------------
|
||||
def import_csv(self, conn_id: str, table: str, csv_path: Path, mode: str = "overwrite",
|
||||
database: str = "", create_table: bool = True, upload_timeout: int = 1800) -> str:
|
||||
with csv_path.open("rb") as handle:
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/import",
|
||||
files={"file": (csv_path.name, handle, "text/csv")},
|
||||
data={
|
||||
"format": "csv",
|
||||
"target_table": table,
|
||||
"mode": mode,
|
||||
"database": database,
|
||||
"mapping": "{}",
|
||||
"create_table": "true" if create_table else "false",
|
||||
},
|
||||
timeout=upload_timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()["job_id"]
|
||||
|
||||
def wait_job(self, job_id: str, interval: int = 2, max_wait: int = 7200) -> dict:
|
||||
deadline = time.time() + max_wait
|
||||
while time.time() < deadline:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/database-transfer-jobs/{job_id}", timeout=self.timeout
|
||||
)
|
||||
resp.raise_for_status()
|
||||
job = resp.json()
|
||||
if job.get("status") in ("success", "failed"):
|
||||
return job
|
||||
time.sleep(interval)
|
||||
raise TimeoutError(f"导入任务 {job_id} 超过 {max_wait}s 仍未完成")
|
||||
|
||||
def run_script(self, conn_id: str, script_id: int | None = None, content: str = "",
|
||||
database: str = "", single_session: bool = False, run_timeout: int = 7200) -> dict:
|
||||
body: dict = {"database": database, "stop_on_error": True, "single_session": single_session}
|
||||
if content:
|
||||
body["content"] = content
|
||||
if script_id is not None:
|
||||
body["script_id"] = int(script_id)
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/run-script",
|
||||
json=body,
|
||||
timeout=run_timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
# --- 数据库读 / 导出(供仓库代理使用)-------------------------------
|
||||
def list_tables(self, conn_id: str, database: str = "") -> list[str]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/tables",
|
||||
params={"database": database},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return [str(item.get("name")) for item in resp.json() if item.get("name")]
|
||||
|
||||
def table_columns(self, conn_id: str, table: str, database: str = "") -> list[dict]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/tables/{table}",
|
||||
params={"database": database},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json().get("columns", [])
|
||||
|
||||
def table_data(self, conn_id: str, table: str, database: str = "", page: int = 1, page_size: int = 50,
|
||||
order_by: str = "", order_dir: str = "asc") -> dict:
|
||||
params = {"database": database, "table": table, "page": page, "page_size": page_size}
|
||||
if order_by:
|
||||
params["order_by"] = order_by
|
||||
params["order_dir"] = "desc" if str(order_dir).lower().startswith("desc") else "asc"
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/table-data", params=params, timeout=self.timeout
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
def submit_export(self, conn_id: str, tables: list[str], fmt: str, database: str = "") -> str:
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/export",
|
||||
json={"format": fmt, "database": database, "tables": tables},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()["job_id"]
|
||||
|
||||
def download_job_file(self, job_id: str, dest: Path) -> None:
|
||||
dest.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self.session.get(
|
||||
f"{self.base}/api/database-transfer-jobs/{job_id}/download", stream=True, timeout=self.timeout
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
with dest.open("wb") as handle:
|
||||
for chunk in resp.iter_content(chunk_size=1024 * 64):
|
||||
if chunk:
|
||||
handle.write(chunk)
|
||||
|
||||
|
||||
def make_client(metrix: MetrixConfig) -> PlatformClient:
|
||||
metrix = metrix.normalized()
|
||||
return PlatformClient(metrix.base_url, metrix.token)
|
||||
|
||||
|
||||
def make_source_downloader(app_config: AppConfig, logger=None):
|
||||
"""Return a file-source downloader matching app_config.source_type. FTP/SFTP use the
|
||||
original RemoteDataDownloader; 'metrix' uses PlatformStorageDownloader. Both share the
|
||||
interface test_connection / list_remote_zip_files / download_to / delete_source_files."""
|
||||
if app_config.source_type == "metrix":
|
||||
return PlatformStorageDownloader(app_config, logger)
|
||||
from app.services.remote_download import RemoteDataDownloader
|
||||
|
||||
return RemoteDataDownloader(app_config.remote_data, logger)
|
||||
|
||||
|
||||
class PlatformStorageDownloader:
|
||||
"""平台储存版下载器,接口与 RemoteDataDownloader 对齐,可被源工厂直接替换。"""
|
||||
|
||||
def __init__(self, app_config: AppConfig, logger=None):
|
||||
self.metrix = app_config.metrix.normalized()
|
||||
self.remote_dir = (app_config.remote_data.remote_dir or "/").strip() or "/"
|
||||
self.logger = logger
|
||||
self.client = make_client(self.metrix)
|
||||
|
||||
def _log(self, message: str) -> None:
|
||||
if self.logger:
|
||||
self.logger(message)
|
||||
|
||||
def test_connection(self) -> None:
|
||||
if not self.metrix.storage_id:
|
||||
raise ValueError("缺少储存连接 ID,请在系统设置的 Metrix 连接中填写")
|
||||
self.client.list_storage_files(self.metrix.storage_id, self.remote_dir, recursive=False)
|
||||
|
||||
def list_remote_zip_files(self, directory: str | None = None) -> list[RemoteFileInfo]:
|
||||
path = self._join(self.remote_dir, directory.strip("/")) if directory else self.remote_dir
|
||||
entries = self.client.list_storage_files(self.metrix.storage_id, path, recursive=True)
|
||||
files: list[RemoteFileInfo] = []
|
||||
for entry in entries:
|
||||
if entry.get("is_dir"):
|
||||
continue
|
||||
name = str(entry.get("name", ""))
|
||||
if not name.lower().endswith(".zip"):
|
||||
continue
|
||||
files.append(self._info(str(entry.get("path", "")), int(entry.get("size", 0) or 0)))
|
||||
return files
|
||||
|
||||
def download_to(self, destination: Path, target_dates: Iterable[date] | None = None) -> RemoteDownloadResult:
|
||||
destination = Path(destination)
|
||||
destination.mkdir(parents=True, exist_ok=True)
|
||||
zip_files = self.list_remote_zip_files()
|
||||
date_filter = set(target_dates or [])
|
||||
|
||||
if date_filter:
|
||||
selected = self._select_by_dates(zip_files, date_filter)
|
||||
elif zip_files:
|
||||
selected, summaries = select_recent_items_by_directory(
|
||||
zip_files,
|
||||
parent_key=lambda item: item.parent,
|
||||
name_key=lambda item: item.name,
|
||||
)
|
||||
for summary in summaries:
|
||||
if summary.skipped_count and summary.start_date and summary.max_date:
|
||||
self._log(
|
||||
f"储存目录 {summary.directory or '.'}: 仅下载 "
|
||||
f"{summary.start_date.isoformat()} 至 {summary.max_date.isoformat()} 的 "
|
||||
f"{summary.selected_count}/{summary.total_count} 个 ZIP,跳过 {summary.skipped_count} 个旧文件"
|
||||
)
|
||||
else:
|
||||
selected = []
|
||||
|
||||
result = RemoteDownloadResult()
|
||||
for remote_file in selected:
|
||||
dest = destination / remote_file.relative_path
|
||||
self._log(f"下载: {remote_file.relative_path}")
|
||||
self.client.download_storage_file(self.metrix.storage_id, remote_file.path, dest)
|
||||
result.file_count += 1
|
||||
result.total_bytes += remote_file.size or (dest.stat().st_size if dest.exists() else 0)
|
||||
result.remote_files.append(remote_file.path)
|
||||
return result
|
||||
|
||||
def delete_source_files(self, remote_files: Iterable[str] | None = None) -> int:
|
||||
files = [path for path in (remote_files or []) if path]
|
||||
if not files:
|
||||
return 0
|
||||
self._log(f"清理储存源文件,共 {len(files)} 个")
|
||||
return self.client.batch_delete_storage(self.metrix.storage_id, files)
|
||||
|
||||
# --- helpers ---------------------------------------------------------
|
||||
def _select_by_dates(self, zip_files: list[RemoteFileInfo], target_dates: set[date]) -> list[RemoteFileInfo]:
|
||||
grouped: dict[str, list[RemoteFileInfo]] = {}
|
||||
for remote_file in zip_files:
|
||||
grouped.setdefault(remote_file.parent, []).append(remote_file)
|
||||
|
||||
selected: list[RemoteFileInfo] = []
|
||||
for parent, files in sorted(grouped.items(), key=lambda item: item[0]):
|
||||
picked = [
|
||||
remote_file
|
||||
for remote_file in files
|
||||
if (date_range := parse_file_date_range(remote_file.name))
|
||||
and (
|
||||
date_range.covers_all(target_dates)
|
||||
if date_range.span_days > 1
|
||||
else date_range.covers_any(target_dates)
|
||||
)
|
||||
]
|
||||
selected.extend(picked)
|
||||
skipped = len(files) - len(picked)
|
||||
if skipped:
|
||||
self._log(
|
||||
f"储存目录 {parent or '.'}: 仅下载目标日期 "
|
||||
f"{min(target_dates).isoformat()} 至 {max(target_dates).isoformat()} 的 "
|
||||
f"{len(picked)}/{len(files)} 个 ZIP,跳过 {skipped} 个非目标文件"
|
||||
)
|
||||
return selected
|
||||
|
||||
@staticmethod
|
||||
def _join(parent: str, child: str) -> str:
|
||||
parent = (parent or "").replace("\\", "/").rstrip("/")
|
||||
if not parent:
|
||||
return child
|
||||
if parent == "/":
|
||||
return f"/{child}"
|
||||
return f"{parent}/{child}"
|
||||
|
||||
def _info(self, remote_path: str, size: int = 0) -> RemoteFileInfo:
|
||||
normalized_root = self.remote_dir.replace("\\", "/").rstrip("/")
|
||||
normalized_path = remote_path.replace("\\", "/")
|
||||
if normalized_root and normalized_root != "/" and normalized_path.startswith(f"{normalized_root}/"):
|
||||
relative_path = normalized_path[len(normalized_root) + 1:]
|
||||
else:
|
||||
relative_path = normalized_path.lstrip("/")
|
||||
relative = Path(relative_path)
|
||||
parent = str(relative.parent).replace("\\", "/")
|
||||
if parent == ".":
|
||||
parent = ""
|
||||
return RemoteFileInfo(
|
||||
path=remote_path,
|
||||
relative_path=relative_path,
|
||||
parent=parent,
|
||||
name=relative.name,
|
||||
size=size,
|
||||
)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""数据仓库抽象:直连 MySQL 或 Metrix 数据库平台。
|
||||
|
||||
`make_warehouse(config)` 按 warehouse_type 返回:
|
||||
- 直连 MySQL: 原版 `DatabaseManager`(已具备下列方法)。
|
||||
- Metrix: `MetrixWarehouse`,用平台数据库 API 实现相同方法,供查看/导出路由透明替换。
|
||||
|
||||
两者都提供: test_connection / get_server_info / get_tables / get_table_info /
|
||||
query_table / truncate_table / drop_table / drop_all_tables / execute_sql。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from app.config import AppConfig
|
||||
from app.database import DatabaseManager
|
||||
from app.services.platform import make_client
|
||||
|
||||
|
||||
def make_warehouse(config: AppConfig):
|
||||
if config.warehouse_type == "metrix":
|
||||
return MetrixWarehouse(config)
|
||||
return DatabaseManager(config)
|
||||
|
||||
|
||||
def _quote_ident(name: str) -> str:
|
||||
return "`" + str(name).replace("`", "``") + "`"
|
||||
|
||||
|
||||
def _quote_value(value: str) -> str:
|
||||
return "'" + str(value).replace("\\", "\\\\").replace("'", "''") + "'"
|
||||
|
||||
|
||||
class MetrixWarehouse:
|
||||
"""用 Metrix 数据库 API 实现 DatabaseManager 的只读/管理子集。"""
|
||||
|
||||
def __init__(self, config: AppConfig):
|
||||
self.metrix = config.metrix.normalized()
|
||||
self.conn_id = self.metrix.database_conn_id
|
||||
self.database = self.metrix.target_database
|
||||
self.client = make_client(self.metrix)
|
||||
|
||||
# --- 连接 / 诊断 -----------------------------------------------------
|
||||
def test_connection(self) -> Tuple[bool, str]:
|
||||
try:
|
||||
self.client.list_tables(self.conn_id, self.database)
|
||||
return True, "连接成功"
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return False, str(exc)
|
||||
|
||||
def get_server_info(self) -> Dict[str, Any]:
|
||||
version = "Metrix"
|
||||
try:
|
||||
res = self.client.run_script(self.conn_id, content="SELECT VERSION() AS v", database=self.database, run_timeout=30)
|
||||
rows = (res.get("results") or [{}])[0].get("rows") or []
|
||||
if rows:
|
||||
version = str(list(rows[0].values())[0])
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return {"version": version, "load_data_infile": True, "load_data_message": "Metrix 平台导入"}
|
||||
|
||||
# --- 表 / 数据 -------------------------------------------------------
|
||||
def get_tables(self) -> List[str]:
|
||||
return self.client.list_tables(self.conn_id, self.database)
|
||||
|
||||
def get_table_info(self, table_name: str) -> Dict[str, Any]:
|
||||
columns = self.client.table_columns(self.conn_id, table_name, self.database)
|
||||
# Map Metrix column shape -> original DESCRIBE-like shape used by the frontend.
|
||||
mapped = [
|
||||
{
|
||||
"Field": col.get("name"),
|
||||
"Type": col.get("type", ""),
|
||||
"Null": "YES" if col.get("nullable", True) else "NO",
|
||||
"Key": "PRI" if col.get("primary_key") else "",
|
||||
"Default": col.get("default"),
|
||||
"Extra": "auto_increment" if col.get("autoincrement") else "",
|
||||
}
|
||||
for col in columns
|
||||
]
|
||||
data = self.client.table_data(self.conn_id, table_name, self.database, page=1, page_size=1)
|
||||
return {"name": table_name, "columns": mapped, "row_count": int(data.get("total") or 0)}
|
||||
|
||||
def query_table(
|
||||
self,
|
||||
table_name: str,
|
||||
page: int = 1,
|
||||
page_size: int = 50,
|
||||
filters: Optional[Dict[str, str]] = None,
|
||||
order_by: Optional[str] = None,
|
||||
order_dir: str = "ASC",
|
||||
) -> Dict[str, Any]:
|
||||
active_filters = {k: v for k, v in (filters or {}).items() if v}
|
||||
if active_filters:
|
||||
return self._query_with_filters(table_name, page, page_size, active_filters, order_by, order_dir)
|
||||
data = self.client.table_data(
|
||||
self.conn_id, table_name, self.database, page=page, page_size=page_size,
|
||||
order_by=order_by or "", order_dir=order_dir,
|
||||
)
|
||||
total = int(data.get("total") or 0)
|
||||
return {
|
||||
"data": data.get("rows", []),
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"total_pages": (total + page_size - 1) // page_size if page_size else 0,
|
||||
}
|
||||
|
||||
def _query_with_filters(self, table_name, page, page_size, filters, order_by, order_dir) -> Dict[str, Any]:
|
||||
where = " AND ".join(f"{_quote_ident(col)} LIKE {_quote_value('%' + str(val) + '%')}" for col, val in filters.items())
|
||||
where_sql = f" WHERE {where}" if where else ""
|
||||
table_sql = _quote_ident(table_name)
|
||||
total_res = self.client.run_script(
|
||||
self.conn_id, content=f"SELECT COUNT(*) AS n FROM {table_sql}{where_sql}",
|
||||
database=self.database, run_timeout=120,
|
||||
)
|
||||
total = int(((total_res.get("results") or [{}])[0].get("rows") or [{}])[0].get("n") or 0)
|
||||
order_sql = ""
|
||||
if order_by:
|
||||
direction = "DESC" if str(order_dir).upper() == "DESC" else "ASC"
|
||||
order_sql = f" ORDER BY {_quote_ident(order_by)} {direction}"
|
||||
offset = max(page - 1, 0) * page_size
|
||||
data_res = self.client.run_script(
|
||||
self.conn_id,
|
||||
content=f"SELECT * FROM {table_sql}{where_sql}{order_sql} LIMIT {int(page_size)} OFFSET {int(offset)}",
|
||||
database=self.database, run_timeout=300,
|
||||
)
|
||||
rows = (data_res.get("results") or [{}])[0].get("rows") or []
|
||||
return {
|
||||
"data": rows,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"total_pages": (total + page_size - 1) // page_size if page_size else 0,
|
||||
}
|
||||
|
||||
# --- 管理操作 --------------------------------------------------------
|
||||
def truncate_table(self, table_name: str) -> bool:
|
||||
self._run(f"TRUNCATE TABLE {_quote_ident(table_name)}")
|
||||
return True
|
||||
|
||||
def drop_table(self, table_name: str) -> bool:
|
||||
self._run(f"DROP TABLE IF EXISTS {_quote_ident(table_name)}")
|
||||
return True
|
||||
|
||||
def drop_all_tables(self) -> Dict[str, Any]:
|
||||
tables = self.get_tables()
|
||||
if not tables:
|
||||
return {"success": True, "dropped_count": 0, "tables": []}
|
||||
drop_sql = "".join(f"DROP TABLE IF EXISTS {_quote_ident(t)};\n" for t in tables)
|
||||
self._run(drop_sql)
|
||||
return {"success": True, "dropped_count": len(tables), "tables": tables}
|
||||
|
||||
def execute_sql(self, sql: str) -> Tuple[bool, Any]:
|
||||
try:
|
||||
res = self.client.run_script(self.conn_id, content=sql, database=self.database, run_timeout=600)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return False, str(exc)
|
||||
if res.get("stopped"):
|
||||
failed = [r for r in res.get("results", []) if not r.get("ok")]
|
||||
return False, (failed[0].get("message") if failed else "SQL 执行失败")
|
||||
results = res.get("results", [])
|
||||
last = results[-1] if results else {}
|
||||
if last.get("rows"):
|
||||
return True, last["rows"]
|
||||
return True, {"affected_rows": sum(int(r.get("affected_rows") or 0) for r in results)}
|
||||
|
||||
def _run(self, content: str) -> None:
|
||||
res = self.client.run_script(self.conn_id, content=content, database=self.database, run_timeout=600)
|
||||
if res.get("stopped"):
|
||||
failed = [r for r in res.get("results", []) if not r.get("ok")]
|
||||
raise RuntimeError(failed[0].get("message") if failed else "SQL 执行失败")
|
||||
Reference in New Issue
Block a user