feat: 双模式集成并精简 API 文档/Token、设置页卡片自适应
- 源/仓库各自可在直连(FTP/SFTP、MySQL)与 Metrix 存储/数据库平台间独立选择,两侧互不依赖 - Metrix 模式下源走平台储存 API、仓库走平台导入 + run-script(single_session)、查看导出代理到平台 - 去掉对外 API 文档与 API Token(前后端 + auth/config 解耦),业务接口仅登录态可访问 - 授权默认到期日改为 2026-12-30 - 设置页卡片改横向自适应(宽屏并排、窄屏换行),处理历史保留卡片收窄
This commit is contained in:
@@ -1,329 +0,0 @@
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import secrets
|
||||
from dataclasses import asdict, dataclass
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
from threading import RLock
|
||||
from typing import Any
|
||||
|
||||
from app.config import BASE_DIR
|
||||
|
||||
|
||||
API_TOKENS_PATH = BASE_DIR / "api_tokens.json"
|
||||
API_TOKEN_PREFIX = "cap_"
|
||||
API_TOKEN_SECRET = "CapaReportApiTokenSecret2026"
|
||||
_STORE_LOCK = RLock()
|
||||
|
||||
|
||||
@dataclass
|
||||
class ApiTokenRecord:
|
||||
id: str
|
||||
name: str
|
||||
token_hash: str
|
||||
prefix: str
|
||||
suffix: str
|
||||
created_at: str
|
||||
expires_at: str | None
|
||||
enabled: bool
|
||||
last_used_at: str | None = None
|
||||
last_used_from: str | None = None
|
||||
token: str | None = None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def ensure_store() -> None:
|
||||
if API_TOKENS_PATH.exists():
|
||||
return
|
||||
with _STORE_LOCK:
|
||||
if API_TOKENS_PATH.exists():
|
||||
return
|
||||
API_TOKENS_PATH.write_text(json.dumps({"tokens": []}, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
|
||||
|
||||
def list_tokens() -> list[dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
return [record_to_public_dict(record) for record in _load_records()]
|
||||
|
||||
|
||||
def export_tokens() -> list[dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
return [record.to_dict() for record in _load_records()]
|
||||
|
||||
|
||||
def import_tokens(items: Any) -> int:
|
||||
if not isinstance(items, list):
|
||||
return 0
|
||||
|
||||
records: list[ApiTokenRecord] = []
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
record = _record_from_dict(item)
|
||||
if record.token_hash:
|
||||
records.append(record)
|
||||
|
||||
with _STORE_LOCK:
|
||||
_save_records(records)
|
||||
return len(records)
|
||||
|
||||
|
||||
def create_token(
|
||||
name: str,
|
||||
expires_in_days: int | None = None,
|
||||
enabled: bool = True,
|
||||
expires_at: str | None = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
raw_token = generate_raw_token()
|
||||
resolved_expires_at = normalize_expiration(expires_at)
|
||||
if resolved_expires_at is None and expires_in_days is not None:
|
||||
resolved_expires_at = expires_at_from_days(expires_in_days)
|
||||
|
||||
record = ApiTokenRecord(
|
||||
id=secrets.token_hex(8),
|
||||
name=name.strip() or "未命名 Token",
|
||||
token_hash=hash_token(raw_token),
|
||||
prefix=raw_token[:12],
|
||||
suffix=raw_token[-12:],
|
||||
created_at=utc_now(),
|
||||
expires_at=resolved_expires_at,
|
||||
enabled=bool(enabled),
|
||||
token=raw_token,
|
||||
)
|
||||
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
records.append(record)
|
||||
_save_records(records)
|
||||
|
||||
return raw_token, record_to_public_dict(record)
|
||||
|
||||
|
||||
def update_token(token_id: str, **changes: Any) -> dict[str, Any]:
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
for index, record in enumerate(records):
|
||||
if record.id != token_id:
|
||||
continue
|
||||
|
||||
if "name" in changes and isinstance(changes["name"], str):
|
||||
record.name = changes["name"].strip() or record.name
|
||||
if "enabled" in changes:
|
||||
record.enabled = bool(changes["enabled"])
|
||||
if "expires_at" in changes:
|
||||
record.expires_at = normalize_expiration(changes["expires_at"])
|
||||
|
||||
records[index] = record
|
||||
_save_records(records)
|
||||
return record_to_public_dict(record)
|
||||
|
||||
raise KeyError(f"Token not found: {token_id}")
|
||||
|
||||
|
||||
def delete_token(token_id: str) -> None:
|
||||
with _STORE_LOCK:
|
||||
records = [record for record in _load_records() if record.id != token_id]
|
||||
_save_records(records)
|
||||
|
||||
|
||||
def delete_tokens(token_ids: list[str]) -> int:
|
||||
token_id_set = {str(token_id).strip() for token_id in token_ids if str(token_id).strip()}
|
||||
if not token_id_set:
|
||||
return 0
|
||||
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
kept_records = [record for record in records if record.id not in token_id_set]
|
||||
_save_records(kept_records)
|
||||
return len(records) - len(kept_records)
|
||||
|
||||
|
||||
def regenerate_token(token_id: str) -> tuple[str, dict[str, Any]]:
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
for index, record in enumerate(records):
|
||||
if record.id != token_id:
|
||||
continue
|
||||
|
||||
raw_token = generate_raw_token()
|
||||
record.token_hash = hash_token(raw_token)
|
||||
record.prefix = raw_token[:12]
|
||||
record.suffix = raw_token[-12:]
|
||||
record.token = raw_token
|
||||
record.created_at = utc_now()
|
||||
record.last_used_at = None
|
||||
record.last_used_from = None
|
||||
records[index] = record
|
||||
_save_records(records)
|
||||
return raw_token, record_to_public_dict(record)
|
||||
|
||||
raise KeyError(f"Token not found: {token_id}")
|
||||
|
||||
|
||||
def verify_api_token(raw_token: str) -> dict[str, Any] | None:
|
||||
ensure_store()
|
||||
token_hash = hash_token(raw_token)
|
||||
now = datetime.now(timezone.utc)
|
||||
with _STORE_LOCK:
|
||||
for record in _load_records():
|
||||
if not record.enabled or record.token_hash != token_hash:
|
||||
continue
|
||||
|
||||
expires_at = parse_datetime(record.expires_at)
|
||||
if expires_at and expires_at < now:
|
||||
continue
|
||||
|
||||
return record_to_context(record)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def touch_token_usage(raw_token: str, source: str | None = None) -> None:
|
||||
token_hash = hash_token(raw_token)
|
||||
now = utc_now()
|
||||
with _STORE_LOCK:
|
||||
records = _load_records()
|
||||
updated = False
|
||||
for index, record in enumerate(records):
|
||||
if record.token_hash != token_hash:
|
||||
continue
|
||||
record.last_used_at = now
|
||||
record.last_used_from = source or record.last_used_from
|
||||
records[index] = record
|
||||
updated = True
|
||||
break
|
||||
if updated:
|
||||
_save_records(records)
|
||||
|
||||
|
||||
def generate_raw_token() -> str:
|
||||
return API_TOKEN_PREFIX + secrets.token_urlsafe(36)
|
||||
|
||||
|
||||
def hash_token(raw_token: str) -> str:
|
||||
digest = hmac.new(API_TOKEN_SECRET.encode(), raw_token.encode(), hashlib.sha256).digest()
|
||||
return base64.urlsafe_b64encode(digest).decode().rstrip("=")
|
||||
|
||||
|
||||
def normalize_expiration(value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
trimmed = str(value).strip()
|
||||
if not trimmed:
|
||||
return None
|
||||
parsed = parse_datetime(trimmed)
|
||||
if parsed is None:
|
||||
raise ValueError("Invalid expiration date")
|
||||
return parsed.isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def expires_at_from_days(days: int) -> str:
|
||||
safe_days = max(int(days), 1)
|
||||
expires_at = datetime.now(timezone.utc) + timedelta(days=safe_days)
|
||||
return expires_at.isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def parse_datetime(value: str | None) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value)
|
||||
except ValueError:
|
||||
try:
|
||||
parsed_date = date.fromisoformat(value)
|
||||
except ValueError:
|
||||
return None
|
||||
parsed = datetime.combine(parsed_date, time.max)
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=timezone.utc)
|
||||
return parsed.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def utc_now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def record_to_public_dict(record: ApiTokenRecord, include_hash: bool = False) -> dict[str, Any]:
|
||||
expires_at = parse_datetime(record.expires_at)
|
||||
data = {
|
||||
"id": record.id,
|
||||
"name": record.name,
|
||||
"prefix": record.prefix,
|
||||
"suffix": record.suffix,
|
||||
"created_at": record.created_at,
|
||||
"expires_at": record.expires_at,
|
||||
"enabled": record.enabled,
|
||||
"last_used_at": record.last_used_at,
|
||||
"last_used_from": record.last_used_from,
|
||||
"expired": bool(expires_at and expires_at < datetime.now(timezone.utc)),
|
||||
"token": record.token,
|
||||
"token_available": bool(record.token),
|
||||
}
|
||||
if include_hash:
|
||||
data["token_hash"] = record.token_hash
|
||||
return data
|
||||
|
||||
|
||||
def record_to_context(record: ApiTokenRecord) -> dict[str, Any]:
|
||||
return {
|
||||
"token_id": record.id,
|
||||
"name": record.name,
|
||||
"created_at": record.created_at,
|
||||
"expires_at": record.expires_at,
|
||||
"token_type": "api_token",
|
||||
}
|
||||
|
||||
|
||||
def _load_records() -> list[ApiTokenRecord]:
|
||||
ensure_store()
|
||||
try:
|
||||
raw = json.loads(API_TOKENS_PATH.read_text(encoding="utf-8"))
|
||||
except json.JSONDecodeError:
|
||||
raw = {"tokens": []}
|
||||
|
||||
tokens = raw.get("tokens", []) if isinstance(raw, dict) else []
|
||||
records: list[ApiTokenRecord] = []
|
||||
for item in tokens:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
record = _record_from_dict(item)
|
||||
if record.token_hash:
|
||||
records.append(record)
|
||||
|
||||
records.sort(key=lambda record: record.created_at, reverse=True)
|
||||
return records
|
||||
|
||||
|
||||
def _record_from_dict(item: dict[str, Any]) -> ApiTokenRecord:
|
||||
raw_token = str(item.get("token") or "").strip() or None
|
||||
token_hash = str(item.get("token_hash") or "").strip()
|
||||
if raw_token and not token_hash:
|
||||
token_hash = hash_token(raw_token)
|
||||
|
||||
prefix = str(item.get("prefix") or "")
|
||||
suffix = str(item.get("suffix") or "")
|
||||
if raw_token:
|
||||
prefix = prefix or raw_token[:12]
|
||||
suffix = suffix or raw_token[-12:]
|
||||
|
||||
return ApiTokenRecord(
|
||||
id=str(item.get("id", "")) or secrets.token_hex(8),
|
||||
name=str(item.get("name", "未命名 Token")),
|
||||
token_hash=token_hash,
|
||||
prefix=prefix,
|
||||
suffix=suffix,
|
||||
created_at=str(item.get("created_at", utc_now())),
|
||||
expires_at=item.get("expires_at"),
|
||||
enabled=bool(item.get("enabled", True)),
|
||||
last_used_at=item.get("last_used_at"),
|
||||
last_used_from=item.get("last_used_from"),
|
||||
token=raw_token,
|
||||
)
|
||||
|
||||
|
||||
def _save_records(records: list[ApiTokenRecord]) -> None:
|
||||
payload = {"tokens": [record.to_dict() for record in records]}
|
||||
API_TOKENS_PATH.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
@@ -11,6 +11,7 @@ from fastapi import HTTPException
|
||||
|
||||
from app import state
|
||||
from app.config import CACHE_DIR, AutoSchedulerConfig
|
||||
from app.services.platform import make_source_downloader
|
||||
from app.services.remote_download import RemoteDataDownloader, RemoteFileInfo
|
||||
from app.utils.file_dates import parse_file_date_range, required_week_days
|
||||
|
||||
@@ -186,7 +187,7 @@ class AutoScheduler:
|
||||
return self._trigger_processing(target_dates, ready_flag, manual)
|
||||
|
||||
target_days = required_week_days(scheduler.week_offset)
|
||||
downloader = RemoteDataDownloader(remote_config)
|
||||
downloader = make_source_downloader(app_config)
|
||||
rj_config = app_config.rj_data.normalized()
|
||||
rj_directories = set(rj_config.weekly_directories) if rj_config.enabled else set()
|
||||
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
"""纯数据处理流水线(容器版,无数据库代码)。
|
||||
|
||||
输入:包含已下载周 ZIP/CSV/Excel 的工作目录。
|
||||
输出:每张暂存表一个规范化 CSV({表名: csv 路径});由平台数据库导入 API 建表入库,
|
||||
再由平台 run-script 跑报表 SQL 生成结果表。这里不含任何数据库/LOAD DATA 代码。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
import zipfile
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Callable
|
||||
|
||||
import chardet
|
||||
import pandas as pd
|
||||
|
||||
from app.utils.file_dates import select_recent_items_by_directory
|
||||
|
||||
ZERO_TEXTS = {"", "-", "--", "—", "–", "NA", "N/A", "NULL", "NONE", "NAN", "\\N"}
|
||||
DATETIME_FORMATS = [
|
||||
"ISO8601",
|
||||
"%Y-%m-%d %H:%M:%S",
|
||||
"%Y-%m-%d %H:%M",
|
||||
"%Y/%m/%d %H:%M:%S",
|
||||
"%Y/%m/%d %H:%M",
|
||||
"%Y-%m-%d",
|
||||
"%Y/%m/%d",
|
||||
"%Y年%m月%d日 %H:%M:%S",
|
||||
"%Y年%m月%d日",
|
||||
"%Y%m%d%H%M%S",
|
||||
"%Y%m%d",
|
||||
]
|
||||
MAX_WORKERS = 8
|
||||
|
||||
|
||||
class CsvProcessor:
|
||||
def __init__(self, work_dir: Path, config: dict, log: Callable[[str], None]):
|
||||
self.work_dir = Path(work_dir)
|
||||
self.config = config
|
||||
self.log = log
|
||||
self.recent_days = int(config.get("recent_days", 7))
|
||||
self.sheet_filter = set(config.get("sheet_filter", []))
|
||||
self.data_dir_to_table = {k.upper(): v for k, v in (config.get("data_dir_to_table") or {}).items()}
|
||||
self.field_map, self.type_map = _build_global_map(config.get("extract_fields", []))
|
||||
rj = config.get("rj") or {}
|
||||
self.rj_enabled = bool(rj.get("enabled"))
|
||||
self.rj_weekly_dirs = list(rj.get("weekly_directories") or [])
|
||||
self.rj_dir_to_table = dict(rj.get("dir_to_table") or {})
|
||||
self.rj_maps = _build_rj_maps(rj.get("table_field_mappings") or {})
|
||||
self.out_dir = self.work_dir / ".out"
|
||||
|
||||
def process(self) -> dict[str, Path]:
|
||||
self._unzip_files()
|
||||
self._excel_to_csv()
|
||||
return self._build_table_csvs()
|
||||
|
||||
# --- step 1: unzip ---------------------------------------------------
|
||||
def _unzip_files(self) -> None:
|
||||
zips = self._filter_recent(list(self.work_dir.rglob("*.zip")), "ZIP")
|
||||
self.log(f"解压 ZIP: {len(zips)} 个")
|
||||
for zip_file in zips:
|
||||
try:
|
||||
_extract_zip(zip_file, self.log)
|
||||
except Exception as exc: # noqa: BLE001 - keep going on a bad archive
|
||||
self.log(f"[WARN] 解压失败 {zip_file.name}: {exc}")
|
||||
|
||||
# --- step 2: excel -> csv -------------------------------------------
|
||||
def _excel_to_csv(self) -> None:
|
||||
excels = self._filter_recent(list(self._scan(self.work_dir, (".xlsx", ".xls"))), "Excel")
|
||||
if not excels:
|
||||
return
|
||||
self.log(f"Excel 转 CSV: {len(excels)} 个文件")
|
||||
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool:
|
||||
futures = {pool.submit(self._one_excel, f): f for f in excels}
|
||||
for future in as_completed(futures):
|
||||
excel = futures[future]
|
||||
try:
|
||||
future.result()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self.log(f"[WARN] Excel 处理失败 {excel.name}: {exc}")
|
||||
|
||||
def _one_excel(self, excel_file: Path) -> None:
|
||||
xl = pd.ExcelFile(excel_file, engine="openpyxl")
|
||||
try:
|
||||
for sheet in xl.sheet_names:
|
||||
if sheet in self.sheet_filter:
|
||||
continue
|
||||
out = excel_file.parent / f"{excel_file.stem}_{sheet}.csv"
|
||||
xl.parse(sheet).to_csv(out, index=False, encoding="utf-8")
|
||||
finally:
|
||||
xl.close()
|
||||
|
||||
# --- step 3: build one normalized CSV per staging table -------------
|
||||
def _build_table_csvs(self) -> dict[str, Path]:
|
||||
data_dirs = self._find_data_dirs()
|
||||
if not data_dirs:
|
||||
self.log("[WARN] 未找到任何数据目录")
|
||||
return {}
|
||||
self.out_dir.mkdir(parents=True, exist_ok=True)
|
||||
result: dict[str, Path] = {}
|
||||
for table, directory in data_dirs.items():
|
||||
csv_files = self._filter_recent(list(self._scan(directory, (".csv",))), "CSV", root=directory)
|
||||
if not csv_files:
|
||||
continue
|
||||
field_map, type_map = self._maps_for_table(table)
|
||||
out_path = self.out_dir / f"{table}.csv"
|
||||
rows = self._write_table_csv(table, csv_files, field_map, type_map, out_path)
|
||||
if rows > 0:
|
||||
result[table] = out_path
|
||||
self.log(f"暂存表 {table}: {rows} 行 -> {out_path.name}")
|
||||
return result
|
||||
|
||||
def _write_table_csv(self, table, csv_files, field_map, type_map, out_path: Path) -> int:
|
||||
# First pass: union of target columns across this table's CSV files.
|
||||
union: list[str] = []
|
||||
seen: set[str] = set()
|
||||
frames: list[tuple[Path, list[str]]] = []
|
||||
for csv_file in csv_files:
|
||||
headers = _read_headers(csv_file)
|
||||
targets = _ordered_targets(headers, field_map)
|
||||
if not targets:
|
||||
continue
|
||||
frames.append((csv_file, headers))
|
||||
for target in targets:
|
||||
if target not in seen:
|
||||
seen.add(target)
|
||||
union.append(target)
|
||||
if not union:
|
||||
return 0
|
||||
|
||||
total = 0
|
||||
header_written = False
|
||||
for csv_file, _headers in frames:
|
||||
df = self._normalize(csv_file, field_map, type_map, union)
|
||||
if df is None or df.empty:
|
||||
continue
|
||||
df.to_csv(out_path, index=False, header=not header_written, mode="w" if not header_written else "a", encoding="utf-8")
|
||||
header_written = True
|
||||
total += len(df)
|
||||
return total
|
||||
|
||||
def _normalize(self, csv_file: Path, field_map, type_map, union: list[str]):
|
||||
try:
|
||||
df = pd.read_csv(
|
||||
csv_file,
|
||||
encoding=_detect_encoding(csv_file),
|
||||
dtype=str,
|
||||
na_values=[""],
|
||||
keep_default_na=False,
|
||||
low_memory=True,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self.log(f"[WARN] 读取 CSV 失败 {csv_file.name}: {exc}")
|
||||
return None
|
||||
|
||||
col_map: dict[str, str] = {}
|
||||
mapped: set[str] = set()
|
||||
for col in df.columns:
|
||||
target = field_map.get(col)
|
||||
if target and target not in mapped:
|
||||
col_map[col] = target
|
||||
mapped.add(target)
|
||||
if not col_map:
|
||||
return None
|
||||
|
||||
out = df[list(col_map.keys())].copy()
|
||||
out.columns = list(col_map.values())
|
||||
out = out.fillna("")
|
||||
for col in out.columns:
|
||||
col_type = type_map.get(col, "string")
|
||||
if col_type == "datetime":
|
||||
out[col] = _convert_datetime(out[col])
|
||||
elif col_type == "int":
|
||||
out[col] = _convert_int(out[col])
|
||||
elif col_type == "float":
|
||||
out[col] = _convert_float(out[col])
|
||||
else:
|
||||
out[col] = out[col].astype("string").str.replace("%", "", regex=False).str.slice(0, 255)
|
||||
# Reindex to the shared union columns; fill missing per type so numeric
|
||||
# staging columns never carry '' (the report SQL re-types them later).
|
||||
for col in union:
|
||||
if col not in out.columns:
|
||||
out[col] = "0" if type_map.get(col) in ("int", "float") else ""
|
||||
return out[union]
|
||||
|
||||
# --- directory detection --------------------------------------------
|
||||
def _find_data_dirs(self) -> dict[str, Path]:
|
||||
data_dirs: dict[str, Path] = {}
|
||||
target_names = set(self.data_dir_to_table.keys())
|
||||
for sub in self.work_dir.rglob("*"):
|
||||
if sub.is_dir() and sub.name.upper() in target_names:
|
||||
table = self.data_dir_to_table[sub.name.upper()]
|
||||
data_dirs.setdefault(table, sub)
|
||||
if self.rj_enabled:
|
||||
self._find_rj_dirs(data_dirs)
|
||||
return data_dirs
|
||||
|
||||
def _find_rj_dirs(self, data_dirs: dict[str, Path]) -> None:
|
||||
for weekly in self.rj_weekly_dirs:
|
||||
path = self.work_dir / weekly
|
||||
if path.exists() and path.is_dir() and path.name in self.rj_dir_to_table:
|
||||
data_dirs.setdefault(self.rj_dir_to_table[path.name], path)
|
||||
if not any(table in data_dirs for table in self.rj_dir_to_table.values()):
|
||||
for sub in self.work_dir.rglob("*"):
|
||||
if sub.is_dir() and sub.name in self.rj_dir_to_table:
|
||||
data_dirs.setdefault(self.rj_dir_to_table[sub.name], sub)
|
||||
|
||||
def _maps_for_table(self, table: str):
|
||||
if table in self.rj_maps:
|
||||
return self.rj_maps[table]
|
||||
return self.field_map, self.type_map
|
||||
|
||||
# --- helpers ---------------------------------------------------------
|
||||
def _scan(self, directory: Path, extensions: tuple[str, ...]):
|
||||
for ext in extensions:
|
||||
yield from directory.rglob(f"*{ext}")
|
||||
|
||||
def _filter_recent(self, files: list[Path], label: str, root: Path | None = None) -> list[Path]:
|
||||
if not files:
|
||||
return files
|
||||
base = (root or self.work_dir).resolve()
|
||||
|
||||
def parent_key(file_path: Path) -> str:
|
||||
try:
|
||||
parent = file_path.parent.resolve().relative_to(base)
|
||||
except ValueError:
|
||||
parent = file_path.parent
|
||||
text = str(parent).replace("\\", "/")
|
||||
return "" if text == "." else text
|
||||
|
||||
selected, summaries = select_recent_items_by_directory(
|
||||
files,
|
||||
parent_key=parent_key,
|
||||
name_key=lambda f: f.name,
|
||||
days=self.recent_days,
|
||||
)
|
||||
for summary in summaries:
|
||||
if summary.skipped_count and summary.start_date and summary.max_date:
|
||||
self.log(
|
||||
f"{label} {summary.directory or '.'}: 取 {summary.start_date}~{summary.max_date} "
|
||||
f"{summary.selected_count}/{summary.total_count},跳过 {summary.skipped_count} 个旧文件"
|
||||
)
|
||||
return sorted(selected)
|
||||
|
||||
|
||||
def _build_global_map(extract_fields: list[dict]) -> tuple[dict[str, str], dict[str, str]]:
|
||||
field_map: dict[str, str] = {}
|
||||
type_map: dict[str, str] = {}
|
||||
for field in extract_fields:
|
||||
target = field.get("Field")
|
||||
if not target:
|
||||
continue
|
||||
type_map[target] = field.get("Type", "string")
|
||||
for source in field.get("Extract", []):
|
||||
field_map[source] = target
|
||||
return field_map, type_map
|
||||
|
||||
|
||||
def _build_rj_maps(table_field_mappings: dict) -> dict[str, tuple[dict[str, str], dict[str, str]]]:
|
||||
maps: dict[str, tuple[dict[str, str], dict[str, str]]] = {}
|
||||
for table, fields in table_field_mappings.items():
|
||||
field_map: dict[str, str] = {}
|
||||
type_map: dict[str, str] = {}
|
||||
for field in fields:
|
||||
source = field.get("Source")
|
||||
target = field.get("Target")
|
||||
if source and target:
|
||||
field_map[source] = target
|
||||
type_map[target] = field.get("Type", "string")
|
||||
maps[table] = (field_map, type_map)
|
||||
return maps
|
||||
|
||||
|
||||
def _ordered_targets(headers: list[str], field_map: dict[str, str]) -> list[str]:
|
||||
targets: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for col in headers:
|
||||
target = field_map.get(col)
|
||||
if target and target not in seen:
|
||||
seen.add(target)
|
||||
targets.append(target)
|
||||
return targets
|
||||
|
||||
|
||||
def _read_headers(csv_file: Path) -> list[str]:
|
||||
try:
|
||||
df = pd.read_csv(csv_file, encoding=_detect_encoding(csv_file), nrows=0, dtype=str)
|
||||
return list(df.columns)
|
||||
except Exception: # noqa: BLE001
|
||||
return []
|
||||
|
||||
|
||||
def _detect_encoding(file_path: Path) -> str:
|
||||
with open(file_path, "rb") as handle:
|
||||
result = chardet.detect(handle.read(8192))
|
||||
encoding = (result.get("encoding") or "utf-8").lower()
|
||||
if "utf" in encoding:
|
||||
return "utf-8"
|
||||
if "gb" in encoding:
|
||||
return "gbk"
|
||||
return "utf-8"
|
||||
|
||||
|
||||
def _extract_zip(zip_file: Path, log: Callable[[str], None]) -> None:
|
||||
for enc in ("utf-8", "gbk", "cp437"):
|
||||
try:
|
||||
with zipfile.ZipFile(zip_file, "r", metadata_encoding=enc) as zf:
|
||||
_extract_members(zf, zip_file.parent, log)
|
||||
return
|
||||
except (UnicodeDecodeError, zipfile.BadZipFile):
|
||||
continue
|
||||
raise RuntimeError("无法解压(编码检测失败)")
|
||||
|
||||
|
||||
def _extract_members(zf: zipfile.ZipFile, target_dir: Path, log: Callable[[str], None]) -> None:
|
||||
root = target_dir.resolve()
|
||||
for member in zf.infolist():
|
||||
name = member.filename.replace("\\", "/")
|
||||
target = (root / name).resolve()
|
||||
try:
|
||||
target.relative_to(root)
|
||||
except ValueError:
|
||||
log(f"[WARN] 跳过不安全的 ZIP 条目: {member.filename}")
|
||||
continue
|
||||
if member.is_dir():
|
||||
target.mkdir(parents=True, exist_ok=True)
|
||||
continue
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
with zf.open(member) as source, target.open("wb") as out:
|
||||
shutil.copyfileobj(source, out)
|
||||
|
||||
|
||||
def _clean_numeric_text(series: pd.Series) -> pd.Series:
|
||||
return (
|
||||
series.str.strip()
|
||||
.str.replace(",", "", regex=False)
|
||||
.str.replace(",", "", regex=False)
|
||||
.str.replace("%", "", regex=False)
|
||||
.str.replace("%", "", regex=False)
|
||||
.str.replace("\t", "", regex=False)
|
||||
.str.replace(" ", "", regex=False)
|
||||
)
|
||||
|
||||
|
||||
def _numeric_series(series: pd.Series) -> pd.Series:
|
||||
text = series.astype("string")
|
||||
has_percent = text.str.contains(r"[%%]", regex=True, na=False)
|
||||
cleaned = _clean_numeric_text(text)
|
||||
zero_mask = cleaned.isna() | cleaned.str.upper().isin(ZERO_TEXTS)
|
||||
numeric = pd.to_numeric(cleaned.mask(zero_mask, "0"), errors="coerce").fillna(0)
|
||||
numeric[has_percent & numeric.notna()] = numeric[has_percent & numeric.notna()] / 100
|
||||
return numeric
|
||||
|
||||
|
||||
def _convert_int(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
rounded = _numeric_series(series).round()
|
||||
return pd.Series([int(v) for v in rounded], index=series.index, dtype=object)
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
|
||||
|
||||
def _convert_float(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
numeric = _numeric_series(series)
|
||||
return pd.Series([float(v) for v in numeric], index=series.index, dtype=object)
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
|
||||
|
||||
def _convert_datetime(series: pd.Series) -> pd.Series:
|
||||
try:
|
||||
valid = series.notna() & (series != "") & (series.astype(str).str.strip() != "")
|
||||
if not valid.any():
|
||||
return pd.Series([None] * len(series), index=series.index)
|
||||
parsed = pd.Series([pd.NaT] * len(series), index=series.index)
|
||||
remaining = valid.copy()
|
||||
for fmt in DATETIME_FORMATS:
|
||||
if not remaining.any():
|
||||
break
|
||||
try:
|
||||
temp = pd.to_datetime(series[remaining], errors="coerce", format=fmt)
|
||||
except Exception: # noqa: BLE001
|
||||
continue
|
||||
ok = temp.notna()
|
||||
if ok.any():
|
||||
idx = remaining[remaining].index[ok]
|
||||
parsed.loc[idx] = temp[ok].values
|
||||
remaining.loc[idx] = False
|
||||
if remaining.any():
|
||||
try:
|
||||
temp = pd.to_datetime(series[remaining], errors="coerce", format="mixed", dayfirst=False)
|
||||
ok = temp.notna()
|
||||
if ok.any():
|
||||
idx = remaining[remaining].index[ok]
|
||||
parsed.loc[idx] = temp[ok].values
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return parsed.dt.strftime("%Y-%m-%d %H:%M:%S")
|
||||
except Exception: # noqa: BLE001
|
||||
return series
|
||||
@@ -12,7 +12,7 @@ from typing import Any
|
||||
from app.config import BASE_DIR
|
||||
|
||||
|
||||
DEFAULT_EXPIRES_ON = date(2026, 6, 20)
|
||||
DEFAULT_EXPIRES_ON = date(2026, 12, 30)
|
||||
EXTEND_DAYS = 30
|
||||
LICENSE_FILE = BASE_DIR / "license.dat"
|
||||
_SECRET = b"CapacityReport local license v1"
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
"""Metrix 仓库模式的处理流水线:CSV 处理 → 平台导入暂存表 → run-script(single_session) 跑报表 SQL。
|
||||
|
||||
仅当 warehouse_type == "metrix" 时使用;直连 MySQL 模式走原版 DataProcessor。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from app.config import SQL_SCRIPT, AppConfig, MetrixConfig
|
||||
from app.processor import ProcessLogger
|
||||
from app.services.csv_processor import CsvProcessor
|
||||
from app.services.platform import make_client
|
||||
|
||||
RJ_DIR_TO_TABLE = {
|
||||
"2.6RJGD": "2_6GRJGD",
|
||||
"2.6RJYD": "2_6GRJYD",
|
||||
"700RJGD": "700MRJGD",
|
||||
"700RJYD": "700MRJYD",
|
||||
}
|
||||
RESULT_TABLES = ["4G_结果表", "5G_结果表"]
|
||||
|
||||
|
||||
def validate_metrix(metrix: MetrixConfig) -> None:
|
||||
missing = []
|
||||
if not metrix.base_url:
|
||||
missing.append("平台地址")
|
||||
if not metrix.token:
|
||||
missing.append("API Token")
|
||||
if not metrix.database_conn_id:
|
||||
missing.append("数据库连接 ID")
|
||||
if missing:
|
||||
raise RuntimeError("Metrix 连接配置不完整: " + ", ".join(missing))
|
||||
|
||||
|
||||
def build_processor_config(app_config: AppConfig) -> dict:
|
||||
metrix = app_config.metrix.normalized()
|
||||
rj = app_config.rj_data.normalized()
|
||||
return {
|
||||
"recent_days": metrix.recent_days,
|
||||
"sheet_filter": list(app_config.sheet_filter),
|
||||
"data_dir_to_table": dict(metrix.data_dir_to_table),
|
||||
"extract_fields": app_config.extract_fields,
|
||||
"rj": {
|
||||
"enabled": rj.enabled,
|
||||
"weekly_directories": rj.weekly_directories,
|
||||
"dir_to_table": RJ_DIR_TO_TABLE,
|
||||
"table_field_mappings": rj.table_field_mappings,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def read_report_sql() -> str:
|
||||
if not SQL_SCRIPT.exists():
|
||||
return ""
|
||||
return SQL_SCRIPT.read_text(encoding="utf-8").strip()
|
||||
|
||||
|
||||
def run_report_sql(app_config: AppConfig, logger: ProcessLogger) -> list[dict]:
|
||||
metrix = app_config.metrix.normalized()
|
||||
validate_metrix(metrix)
|
||||
report_sql = read_report_sql()
|
||||
if not report_sql:
|
||||
raise RuntimeError("报表 SQL(ReportScript.sql)为空或不存在")
|
||||
client = make_client(metrix)
|
||||
logger.info("执行报表 SQL(single_session)...")
|
||||
result = client.run_script(
|
||||
metrix.database_conn_id,
|
||||
content=report_sql,
|
||||
database=metrix.target_database,
|
||||
single_session=True,
|
||||
run_timeout=7200,
|
||||
)
|
||||
statements = result.get("results", [])
|
||||
failed = [item for item in statements if not item.get("ok")]
|
||||
if result.get("stopped") or failed:
|
||||
for item in failed[:5]:
|
||||
logger.error(f"[SQL] 第 {item.get('index')} 条失败: {item.get('message')}")
|
||||
raise RuntimeError("报表 SQL 执行失败")
|
||||
logger.success(f"报表 SQL 执行完成,共 {len(statements)} 条语句")
|
||||
return statements
|
||||
|
||||
|
||||
def run_import_and_report(work_dir: Path, app_config: AppConfig, logger: ProcessLogger) -> dict:
|
||||
"""处理工作目录数据 → 平台导入暂存表 → 跑报表 SQL。失败抛 RuntimeError。"""
|
||||
metrix = app_config.metrix.normalized()
|
||||
validate_metrix(metrix)
|
||||
|
||||
logger.set_stage("converting")
|
||||
tables = CsvProcessor(work_dir, build_processor_config(app_config), logger.info).process()
|
||||
if not tables:
|
||||
raise RuntimeError("处理后没有产出任何暂存表数据")
|
||||
|
||||
client = make_client(metrix)
|
||||
conn_id = metrix.database_conn_id
|
||||
target_db = metrix.target_database
|
||||
|
||||
# 导入前 DROP 旧暂存表,让自动建表按当周实际列重建。
|
||||
logger.set_stage("importing")
|
||||
drop_sql = "".join(f"DROP TABLE IF EXISTS `{table}`;\n" for table in tables)
|
||||
drop_result = client.run_script(conn_id, content=drop_sql, database=target_db, run_timeout=600)
|
||||
if drop_result.get("stopped"):
|
||||
raise RuntimeError("清理旧暂存表失败")
|
||||
|
||||
for table, csv_path in tables.items():
|
||||
logger.info(f"导入暂存表 {table} ...")
|
||||
job_id = client.import_csv(conn_id, table, csv_path, mode="overwrite", database=target_db, create_table=True)
|
||||
job = client.wait_job(job_id)
|
||||
if job.get("status") != "success":
|
||||
raise RuntimeError(f"暂存表 {table} 导入失败: {job.get('error_code') or job.get('status')}")
|
||||
logger.success(f"暂存表 {table} 导入完成")
|
||||
|
||||
logger.set_stage("scripting")
|
||||
statements = run_report_sql(app_config, logger)
|
||||
return {"tables": list(tables.keys()), "statements": len(statements)}
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Metrix 平台集成:API 客户端 + 储存下载器。
|
||||
|
||||
当 source_type/warehouse_type 选 "metrix" 时,源数据走平台储存模块、数据仓库走平台数据库模块。
|
||||
连接信息(地址/token/storage_id/database_conn_id/target_database)来自 Configure.json 的 Metrix 段。
|
||||
储存下载器与 RemoteDataDownloader 接口一致,可被源工厂直接替换。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import date
|
||||
from pathlib import Path
|
||||
from typing import Iterable
|
||||
|
||||
import requests
|
||||
|
||||
from app.config import AppConfig, MetrixConfig
|
||||
from app.services.remote_download import RemoteDownloadResult, RemoteFileInfo
|
||||
from app.utils.file_dates import parse_file_date_range, select_recent_items_by_directory
|
||||
|
||||
|
||||
class PlatformClient:
|
||||
"""平台储存 + 数据库模块的最小 API 封装(Bearer Token 鉴权)。"""
|
||||
|
||||
def __init__(self, base_url: str, token: str, timeout: int = 60):
|
||||
if not base_url:
|
||||
raise ValueError("缺少平台地址,请在系统设置的 Metrix 连接中填写")
|
||||
if not token:
|
||||
raise ValueError("缺少平台 API Token,请在系统设置的 Metrix 连接中填写")
|
||||
self.base = base_url.rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.session = requests.Session()
|
||||
self.session.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
# --- 储存模块 --------------------------------------------------------
|
||||
def list_storage_files(self, storage_id: str, path: str = "/", recursive: bool = True) -> list[dict]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/storages/{storage_id}/files",
|
||||
params={"path": path, "recursive": "true" if recursive else "false"},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json().get("entries", [])
|
||||
|
||||
def download_storage_file(self, storage_id: str, path: str, dest: Path) -> None:
|
||||
dest.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self.session.get(
|
||||
f"{self.base}/api/storages/{storage_id}/download",
|
||||
params={"path": path},
|
||||
stream=True,
|
||||
timeout=self.timeout,
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
with dest.open("wb") as handle:
|
||||
for chunk in resp.iter_content(chunk_size=1024 * 64):
|
||||
if chunk:
|
||||
handle.write(chunk)
|
||||
|
||||
def batch_delete_storage(self, storage_id: str, paths: list[str]) -> int:
|
||||
deleted = 0
|
||||
for start in range(0, len(paths), 100):
|
||||
chunk = [p for p in paths[start:start + 100] if p]
|
||||
if not chunk:
|
||||
continue
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/storages/{storage_id}/batch-delete",
|
||||
json={"paths": chunk},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
deleted += len(chunk)
|
||||
return deleted
|
||||
|
||||
# --- 数据库模块 ------------------------------------------------------
|
||||
def import_csv(self, conn_id: str, table: str, csv_path: Path, mode: str = "overwrite",
|
||||
database: str = "", create_table: bool = True, upload_timeout: int = 1800) -> str:
|
||||
with csv_path.open("rb") as handle:
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/import",
|
||||
files={"file": (csv_path.name, handle, "text/csv")},
|
||||
data={
|
||||
"format": "csv",
|
||||
"target_table": table,
|
||||
"mode": mode,
|
||||
"database": database,
|
||||
"mapping": "{}",
|
||||
"create_table": "true" if create_table else "false",
|
||||
},
|
||||
timeout=upload_timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()["job_id"]
|
||||
|
||||
def wait_job(self, job_id: str, interval: int = 2, max_wait: int = 7200) -> dict:
|
||||
deadline = time.time() + max_wait
|
||||
while time.time() < deadline:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/database-transfer-jobs/{job_id}", timeout=self.timeout
|
||||
)
|
||||
resp.raise_for_status()
|
||||
job = resp.json()
|
||||
if job.get("status") in ("success", "failed"):
|
||||
return job
|
||||
time.sleep(interval)
|
||||
raise TimeoutError(f"导入任务 {job_id} 超过 {max_wait}s 仍未完成")
|
||||
|
||||
def run_script(self, conn_id: str, script_id: int | None = None, content: str = "",
|
||||
database: str = "", single_session: bool = False, run_timeout: int = 7200) -> dict:
|
||||
body: dict = {"database": database, "stop_on_error": True, "single_session": single_session}
|
||||
if content:
|
||||
body["content"] = content
|
||||
if script_id is not None:
|
||||
body["script_id"] = int(script_id)
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/run-script",
|
||||
json=body,
|
||||
timeout=run_timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
# --- 数据库读 / 导出(供仓库代理使用)-------------------------------
|
||||
def list_tables(self, conn_id: str, database: str = "") -> list[str]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/tables",
|
||||
params={"database": database},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return [str(item.get("name")) for item in resp.json() if item.get("name")]
|
||||
|
||||
def table_columns(self, conn_id: str, table: str, database: str = "") -> list[dict]:
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/tables/{table}",
|
||||
params={"database": database},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json().get("columns", [])
|
||||
|
||||
def table_data(self, conn_id: str, table: str, database: str = "", page: int = 1, page_size: int = 50,
|
||||
order_by: str = "", order_dir: str = "asc") -> dict:
|
||||
params = {"database": database, "table": table, "page": page, "page_size": page_size}
|
||||
if order_by:
|
||||
params["order_by"] = order_by
|
||||
params["order_dir"] = "desc" if str(order_dir).lower().startswith("desc") else "asc"
|
||||
resp = self.session.get(
|
||||
f"{self.base}/api/databases/{conn_id}/table-data", params=params, timeout=self.timeout
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
def submit_export(self, conn_id: str, tables: list[str], fmt: str, database: str = "") -> str:
|
||||
resp = self.session.post(
|
||||
f"{self.base}/api/databases/{conn_id}/export",
|
||||
json={"format": fmt, "database": database, "tables": tables},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()["job_id"]
|
||||
|
||||
def download_job_file(self, job_id: str, dest: Path) -> None:
|
||||
dest.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self.session.get(
|
||||
f"{self.base}/api/database-transfer-jobs/{job_id}/download", stream=True, timeout=self.timeout
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
with dest.open("wb") as handle:
|
||||
for chunk in resp.iter_content(chunk_size=1024 * 64):
|
||||
if chunk:
|
||||
handle.write(chunk)
|
||||
|
||||
|
||||
def make_client(metrix: MetrixConfig) -> PlatformClient:
|
||||
metrix = metrix.normalized()
|
||||
return PlatformClient(metrix.base_url, metrix.token)
|
||||
|
||||
|
||||
def make_source_downloader(app_config: AppConfig, logger=None):
|
||||
"""Return a file-source downloader matching app_config.source_type. FTP/SFTP use the
|
||||
original RemoteDataDownloader; 'metrix' uses PlatformStorageDownloader. Both share the
|
||||
interface test_connection / list_remote_zip_files / download_to / delete_source_files."""
|
||||
if app_config.source_type == "metrix":
|
||||
return PlatformStorageDownloader(app_config, logger)
|
||||
from app.services.remote_download import RemoteDataDownloader
|
||||
|
||||
return RemoteDataDownloader(app_config.remote_data, logger)
|
||||
|
||||
|
||||
class PlatformStorageDownloader:
|
||||
"""平台储存版下载器,接口与 RemoteDataDownloader 对齐,可被源工厂直接替换。"""
|
||||
|
||||
def __init__(self, app_config: AppConfig, logger=None):
|
||||
self.metrix = app_config.metrix.normalized()
|
||||
self.remote_dir = (app_config.remote_data.remote_dir or "/").strip() or "/"
|
||||
self.logger = logger
|
||||
self.client = make_client(self.metrix)
|
||||
|
||||
def _log(self, message: str) -> None:
|
||||
if self.logger:
|
||||
self.logger(message)
|
||||
|
||||
def test_connection(self) -> None:
|
||||
if not self.metrix.storage_id:
|
||||
raise ValueError("缺少储存连接 ID,请在系统设置的 Metrix 连接中填写")
|
||||
self.client.list_storage_files(self.metrix.storage_id, self.remote_dir, recursive=False)
|
||||
|
||||
def list_remote_zip_files(self, directory: str | None = None) -> list[RemoteFileInfo]:
|
||||
path = self._join(self.remote_dir, directory.strip("/")) if directory else self.remote_dir
|
||||
entries = self.client.list_storage_files(self.metrix.storage_id, path, recursive=True)
|
||||
files: list[RemoteFileInfo] = []
|
||||
for entry in entries:
|
||||
if entry.get("is_dir"):
|
||||
continue
|
||||
name = str(entry.get("name", ""))
|
||||
if not name.lower().endswith(".zip"):
|
||||
continue
|
||||
files.append(self._info(str(entry.get("path", "")), int(entry.get("size", 0) or 0)))
|
||||
return files
|
||||
|
||||
def download_to(self, destination: Path, target_dates: Iterable[date] | None = None) -> RemoteDownloadResult:
|
||||
destination = Path(destination)
|
||||
destination.mkdir(parents=True, exist_ok=True)
|
||||
zip_files = self.list_remote_zip_files()
|
||||
date_filter = set(target_dates or [])
|
||||
|
||||
if date_filter:
|
||||
selected = self._select_by_dates(zip_files, date_filter)
|
||||
elif zip_files:
|
||||
selected, summaries = select_recent_items_by_directory(
|
||||
zip_files,
|
||||
parent_key=lambda item: item.parent,
|
||||
name_key=lambda item: item.name,
|
||||
)
|
||||
for summary in summaries:
|
||||
if summary.skipped_count and summary.start_date and summary.max_date:
|
||||
self._log(
|
||||
f"储存目录 {summary.directory or '.'}: 仅下载 "
|
||||
f"{summary.start_date.isoformat()} 至 {summary.max_date.isoformat()} 的 "
|
||||
f"{summary.selected_count}/{summary.total_count} 个 ZIP,跳过 {summary.skipped_count} 个旧文件"
|
||||
)
|
||||
else:
|
||||
selected = []
|
||||
|
||||
result = RemoteDownloadResult()
|
||||
for remote_file in selected:
|
||||
dest = destination / remote_file.relative_path
|
||||
self._log(f"下载: {remote_file.relative_path}")
|
||||
self.client.download_storage_file(self.metrix.storage_id, remote_file.path, dest)
|
||||
result.file_count += 1
|
||||
result.total_bytes += remote_file.size or (dest.stat().st_size if dest.exists() else 0)
|
||||
result.remote_files.append(remote_file.path)
|
||||
return result
|
||||
|
||||
def delete_source_files(self, remote_files: Iterable[str] | None = None) -> int:
|
||||
files = [path for path in (remote_files or []) if path]
|
||||
if not files:
|
||||
return 0
|
||||
self._log(f"清理储存源文件,共 {len(files)} 个")
|
||||
return self.client.batch_delete_storage(self.metrix.storage_id, files)
|
||||
|
||||
# --- helpers ---------------------------------------------------------
|
||||
def _select_by_dates(self, zip_files: list[RemoteFileInfo], target_dates: set[date]) -> list[RemoteFileInfo]:
|
||||
grouped: dict[str, list[RemoteFileInfo]] = {}
|
||||
for remote_file in zip_files:
|
||||
grouped.setdefault(remote_file.parent, []).append(remote_file)
|
||||
|
||||
selected: list[RemoteFileInfo] = []
|
||||
for parent, files in sorted(grouped.items(), key=lambda item: item[0]):
|
||||
picked = [
|
||||
remote_file
|
||||
for remote_file in files
|
||||
if (date_range := parse_file_date_range(remote_file.name))
|
||||
and (
|
||||
date_range.covers_all(target_dates)
|
||||
if date_range.span_days > 1
|
||||
else date_range.covers_any(target_dates)
|
||||
)
|
||||
]
|
||||
selected.extend(picked)
|
||||
skipped = len(files) - len(picked)
|
||||
if skipped:
|
||||
self._log(
|
||||
f"储存目录 {parent or '.'}: 仅下载目标日期 "
|
||||
f"{min(target_dates).isoformat()} 至 {max(target_dates).isoformat()} 的 "
|
||||
f"{len(picked)}/{len(files)} 个 ZIP,跳过 {skipped} 个非目标文件"
|
||||
)
|
||||
return selected
|
||||
|
||||
@staticmethod
|
||||
def _join(parent: str, child: str) -> str:
|
||||
parent = (parent or "").replace("\\", "/").rstrip("/")
|
||||
if not parent:
|
||||
return child
|
||||
if parent == "/":
|
||||
return f"/{child}"
|
||||
return f"{parent}/{child}"
|
||||
|
||||
def _info(self, remote_path: str, size: int = 0) -> RemoteFileInfo:
|
||||
normalized_root = self.remote_dir.replace("\\", "/").rstrip("/")
|
||||
normalized_path = remote_path.replace("\\", "/")
|
||||
if normalized_root and normalized_root != "/" and normalized_path.startswith(f"{normalized_root}/"):
|
||||
relative_path = normalized_path[len(normalized_root) + 1:]
|
||||
else:
|
||||
relative_path = normalized_path.lstrip("/")
|
||||
relative = Path(relative_path)
|
||||
parent = str(relative.parent).replace("\\", "/")
|
||||
if parent == ".":
|
||||
parent = ""
|
||||
return RemoteFileInfo(
|
||||
path=remote_path,
|
||||
relative_path=relative_path,
|
||||
parent=parent,
|
||||
name=relative.name,
|
||||
size=size,
|
||||
)
|
||||
Reference in New Issue
Block a user