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:
2026-06-24 05:34:28 +08:00
parent f708d6947a
commit dab672621c
34 changed files with 1586 additions and 1842 deletions
-135
View File
@@ -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 文档仅登录后可访问。",
}
+41 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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),
+8 -4
View File
@@ -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:
+18 -8
View File
@@ -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),
-6
View File
@@ -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
View File
@@ -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
View File
@@ -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()
-329
View File
@@ -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")
+2 -1
View File
@@ -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()
+402
View File
@@ -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
+1 -1
View File
@@ -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"
+114
View File
@@ -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)}
+315
View File
@@ -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,
)
+170
View File
@@ -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 执行失败")