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
-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,
)