feat: 导入 CapacityReport 初始源码

This commit is contained in:
Nixevol
2026-09-24 06:16:41 +08:00
commit 89cca70430
134 changed files with 38799 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
"""Service helpers."""
+669
View File
@@ -0,0 +1,669 @@
from __future__ import annotations
import json
import threading
from collections import defaultdict
from dataclasses import dataclass
from datetime import date, datetime, timedelta
from typing import Any
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
READY_DIR = CACHE_DIR / "auto_scheduler"
READY_FLAG = READY_DIR / "ready.flag"
DISABLED_CHECK_SECONDS = 60
STARTUP_CHECK_SECONDS = 5
FAILURE_RESULTS = {"scan_failed", "trigger_failed", "source_cleanup_failed", "failed", "invalid_flag"}
@dataclass(frozen=True)
class AutoDirectoryReadyStatus:
"""自动粒度目录就绪状态,按最新文件自动识别日粒度或周粒度。"""
directory: str
ready: bool
granularity: str | None = None
found_days: list[date] | None = None
missing_days: list[date] | None = None
file_name: str | None = None
file_count: int = 0
error: str | None = None
skipped: bool = False
skip_reason: str | None = None
def to_dict(self) -> dict[str, Any]:
found_days = self.found_days or []
missing_days = self.missing_days or []
return {
"ready": self.ready,
"granularity": self.granularity,
"found_days": [item.isoformat() for item in found_days],
"missing_days": [item.isoformat() for item in missing_days],
"found_count": len(found_days),
"required_count": len(found_days) + len(missing_days),
"file_name": self.file_name,
"file_count": self.file_count,
"error": self.error,
"skipped": self.skipped,
"skip_reason": self.skip_reason,
}
@dataclass(frozen=True)
class DirectoryReadyStatus:
directory: str
ready: bool
found_days: list[date]
missing_days: list[date]
file_count: int
error: str | None = None
skipped: bool = False
skip_reason: str | None = None
def to_dict(self) -> dict[str, Any]:
return {
"ready": self.ready,
"found_days": [item.isoformat() for item in self.found_days],
"missing_days": [item.isoformat() for item in self.missing_days],
"found_count": len(self.found_days),
"required_count": len(self.found_days) + len(self.missing_days),
"file_count": self.file_count,
"error": self.error,
"skipped": self.skipped,
"skip_reason": self.skip_reason,
}
class AutoScheduler:
def __init__(self) -> None:
self._stop_event = threading.Event()
self._check_lock = threading.Lock()
self._status_lock = threading.Lock()
self._thread: threading.Thread | None = None
self._running = False
self._last_check_at: datetime | None = None
self._next_check_at: datetime | None = None
self._last_result = "not_started"
self._last_message = "自动调度器尚未检查"
self._directory_status: dict[str, dict[str, Any]] = {}
self._task_running = False
self._task_id: str | None = None
self._failure_count = 0
def start(self) -> None:
if self._thread and self._thread.is_alive():
return
self._running = True
self._stop_event.clear()
self._set_next_check(datetime.now() + timedelta(seconds=STARTUP_CHECK_SECONDS))
self._thread = threading.Thread(target=self._run_loop, name="auto-scheduler", daemon=True)
self._thread.start()
def stop(self) -> None:
self._running = False
self._stop_event.set()
if self._thread and self._thread.is_alive():
self._thread.join(timeout=5)
def get_status(self) -> dict[str, Any]:
app_config = state.current_config()
config = app_config.remote_data.normalized()
scheduler = config.auto_scheduler.normalized()
data_mappings = app_config.data_mappings.normalized()
auto_directories = [
item["path"]
for item in data_mappings.directories
if item.get("ready_rule") == "auto"
]
target_days = required_week_days(scheduler.week_offset)
ready_flag = self._read_ready_flag()
with self._status_lock:
return {
"enabled": scheduler.enabled,
"running": self._running,
"check_interval_hours": scheduler.check_interval_hours,
"expected_directories": scheduler.expected_directories,
"week_offset": scheduler.week_offset,
"auto_delete_source": config.auto_delete_source,
"auto_ready_directories": auto_directories,
"next_check_at": self._format_dt(self._next_check_at),
"last_check_at": self._format_dt(self._last_check_at),
"last_result": self._last_result,
"last_message": self._last_message,
"failure_count": self._failure_count,
"task_running": self._task_running,
"task_id": self._task_id,
"ready_flag": ready_flag,
"target_week": {
"start": target_days[0].isoformat(),
"end": target_days[-1].isoformat(),
"days": [item.isoformat() for item in target_days],
},
"directory_status": self._directory_status,
}
def check_and_run(self, manual: bool = False) -> dict[str, Any]:
if not self._check_lock.acquire(blocking=False):
return self._finish_check("busy", "自动调度器正在检查中", manual)
try:
return self._check_and_run_locked(manual)
finally:
self._check_lock.release()
def _run_loop(self) -> None:
while not self._stop_event.wait(self._seconds_until_next_check()):
self.check_and_run(manual=False)
def _check_and_run_locked(self, manual: bool) -> dict[str, Any]:
now = datetime.now()
app_config = state.current_config()
remote_config = app_config.remote_data.normalized()
scheduler = remote_config.auto_scheduler.normalized()
self._last_check_at = now
if not scheduler.enabled:
self._set_next_check(now + timedelta(seconds=DISABLED_CHECK_SECONDS))
return self._finish_check("disabled", "自动调度未启用", manual)
self._set_next_check(now + timedelta(hours=scheduler.check_interval_hours))
if not remote_config.enabled:
return self._finish_check("remote_disabled", "远程数据源未启用", manual)
if state.global_task_lock["locked"]:
return self._finish_check("task_running", "已有任务在运行,本轮自动调度跳过", manual)
ready_flag = self._read_ready_flag()
if ready_flag["exists"]:
if ready_flag.get("invalid"):
self._clear_ready_flag()
return self._finish_check("invalid_flag", "就绪标识格式错误,已清除,本轮不触发处理", manual)
target_dates = self._target_dates_from_flag(ready_flag, scheduler)
return self._trigger_processing(target_dates, ready_flag, manual)
target_days = required_week_days(scheduler.week_offset)
downloader = make_source_downloader(app_config)
data_mappings = app_config.data_mappings.normalized()
auto_directories = {
item["path"]
for item in data_mappings.directories
if item.get("ready_rule") == "auto"
}
# 自动粒度目录单独判断,避免被普通 7 天规则误拦。
directory_status = self._check_remote_ready(
downloader,
scheduler,
target_days,
excluded_directories=auto_directories,
)
error_count = sum(1 for item in directory_status.values() if item.error)
if error_count:
self._set_combined_directory_status(directory_status, {})
return self._finish_check(
"scan_failed",
f"远程目录扫描失败,{error_count}/{len(directory_status)} 个目录无法访问或扫描失败",
manual,
)
active_status = [item for item in directory_status.values() if not item.skipped]
skipped_count = len(directory_status) - len(active_status)
# 检查自动粒度目录,按最新文件自动判断日粒度或周粒度。
auto_status = self._check_auto_ready(downloader, target_days, sorted(auto_directories))
auto_error_count = sum(1 for item in auto_status.values() if item.error)
self._set_combined_directory_status(directory_status, auto_status)
if auto_error_count:
return self._finish_check(
"scan_failed",
f"自动粒度远程目录扫描失败,{auto_error_count}/{len(auto_status)} 个目录无法访问或扫描失败",
manual,
)
active_auto_status = [item for item in auto_status.values() if not item.skipped]
daily_ready = all(item.ready for item in active_status) if active_status else not directory_status and bool(active_auto_status)
auto_ready = all(item.ready for item in active_auto_status)
auto_status_text = ""
if active_auto_status:
auto_ready_count = sum(1 for item in active_auto_status if item.ready)
auto_total = len(active_auto_status)
if not auto_ready:
auto_status_text = f",自动粒度数据 {auto_ready_count}/{auto_total} 个有效目录就绪"
# 两个条件都满足才触发
if daily_ready and auto_ready:
self._mark_ready(target_days, directory_status, auto_status)
skipped_text = f",已跳过 {skipped_count} 个停推目录" if skipped_count else ""
auto_text = ",自动粒度数据已就绪" if active_auto_status else ""
return self._finish_check(
"marked_ready",
f"远程数据已满足目标周 7 天{skipped_text}{auto_text},已写入就绪标识,下次检查将自动处理",
manual,
)
if not active_status and directory_status:
return self._finish_check(
"waiting",
f"远程数据未就绪,{len(directory_status)} 个目录均为空,已视为停推但不会触发处理",
manual,
)
if not active_status:
return self._finish_check(
"waiting",
f"远程普通日数据未发现有效目录,无法触发处理{auto_status_text}",
manual,
)
ready_count = sum(1 for item in active_status if item.ready)
skipped_text = f",跳过 {skipped_count} 个停推目录" if skipped_count else ""
return self._finish_check(
"waiting",
f"远程数据未就绪,{ready_count}/{len(active_status)} 个有效目录满足目标周 7 天{skipped_text}{auto_status_text}",
manual,
)
def _check_remote_ready(
self,
downloader: RemoteDataDownloader,
scheduler: AutoSchedulerConfig,
target_days: list[date],
excluded_directories: set[str] | None = None,
) -> dict[str, DirectoryReadyStatus]:
excluded = {self._normalize_directory_name(item) for item in (excluded_directories or set())}
expected_directories = scheduler.expected_directories
if expected_directories:
return {
directory: self._directory_ready_status(
directory,
self._safe_list_remote_zip_files(downloader, directory),
target_days,
)
for directory in expected_directories
if not self._is_excluded_directory(directory, excluded)
}
files = self._safe_list_remote_zip_files(downloader, None)
if files is None:
return {
".": DirectoryReadyStatus(
directory=".",
ready=False,
found_days=[],
missing_days=target_days,
file_count=0,
error="远程目录扫描失败",
)
}
if not files:
return {
".": DirectoryReadyStatus(
directory=".",
ready=False,
found_days=[],
missing_days=target_days,
file_count=0,
error="远程目录未找到 ZIP 文件",
)
}
grouped: dict[str, list[RemoteFileInfo]] = defaultdict(list)
for remote_file in files:
parent = remote_file.parent or "."
if self._is_excluded_directory(parent, excluded):
continue
grouped[parent].append(remote_file)
if not grouped:
return {}
return {
directory: self._directory_ready_status(directory, directory_files, target_days)
for directory, directory_files in sorted(grouped.items(), key=lambda item: item[0])
}
def _safe_list_remote_zip_files(
self,
downloader: RemoteDataDownloader,
directory: str | None,
) -> list[RemoteFileInfo] | None:
try:
return downloader.list_remote_zip_files(directory)
except Exception:
return None
def _directory_ready_status(
self,
directory: str,
files: list[RemoteFileInfo] | None,
target_days: list[date],
) -> DirectoryReadyStatus:
if files is None:
return DirectoryReadyStatus(
directory=directory,
ready=False,
found_days=[],
missing_days=target_days,
file_count=0,
error="远程目录不存在或无法访问",
)
if not files:
return DirectoryReadyStatus(
directory=directory,
ready=True,
found_days=[],
missing_days=[],
file_count=0,
skipped=True,
skip_reason="目录为空,视为已停推并跳过",
)
required = set(target_days)
found = {
target_day
for remote_file in files
if (date_range := parse_file_date_range(remote_file.name))
for target_day in date_range.covered_days()
if target_day in required
}
missing = [item for item in target_days if item not in found]
return DirectoryReadyStatus(
directory=directory,
ready=not missing,
found_days=sorted(found),
missing_days=missing,
file_count=len(files),
)
def _check_auto_ready(
self,
downloader: RemoteDataDownloader,
target_days: list[date],
directories: list[str],
) -> dict[str, AutoDirectoryReadyStatus]:
"""检查自动粒度目录是否就绪,自动识别目录最新文件是日粒度还是周粒度。"""
if not directories:
return {}
result: dict[str, AutoDirectoryReadyStatus] = {}
for directory in directories:
result[directory] = self._check_single_auto_directory(downloader, directory, target_days)
return result
def _check_single_auto_directory(
self,
downloader: RemoteDataDownloader,
directory: str,
target_days: list[date],
) -> AutoDirectoryReadyStatus:
"""检查单个自动粒度目录是否包含目标周数据。"""
files = self._safe_list_remote_zip_files(downloader, directory)
if files is None:
return AutoDirectoryReadyStatus(
directory=directory,
ready=False,
error="远程目录不存在或无法访问",
)
if not files:
return AutoDirectoryReadyStatus(
directory=directory,
ready=True,
file_count=0,
skipped=True,
skip_reason="目录为空,视为已停推并跳过",
)
parsed_files = [
(remote_file, date_range)
for remote_file in files
if (date_range := parse_file_date_range(remote_file.name))
]
if not parsed_files:
return AutoDirectoryReadyStatus(
directory=directory,
ready=False,
found_days=[],
missing_days=target_days,
file_count=len(files),
error="目录中未找到可识别日期的 ZIP 文件",
)
latest_file, latest_range = max(
parsed_files,
key=lambda item: (item[1].start, item[1].end_exclusive, item[0].name),
)
granularity = "daily" if latest_range.span_days <= 1 else "weekly"
required = set(target_days)
found = {
target_day
for _, date_range in parsed_files
for target_day in date_range.covered_days()
if target_day in required
}
missing = [item for item in target_days if item not in found]
if granularity == "daily":
return AutoDirectoryReadyStatus(
directory=directory,
ready=not missing,
granularity=granularity,
found_days=sorted(found),
missing_days=missing,
file_name=latest_file.name,
file_count=len(files),
)
for remote_file in files:
date_range = parse_file_date_range(remote_file.name)
if date_range and date_range.covers_all(required):
return AutoDirectoryReadyStatus(
directory=directory,
ready=True,
granularity=granularity,
found_days=target_days,
missing_days=[],
file_name=remote_file.name,
file_count=len(files),
)
return AutoDirectoryReadyStatus(
directory=directory,
ready=False,
granularity=granularity,
found_days=sorted(found),
missing_days=missing,
file_name=latest_file.name,
file_count=len(files),
)
def _trigger_processing(
self,
target_dates: list[date],
ready_flag: dict[str, Any],
manual: bool,
) -> dict[str, Any]:
from app.api.routers.remote import start_remote_processing_job
try:
result = start_remote_processing_job(
source="scheduler",
on_finish=self._on_processing_finish,
target_dates=target_dates,
)
except HTTPException as exc:
return self._finish_check("trigger_failed", str(exc.detail), manual)
except Exception as exc:
return self._finish_check("trigger_failed", f"自动调度触发失败: {exc}", manual)
task_id = result.get("task_id")
with self._status_lock:
self._task_running = True
self._task_id = str(task_id) if task_id else None
self._last_result = "triggered"
self._last_message = "已根据就绪标识触发远程下载并处理"
return {
"success": True,
"result": "triggered",
"message": "已根据就绪标识触发远程下载并处理",
"task_id": task_id,
"ready_flag": ready_flag,
"status": self.get_status(),
}
def _mark_ready(
self,
target_days: list[date],
directory_status: dict[str, DirectoryReadyStatus],
auto_status: dict[str, AutoDirectoryReadyStatus] | None = None,
) -> None:
READY_DIR.mkdir(parents=True, exist_ok=True)
payload = {
"ready_at": datetime.now().isoformat(timespec="seconds"),
"week_start": target_days[0].isoformat(),
"week_end": target_days[-1].isoformat(),
"target_dates": [item.isoformat() for item in target_days],
"directories": {
name: item.to_dict()
for name, item in directory_status.items()
},
}
if auto_status:
payload["auto_directories"] = {
name: item.to_dict()
for name, item in auto_status.items()
}
READY_FLAG.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
def _clear_ready_flag(self) -> None:
try:
READY_FLAG.unlink(missing_ok=True)
except OSError:
pass
def _read_ready_flag(self) -> dict[str, Any]:
if not READY_FLAG.exists():
return {"exists": False}
try:
data = json.loads(READY_FLAG.read_text(encoding="utf-8"))
if isinstance(data, dict):
return {"exists": True, **data}
except Exception as exc:
return {"exists": True, "invalid": True, "error": str(exc)}
return {"exists": True, "invalid": True, "error": "ready.flag 格式错误"}
def _target_dates_from_flag(
self,
ready_flag: dict[str, Any],
scheduler: AutoSchedulerConfig,
) -> list[date]:
raw_dates = ready_flag.get("target_dates")
if isinstance(raw_dates, list):
dates: list[date] = []
for raw_date in raw_dates:
try:
dates.append(date.fromisoformat(str(raw_date)))
except ValueError:
continue
if dates:
return sorted(dates)
return required_week_days(scheduler.week_offset)
def _on_processing_finish(self, task_id: str, status: str) -> None:
if status == "completed":
self._clear_ready_flag()
message = "自动调度任务处理成功,已清除就绪标识"
result = "completed"
else:
message = "自动调度任务未完成,就绪标识已保留,下次检查会重试"
result = status or "failed"
with self._status_lock:
self._task_running = False
self._task_id = task_id
self._last_result = result
self._last_message = message
self._failure_count = 0 if result == "completed" else self._failure_count + 1
def _finish_check(self, result: str, message: str, manual: bool) -> dict[str, Any]:
with self._status_lock:
self._last_result = result
self._last_message = message
if result in FAILURE_RESULTS:
self._failure_count += 1
elif result not in {"busy", "task_running"}:
self._failure_count = 0
if result not in {"triggered", "task_running"}:
self._task_running = False
return {
"success": result not in FAILURE_RESULTS,
"manual": manual,
"result": result,
"message": message,
"status": self.get_status(),
}
def _set_combined_directory_status(
self,
directory_status: dict[str, DirectoryReadyStatus],
auto_status: dict[str, AutoDirectoryReadyStatus],
) -> None:
with self._status_lock:
combined = {
name: item.to_dict()
for name, item in directory_status.items()
}
combined.update(
{
name: item.to_dict()
for name, item in auto_status.items()
}
)
self._directory_status = combined
@staticmethod
def _normalize_directory_name(directory: str) -> str:
normalized = str(directory or "").replace("\\", "/").strip().strip("/")
return "." if normalized in {"", "."} else normalized
@classmethod
def _is_excluded_directory(cls, directory: str, excluded: set[str]) -> bool:
if not excluded:
return False
normalized = cls._normalize_directory_name(directory)
return any(
normalized == item or normalized.startswith(f"{item}/")
for item in excluded
if item != "."
)
def _set_next_check(self, value: datetime) -> None:
with self._status_lock:
self._next_check_at = value
def _seconds_until_next_check(self) -> float:
with self._status_lock:
next_check_at = self._next_check_at
if not next_check_at:
return STARTUP_CHECK_SECONDS
return max((next_check_at - datetime.now()).total_seconds(), 0)
@staticmethod
def _format_dt(value: datetime | None) -> str | None:
return value.isoformat(timespec="seconds") if value else None
+648
View File
@@ -0,0 +1,648 @@
from __future__ import annotations
import csv
import re
import stat
import zipfile
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any, Callable, Iterable
import chardet
import pymysql
from app.config import AppConfig, CellDataConfig
from app.services.remote_download import RemoteDataDownloader, RemoteFileInfo
LogFn = Callable[[str], None]
CELLINFO_COLUMNS = [
"CGI",
"eNodeBID",
"CellID",
"PLMN",
"基站名称",
"小区名称",
"频点",
"带宽",
"制式",
"功率",
"网络",
]
@dataclass(frozen=True)
class SelectedZip:
scan_path: str
band: str
remote_file: RemoteFileInfo
timestamp: str
@dataclass
class CellDataResult:
selected_files: int = 0
parsed_rows: int = 0
imported_rows: int = 0
skipped_rows: int = 0
class CellDataProcessor:
def __init__(self, app_config: AppConfig, work_dir: Path, logger=None):
self.app_config = app_config
self.config = app_config.cell_data.normalized()
self.work_dir = work_dir
self.logger = logger
self.downloader = RemoteDataDownloader(self.config.remote_data, self._log)
self.year_dir_re = re.compile(self.config.year_dir_regex)
self.month_dir_re = re.compile(self.config.month_dir_regex)
self.day_dir_re = re.compile(self.config.day_dir_regex)
self.file_name_re = re.compile(self.config.file_name_regex)
self.file_time_re = re.compile(self.config.file_time_regex)
def run(self) -> CellDataResult:
if not self.config.remote_data.enabled:
self._log("CellData 数据源未启用,跳过")
return CellDataResult()
self._set_stage("locating")
selected = self._select_latest_zip_files()
if not selected:
raise RuntimeError("未找到可处理的 CellData ZIP 文件")
self._log(f"已选择 {len(selected)} 个 CellData ZIP 文件")
self._set_stage("downloading")
local_files = self._download_selected(selected)
self._set_stage("parsing")
result = CellDataResult(selected_files=len(selected))
rows = self._parse_zip_files(local_files, result)
if not rows:
raise RuntimeError("CellData ZIP 中未解析到有效数据")
self._set_stage("importing")
result.imported_rows = self._replace_cellinfo(rows)
self._log(f"CellData 导入完成,共 {result.imported_rows} 行")
return result
def run_local(self, upload_root: Path) -> CellDataResult:
self._set_stage("locating")
local_files = self._select_local_zip_files(upload_root)
if not local_files:
raise RuntimeError("未找到可处理的 CellData ZIP 文件")
self._log(f"已选择 {len(local_files)} 个 CellData ZIP 文件")
self._set_stage("parsing")
result = CellDataResult(selected_files=len(local_files))
rows = self._parse_zip_files(local_files, result)
if not rows:
raise RuntimeError("CellData ZIP 中未解析到有效数据")
self._set_stage("importing")
result.imported_rows = self._replace_cellinfo(rows)
self._log(f"CellData 导入完成,共 {result.imported_rows} 行")
return result
def _select_latest_zip_files(self) -> list[SelectedZip]:
selected: list[SelectedZip] = []
for template in self.config.scan_paths:
scan_path = self._resolve_scan_path(template)
band_dirs = [
entry
for entry in self._list_dir(scan_path)
if entry["type"] == "directory"
]
for band_dir in band_dirs:
band = str(band_dir["name"])
files = [
entry
for entry in self._list_dir(str(band_dir["path"]))
if entry["type"] == "file" and self.file_name_re.search(str(entry["name"]))
]
latest = self._latest_file(files)
if latest:
selected.append(
SelectedZip(
scan_path=scan_path,
band=band,
remote_file=self._remote_file_info(str(latest["path"]), int(latest.get("size") or 0)),
timestamp=str(latest["timestamp"]),
)
)
self._log(f"{band}: {latest['name']}")
return selected
def _resolve_scan_path(self, template: str) -> str:
path = self._replace_date_placeholders(str(template).replace("\\", "/").strip())
if not path.startswith("/"):
path = self._join_remote_path(self.config.remote_data.remote_dir, path)
parts = [part for part in path.split("/") if part]
replacements = [
("{maxyear}", "year", self.year_dir_re, "年份"),
("{maxmonth}", "month", self.month_dir_re, "月份"),
("{maxday}", "day", self.day_dir_re, "日期"),
]
for token, group_name, pattern, label in replacements:
token_index = next((index for index, part in enumerate(parts) if token in part), -1)
if token_index < 0:
continue
parent = "/" + "/".join(parts[:token_index]) if token_index else "/"
values: list[int] = []
for entry in self._list_dir(parent):
if entry["type"] != "directory":
continue
match = pattern.search(str(entry["name"]))
if match:
values.append(int(match.group(group_name)))
if not values:
raise RuntimeError(f"未找到{label}目录: {parent}")
parts[token_index] = parts[token_index].replace(token, str(max(values)))
return "/" + "/".join(parts)
def _replace_date_placeholders(self, path: str) -> str:
now = datetime.now()
return (
path.replace("{yyyy}", now.strftime("%Y"))
.replace("{yyyymm}", now.strftime("%Y%m"))
.replace("{yyyymmdd}", now.strftime("%Y%m%d"))
)
def _latest_file(self, files: list[dict[str, Any]]) -> dict[str, Any] | None:
candidates: list[dict[str, Any]] = []
for item in files:
match = self.file_time_re.search(str(item["name"]))
if not match:
continue
timestamp = match.group("timestamp")
candidates.append({**item, "timestamp": timestamp})
if not candidates:
return None
return max(candidates, key=lambda item: str(item["timestamp"]))
def _download_selected(self, selected: list[SelectedZip]) -> list[tuple[SelectedZip, Path]]:
target_dir = self.work_dir / "cell_data"
target_dir.mkdir(parents=True, exist_ok=True)
local_files: list[tuple[SelectedZip, Path]] = []
if self.config.remote_data.protocol == "ftp":
with self.downloader._ftp_client() as ftp: # noqa: SLF001 - reuse existing connection helpers
for item in selected:
local_path = target_dir / item.band / item.remote_file.name
self.downloader._download_ftp_file(ftp, item.remote_file.path, local_path, self._download_result(), item.remote_file.size) # noqa: SLF001
local_files.append((item, local_path))
return local_files
ssh = self.downloader._sftp_ssh_client() # noqa: SLF001
try:
with ssh.open_sftp() as sftp:
for item in selected:
local_path = target_dir / item.band / item.remote_file.name
self.downloader._download_sftp_file(sftp, item.remote_file.path, local_path, self._download_result(), item.remote_file.size) # noqa: SLF001
local_files.append((item, local_path))
finally:
ssh.close()
return local_files
def _select_local_zip_files(self, upload_root: Path) -> list[tuple[SelectedZip, Path]]:
grouped: dict[str, list[Path]] = {}
for path in upload_root.rglob("*.zip"):
if not self.file_name_re.search(path.name):
continue
try:
parent = str(path.parent.relative_to(upload_root)).replace("\\", "/")
except ValueError:
parent = ""
grouped.setdefault("" if parent == "." else parent, []).append(path)
selected: list[tuple[SelectedZip, Path]] = []
for parent, paths in sorted(grouped.items(), key=lambda item: item[0]):
candidates = []
for path in paths:
match = self.file_time_re.search(path.name)
if match:
candidates.append((match.group("timestamp"), path))
if not candidates:
continue
timestamp, path = max(candidates, key=lambda item: item[0])
band = Path(parent).name if parent else ""
selected_zip = SelectedZip(
scan_path=str(upload_root),
band=band,
remote_file=RemoteFileInfo(
path=str(path),
relative_path=str(path.relative_to(upload_root)).replace("\\", "/"),
parent=parent,
name=path.name,
size=path.stat().st_size,
),
timestamp=timestamp,
)
if not band:
self._log(f"未从目录名识别频段: {path.name}")
else:
self._log(f"{band}: {path.name}")
selected.append((selected_zip, path))
return selected
@staticmethod
def _download_result():
from app.services.remote_download import RemoteDownloadResult
return RemoteDownloadResult()
def _parse_zip_files(self, local_files: list[tuple[SelectedZip, Path]], result: CellDataResult) -> list[dict[str, str]]:
mapping = self.config.mapping
key_config = mapping["key"]
key_field = str(key_config["field"])
rows_by_key: dict[str, dict[str, str]] = {}
self._log(f"开始解压并解析 {len(local_files)} 个 ZIP 文件...")
for selected, local_path in local_files:
band_label = selected.band or "未识别频段"
size_kb = local_path.stat().st_size / 1024 if local_path.exists() else 0
self._log(f"[{band_label}] 解压 {local_path.name}({size_kb:.0f} KB)...")
zip_parsed_before = result.parsed_rows
zip_skipped_before = result.skipped_rows
with zipfile.ZipFile(local_path) as zf:
sources = list(mapping["sources"])
csv_entries = [info for info in zf.infolist() if Path(info.filename).name.lower().endswith(".csv")]
self._log(f" 压缩包内含 {len(csv_entries)} 个 CSV 文件")
for info in csv_entries:
name = Path(info.filename).name
matching_sources = [
source
for source in sources
if name.startswith(source["file_prefix"]) and (not selected.band or source["band"] == selected.band)
]
if not matching_sources:
self._log(f" 跳过未匹配规则的文件: {name}")
continue
if not selected.band and len(matching_sources) > 1:
self._log(f" 跳过无法识别频段的文件: {name}")
continue
for source in matching_sources:
raw = zf.read(info.filename)
text = self._decode_csv(raw)
reader = csv.DictReader(text.splitlines())
added = 0
skipped = 0
for csv_row in reader:
row = self._map_row(source["fields"], csv_row)
key = self._render_expr(str(key_config["expr"]), row)
if not key or "--" in key:
result.skipped_rows += 1
skipped += 1
continue
row[key_field] = key
rows_by_key[key] = {column: row.get(column, "") for column in CELLINFO_COLUMNS}
result.parsed_rows += 1
added += 1
self._log(f" 解析 {name}(频段 {source.get('band', '') or '通用'}):有效 {added} 行,跳过 {skipped} 行")
self._log(
f"[{band_label}] {local_path.name} 解析完成:"
f"本包有效 {result.parsed_rows - zip_parsed_before} 行,跳过 {result.skipped_rows - zip_skipped_before} 行"
)
self._log(
f"全部解析完成:累计有效 {result.parsed_rows} 行,按 {key_field} 去重后 {len(rows_by_key)} 行,"
f"累计跳过 {result.skipped_rows} 行"
)
return list(rows_by_key.values())
def _map_row(self, fields: dict[str, Any], csv_row: dict[str, str]) -> dict[str, str]:
row: dict[str, str] = {}
for target, rule in fields.items():
if isinstance(rule, str):
row[target] = str(csv_row.get(rule, "") or "").strip()
elif isinstance(rule, dict) and "value" in rule:
row[target] = str(rule.get("value", "") or "").strip()
else:
row[target] = ""
return row
@staticmethod
def _render_expr(expr: str, row: dict[str, str]) -> str:
def replace(match):
return row.get(match.group(1), "")
return re.sub(r"\{([^{}]+)\}", replace, expr).strip()
@staticmethod
def _decode_csv(raw: bytes) -> str:
candidates = ["utf-8-sig", "utf-8", "gb18030", "gbk"]
detected = (chardet.detect(raw).get("encoding") or "").lower()
if detected and detected not in candidates:
candidates.append(detected)
for encoding in dict.fromkeys(candidates):
try:
return raw.decode(encoding).lstrip("\ufeff")
except (LookupError, UnicodeDecodeError):
pass
return raw.decode("utf-8", errors="replace").lstrip("\ufeff")
def _replace_cellinfo(self, rows: list[dict[str, str]]) -> int:
mysql = self.config.mysql.normalized()
table = str(self.config.mapping.get("target_table") or "cellinfo")
self._log(f"准备写入表 `{table}`(库 {mysql.dbname}@{mysql.host}:{mysql.port}),共 {len(rows)} 行")
conn = pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
database=mysql.dbname,
charset="utf8mb4",
cursorclass=pymysql.cursors.DictCursor,
autocommit=False,
)
self._log("已连接 CellData 数据库")
try:
with conn.cursor() as cursor:
self._ensure_cellinfo_table(cursor, table)
self._log(f"已确认表结构 `{table}`")
cursor.execute(f"TRUNCATE TABLE `{table}`")
self._log(f"已清空表 `{table}`(TRUNCATE)")
placeholders = ", ".join(["%s"] * len(CELLINFO_COLUMNS))
columns = ", ".join(f"`{column}`" for column in CELLINFO_COLUMNS)
values = [tuple(row.get(column, "") for column in CELLINFO_COLUMNS) for row in rows]
total = len(values)
for start in range(0, total, 1000):
cursor.executemany(
f"INSERT INTO `{table}` ({columns}) VALUES ({placeholders})",
values[start:start + 1000],
)
self._log(f"写入中 {min(start + 1000, total)}/{total} 行...")
conn.commit()
self._log(f"已提交,成功写入 {len(rows)} 行到 `{table}`")
return len(rows)
except Exception as exc:
conn.rollback()
self._log(f"写入失败,已回滚: {exc}")
raise
finally:
conn.close()
@staticmethod
def _ensure_cellinfo_table(cursor, table: str) -> None:
cursor.execute(
f"""
CREATE TABLE IF NOT EXISTS `{table}` (
`CGI` varchar(120) DEFAULT NULL,
`eNodeBID` int DEFAULT NULL,
`CellID` int DEFAULT NULL,
`PLMN` varchar(100) DEFAULT NULL,
`基站名称` varchar(200) DEFAULT NULL,
`小区名称` varchar(200) DEFAULT NULL,
`频点` varchar(50) DEFAULT NULL,
`带宽` varchar(20) DEFAULT NULL,
`制式` varchar(50) DEFAULT NULL,
`功率` varchar(100) DEFAULT NULL,
`网络` varchar(20) DEFAULT NULL,
KEY `CGI` (`CGI`)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci
"""
)
def _list_dir(self, remote_path: str) -> list[dict[str, Any]]:
if self.config.remote_data.protocol == "ftp":
with self.downloader._ftp_client() as ftp: # noqa: SLF001
return [
{
"name": name,
"type": "directory" if entry_type == "dir" else "file",
"size": size,
"path": self._join_remote_path(remote_path, name),
}
for name, entry_type, size in self.downloader._list_ftp_entries(ftp, remote_path) # noqa: SLF001
if name not in {".", ".."}
]
ssh = self.downloader._sftp_ssh_client() # noqa: SLF001
try:
with ssh.open_sftp() as sftp:
return [
{
"name": item.filename,
"type": "directory" if stat.S_ISDIR(item.st_mode) else "file",
"size": int(getattr(item, "st_size", 0) or 0),
"path": self._join_remote_path(remote_path, item.filename),
}
for item in sftp.listdir_attr(remote_path)
if item.filename not in {".", ".."}
]
finally:
ssh.close()
def _remote_file_info(self, remote_path: str, size: int) -> RemoteFileInfo:
return self.downloader._remote_file_info(remote_path, size) # noqa: SLF001
@staticmethod
def _join_remote_path(parent: str, child: str) -> str:
parent = parent.replace("\\", "/").rstrip("/")
child = child.replace("\\", "/").strip("/")
if not parent:
return child
if parent == "/":
return f"/{child}"
return f"{parent}/{child}"
def _set_stage(self, stage: str) -> None:
if hasattr(self.logger, "set_stage"):
self.logger.set_stage(stage)
def _log(self, message: str) -> None:
if self.logger:
if hasattr(self.logger, "info"):
self.logger.info(message)
else:
self.logger(message)
def refresh_cell_data(app_config: AppConfig, work_dir: Path, logger=None) -> CellDataResult:
return CellDataProcessor(app_config, work_dir, logger).run()
def execute_celldata_script(script_path: Path, app_config: AppConfig, logger=None) -> None:
try:
from app.db_init import ensure_required_tables
ensure_required_tables(app_config, logger)
except Exception as exc: # noqa: BLE001
if logger:
logger.info(f"[前置检查] 执行异常:{exc}")
if not script_path.exists():
if logger:
logger.info("CellData 脚本文件不存在,跳过")
return
sql_text = script_path.read_text(encoding="utf-8").strip()
if not sql_text:
if logger:
logger.info("CellData 脚本为空,跳过")
return
from app.processor import DataProcessor
statements = DataProcessor.parse_sql_script(sql_text)
if not statements:
if logger:
logger.info("CellData 脚本中没有有效语句,跳过")
return
mysql = app_config.cell_data.mysql.normalized()
conn = pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
database=mysql.dbname,
charset="utf8mb4",
cursorclass=pymysql.cursors.DictCursor,
autocommit=False,
)
try:
with conn.cursor() as cursor:
for i, sql in enumerate(statements, 1):
preview = sql[:80].replace("\n", " ")
if logger:
logger.info(f"CellData SQL ({i}/{len(statements)}): {preview}...")
cursor.execute(sql)
affected = cursor.rowcount if cursor.rowcount >= 0 else 0
if affected > 0 and logger:
logger.info(f"完成,影响 {affected} 行")
conn.commit()
if logger:
logger.success(f"CellData 脚本执行完成,共 {len(statements)} 条语句")
except Exception:
conn.rollback()
raise
finally:
conn.close()
def copy_celldata_tables_to_capacity(app_config: AppConfig, logger=None) -> int:
cd_mysql = app_config.cell_data.mysql.normalized()
cap_mysql = app_config.mysql.normalized()
if (
cd_mysql.host == cap_mysql.host
and cd_mysql.port == cap_mysql.port
and cd_mysql.dbname == cap_mysql.dbname
):
if logger:
logger.info("CellData 与容量数据库相同,跳过表复制")
return 0
cd_conn = pymysql.connect(
host=cd_mysql.host, port=cd_mysql.port,
user=cd_mysql.user, password=cd_mysql.passwd,
database=cd_mysql.dbname, charset="utf8mb4",
cursorclass=pymysql.cursors.DictCursor,
)
cap_conn = pymysql.connect(
host=cap_mysql.host, port=cap_mysql.port,
user=cap_mysql.user, password=cap_mysql.passwd,
database=cap_mysql.dbname, charset="utf8mb4",
cursorclass=pymysql.cursors.DictCursor,
local_infile=True, autocommit=False,
)
try:
tables = _list_tables(cd_conn)
if not tables:
if logger:
logger.info("CellData 数据库中没有表,跳过复制")
return 0
target_charset, target_collation = _database_charset_and_collation(cap_conn)
copied = 0
for table in tables:
rows = _copy_one_table(
cd_conn,
cap_conn,
table,
target_charset,
target_collation,
logger,
)
if rows >= 0:
copied += 1
cap_conn.commit()
if logger:
logger.success(f"已将 {copied} 张表从 CellData 复制到容量数据库")
return copied
except Exception:
cap_conn.rollback()
raise
finally:
cd_conn.close()
cap_conn.close()
def _list_tables(conn) -> list[str]:
with conn.cursor() as cur:
cur.execute("SHOW TABLES")
return [list(row.values())[0] for row in cur.fetchall()]
def _database_charset_and_collation(conn) -> tuple[str, str]:
with conn.cursor() as cur:
cur.execute(
"SELECT @@character_set_database AS `charset_name`, "
"@@collation_database AS `collation_name`"
)
row = cur.fetchone()
charset = str(row["charset_name"])
collation = str(row["collation_name"])
if not re.fullmatch(r"[A-Za-z0-9_]+", charset) or not re.fullmatch(r"[A-Za-z0-9_]+", collation):
raise RuntimeError("目标数据库字符集或排序规则名称无效")
return charset, collation
def _copy_one_table(
src_conn,
dst_conn,
table: str,
target_charset: str,
target_collation: str,
logger=None,
) -> int:
with src_conn.cursor() as cur:
cur.execute(f"SHOW CREATE TABLE `{table}`")
row = cur.fetchone()
create_sql = row.get("Create Table") or list(row.values())[1]
with src_conn.cursor() as cur:
cur.execute(f"SELECT COUNT(*) AS `cnt` FROM `{table}`")
count = cur.fetchone()["cnt"]
with dst_conn.cursor() as cur:
cur.execute(f"DROP TABLE IF EXISTS `{table}`")
cur.execute(create_sql)
cur.execute(
f"ALTER TABLE `{table}` CONVERT TO CHARACTER SET {target_charset} "
f"COLLATE {target_collation}"
)
if count == 0:
if logger:
logger.info(f"复制表 {table}: 0 行(空表)")
return 0
with src_conn.cursor() as src_cur:
src_cur.execute(f"SELECT * FROM `{table}`")
columns = [desc[0] for desc in src_cur.description]
col_list = ", ".join(f"`{c}`" for c in columns)
placeholders = ", ".join(["%s"] * len(columns))
insert_sql = f"INSERT INTO `{table}` ({col_list}) VALUES ({placeholders})"
batch: list[tuple] = []
inserted = 0
with dst_conn.cursor() as dst_cur:
for row in src_cur:
batch.append(tuple(row.values()))
if len(batch) >= 5000:
dst_cur.executemany(insert_sql, batch)
inserted += len(batch)
batch.clear()
if batch:
dst_cur.executemany(insert_sql, batch)
inserted += len(batch)
if logger:
logger.info(f"复制表 {table}: {inserted} 行")
return inserted
+412
View File
@@ -0,0 +1,412 @@
"""纯数据处理流水线(容器版,无数据库代码)。
输入:包含已下载周 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.directories = _build_directory_mappings(config.get("directories") or [])
self.field_map, self.type_map = _build_global_map(config.get("extract_fields", []))
self.table_maps = _build_table_maps(config.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, directories in data_dirs.items():
csv_files: list[Path] = []
for directory in directories:
csv_files.extend(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, list[Path]]:
data_dirs: dict[str, list[Path]] = {}
for item in self.directories:
directory = self.work_dir / item["path"]
if directory.exists() and directory.is_dir():
_add_data_dir(data_dirs, item["table"], directory)
return data_dirs
def _maps_for_table(self, table: str):
if table in self.table_maps:
return self.table_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_directory_mappings(items: list[dict]) -> list[dict[str, str]]:
mappings: list[dict[str, str]] = []
seen: set[tuple[str, str]] = set()
for item in items:
if not isinstance(item, dict):
continue
path = str(item.get("path", "")).replace("\\", "/").strip().strip("/")
table = str(item.get("table", "")).strip()
if not path or not table:
continue
key = (path, table)
if key in seen:
continue
seen.add(key)
mappings.append({"path": path, "table": table})
return mappings
def _add_data_dir(data_dirs: dict[str, list[Path]], table: str, directory: Path) -> None:
existing = data_dirs.setdefault(table, [])
resolved = directory.resolve()
if all(path.resolve() != resolved for path in existing):
existing.append(directory)
def _build_table_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
+190
View File
@@ -0,0 +1,190 @@
"""本地授权期限校验。"""
import base64
import hashlib
import hmac
import json
import re
from dataclasses import dataclass
from datetime import date, datetime, timedelta
from pathlib import Path
from typing import Any
from app.config import BASE_DIR
DEFAULT_EXPIRES_ON = date(2026, 12, 30)
EXTEND_DAYS = 30
LICENSE_FILE = BASE_DIR / "license.dat"
_SECRET = b"CapacityReport local license v1"
_ZIP_DATE_RE = re.compile(r"(?<!\d)(20\d{10}(?:\d{2})?)(?!\d)")
class LicenseError(Exception):
"""授权校验错误。"""
code = "LICENSE_ERROR"
def to_detail(self) -> dict[str, Any]:
return {"code": self.code, "message": str(self)}
class LicenseExpiredError(LicenseError):
"""数据日期超过授权到期日期。"""
code = "LICENSE_EXPIRED"
def __init__(self, expires_on: date, current_date: date):
self.expires_on = expires_on
self.current_date = current_date
super().__init__(
f"授权已过期:数据日期 {current_date.isoformat()} 已超过到期日期 {expires_on.isoformat()}"
)
def to_detail(self) -> dict[str, Any]:
return {
"code": self.code,
"message": str(self),
"expires_on": self.expires_on.isoformat(),
"current_date": self.current_date.isoformat(),
"key_label": format_key_label(self.expires_on),
}
class InvalidActivationCodeError(LicenseError):
"""激活码错误。"""
code = "LICENSE_INVALID"
def __init__(self, expires_on: date):
self.expires_on = expires_on
super().__init__("激活码无效,请按当前 key 重新计算后输入")
def to_detail(self) -> dict[str, Any]:
return {
"code": self.code,
"message": str(self),
"expires_on": self.expires_on.isoformat(),
"key_label": format_key_label(self.expires_on),
}
@dataclass(frozen=True)
class LicenseInfo:
expires_on: date
current_date: date | None = None
zip_count: int = 0
@property
def key_label(self) -> str:
return format_key_label(self.expires_on)
def to_dict(self) -> dict[str, Any]:
return {
"expires_on": self.expires_on.isoformat(),
"key_label": self.key_label,
"current_date": self.current_date.isoformat() if self.current_date else None,
"zip_count": self.zip_count,
}
def get_license_info() -> LicenseInfo:
return LicenseInfo(expires_on=read_expires_on())
def activate(code: str) -> LicenseInfo:
expires_on = read_expires_on()
expected = activation_hash(expires_on)
normalized_code = (code or "").strip().lower()
if not hmac.compare_digest(normalized_code, expected):
raise InvalidActivationCodeError(expires_on)
new_expires_on = expires_on + timedelta(days=EXTEND_DAYS)
write_expires_on(new_expires_on)
return LicenseInfo(expires_on=new_expires_on)
def check_processing_allowed(work_dir: Path) -> LicenseInfo:
expires_on = read_expires_on()
zip_count, current_date = extract_max_zip_date(work_dir)
info = LicenseInfo(expires_on=expires_on, current_date=current_date, zip_count=zip_count)
if current_date and current_date > expires_on:
raise LicenseExpiredError(expires_on, current_date)
return info
def extract_max_zip_date(work_dir: Path) -> tuple[int, date | None]:
max_date: date | None = None
zip_count = 0
for zip_file in work_dir.rglob("*.zip"):
zip_count += 1
for raw_value in _ZIP_DATE_RE.findall(zip_file.name):
parsed_date = _parse_zip_timestamp(raw_value)
if parsed_date and (max_date is None or parsed_date > max_date):
max_date = parsed_date
return zip_count, max_date
def activation_hash(expires_on: date) -> str:
return hashlib.sha256(format_key_label(expires_on).encode("utf-8")).hexdigest()
def format_key_label(value: date) -> str:
return value.strftime("%Y/%m/%d")
def read_expires_on() -> date:
if not LICENSE_FILE.exists():
write_expires_on(DEFAULT_EXPIRES_ON)
return DEFAULT_EXPIRES_ON
try:
encrypted = base64.urlsafe_b64decode(LICENSE_FILE.read_text(encoding="utf-8").encode("ascii"))
raw = _xor_bytes(encrypted)
data = json.loads(raw.decode("utf-8"))
payload = data["payload"]
signature = data["signature"]
payload_raw = _dump_json(payload)
expected_signature = hmac.new(_SECRET, payload_raw, hashlib.sha256).hexdigest()
if not hmac.compare_digest(signature, expected_signature):
raise ValueError("signature mismatch")
return date.fromisoformat(str(payload["expires_on"]))
except Exception:
write_expires_on(DEFAULT_EXPIRES_ON)
return DEFAULT_EXPIRES_ON
def write_expires_on(expires_on: date) -> None:
payload = {"expires_on": expires_on.isoformat()}
payload_raw = _dump_json(payload)
data = {
"payload": payload,
"signature": hmac.new(_SECRET, payload_raw, hashlib.sha256).hexdigest(),
}
encrypted = _xor_bytes(_dump_json(data))
LICENSE_FILE.write_text(base64.urlsafe_b64encode(encrypted).decode("ascii"), encoding="utf-8")
def _parse_zip_timestamp(value: str) -> date | None:
fmt = "%Y%m%d%H%M%S" if len(value) == 14 else "%Y%m%d%H%M"
try:
return datetime.strptime(value, fmt).date()
except ValueError:
return None
def _dump_json(data: dict[str, Any]) -> bytes:
return json.dumps(data, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode("utf-8")
def _xor_bytes(data: bytes) -> bytes:
output = bytearray()
counter = 0
while len(output) < len(data):
block = hashlib.sha256(_SECRET + counter.to_bytes(4, "big")).digest()
output.extend(block)
counter += 1
return bytes(value ^ key for value, key in zip(data, output))
+103
View File
@@ -0,0 +1,103 @@
"""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
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()
mappings = app_config.data_mappings.normalized()
return {
"recent_days": metrix.recent_days,
"sheet_filter": list(app_config.sheet_filter),
"directories": list(mappings.directories),
"extract_fields": app_config.extract_fields,
"table_field_mappings": mappings.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)}
+312
View File
@@ -0,0 +1,312 @@
"""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, content: str = "",
database: str = "", single_session: bool = False, run_timeout: int = 7200) -> dict:
# 始终按传入的 SQL 文本执行(来自本地 ReportScript.sql),不走 Metrix 库内脚本(script_id)那条路。
body: dict = {"content": content, "database": database, "stop_on_error": True, "single_session": single_session}
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,
)
+488
View File
@@ -0,0 +1,488 @@
from __future__ import annotations
import stat
from collections import defaultdict
from dataclasses import dataclass, field
from datetime import date
from ftplib import FTP
from pathlib import Path
from typing import Callable, Iterable
from app.config import RemoteDataConfig
from app.utils.file_dates import parse_file_date_range, select_recent_items_by_directory
LogFn = Callable[[str], None]
@dataclass
class RemoteDownloadResult:
file_count: int = 0
total_bytes: int = 0
remote_files: list[str] = field(default_factory=list)
@dataclass(frozen=True)
class RemoteFileInfo:
path: str
relative_path: str
parent: str
name: str
size: int = 0
class RemoteDownloadError(RuntimeError):
pass
class RemoteDataDownloader:
def __init__(self, config: RemoteDataConfig, logger: LogFn | None = None):
self.config = config.normalized()
self.logger = logger
def test_connection(self) -> None:
self._validate_config()
if self.config.protocol == "ftp":
with self._ftp_client() as ftp:
ftp.cwd(self.config.remote_dir)
return
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
sftp.stat(self.config.remote_dir)
finally:
ssh.close()
def download_to(self, destination: Path, target_dates: Iterable[date] | None = None) -> RemoteDownloadResult:
self._validate_config()
destination.mkdir(parents=True, exist_ok=True)
zip_files = self.list_remote_zip_files()
date_filter = set(target_dates or [])
if date_filter:
selected_files = self._select_files_by_dates(zip_files, date_filter)
return self._download_selected_files(destination, selected_files)
if zip_files:
selected_files, 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.max_date and summary.start_date and summary.skipped_count:
self._log(
f"远程目录 {summary.directory or '.'}: 仅下载 "
f"{summary.start_date.isoformat()} 至 {summary.max_date.isoformat()} "
f"的 {summary.selected_count}/{summary.total_count} 个 ZIP 文件,"
f"跳过 {summary.skipped_count} 个旧文件"
)
return self._download_selected_files(destination, selected_files)
if self.config.protocol == "ftp":
return self._download_ftp(destination)
return self._download_sftp(destination)
def _select_files_by_dates(self, zip_files: list[RemoteFileInfo], target_dates: set[date]) -> list[RemoteFileInfo]:
selected_files: list[RemoteFileInfo] = []
grouped: dict[str, list[RemoteFileInfo]] = defaultdict(list)
for remote_file in zip_files:
grouped[remote_file.parent].append(remote_file)
for parent, files in sorted(grouped.items(), key=lambda item: item[0]):
selected = [
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_files.extend(selected)
skipped_count = len(files) - len(selected)
if skipped_count:
first_day = min(target_dates).isoformat()
last_day = max(target_dates).isoformat()
self._log(
f"远程目录 {parent or '.'}: 仅下载调度目标日期 "
f"{first_day} 至 {last_day} 的 {len(selected)}/{len(files)} 个 ZIP 文件,"
f"跳过 {skipped_count} 个非目标日期文件"
)
return selected_files
def list_remote_zip_files(self, directory: str | None = None) -> list[RemoteFileInfo]:
self._validate_config()
remote_dir = (
self._join_remote_path(self.config.remote_dir, directory.strip("/"))
if directory
else self.config.remote_dir
)
if self.config.protocol == "ftp":
return self._list_ftp_zip_files(remote_dir)
return self._list_sftp_zip_files(remote_dir)
def delete_source_files(self, remote_files: Iterable[str] | None = None) -> int:
self._validate_config()
source_files = list(remote_files or [])
if source_files:
if self.config.protocol == "ftp":
return self._delete_ftp_file_paths(source_files)
return self._delete_sftp_file_paths(source_files)
if self.config.protocol == "ftp":
return self._delete_ftp_source_files()
return self._delete_sftp_source_files()
def _validate_config(self) -> None:
if self.config.protocol not in {"ftp", "sftp"}:
raise RemoteDownloadError("远程协议只支持 FTP 或 SFTP")
if not self.config.host:
raise RemoteDownloadError("请填写远程服务器地址")
if not self.config.user:
raise RemoteDownloadError("请填写远程服务器用户名")
if not self.config.remote_dir:
raise RemoteDownloadError("请填写远程数据目录")
def _log(self, message: str) -> None:
if self.logger:
self.logger(message)
def _ftp_client(self) -> FTP:
ftp = FTP()
ftp.connect(self.config.host, self.config.port, timeout=self.config.timeout)
ftp.login(self.config.user, self.config.passwd)
ftp.set_pasv(self.config.passive)
return ftp
def _download_ftp(self, destination: Path) -> RemoteDownloadResult:
result = RemoteDownloadResult()
with self._ftp_client() as ftp:
self._log(f"已连接 FTP: {self.config.host}:{self.config.port}")
self._download_ftp_dir(ftp, self.config.remote_dir, destination, result)
return result
def _download_ftp_dir(
self,
ftp: FTP,
remote_dir: str,
local_dir: Path,
result: RemoteDownloadResult,
) -> None:
local_dir.mkdir(parents=True, exist_ok=True)
entries = self._list_ftp_entries(ftp, remote_dir)
for name, entry_type, size in entries:
if name in {".", ".."}:
continue
remote_path = self._join_remote_path(remote_dir, name)
local_path = local_dir / name
if entry_type == "dir":
self._download_ftp_dir(ftp, remote_path, local_path, result)
continue
if entry_type == "unknown" and self._ftp_is_dir(ftp, remote_path):
self._download_ftp_dir(ftp, remote_path, local_path, result)
continue
self._download_ftp_file(ftp, remote_path, local_path, result, size)
def _list_ftp_entries(self, ftp: FTP, remote_dir: str) -> list[tuple[str, str, int]]:
try:
return [
(
name,
facts.get("type", "unknown"),
int(facts.get("size", "0") or 0),
)
for name, facts in ftp.mlsd(remote_dir)
]
except Exception:
names = ftp.nlst(remote_dir)
entries = []
for name in names:
clean_name = Path(name.replace("\\", "/")).name
if clean_name:
entries.append((clean_name, "unknown", 0))
return entries
def _ftp_is_dir(self, ftp: FTP, remote_path: str) -> bool:
current = ftp.pwd()
try:
ftp.cwd(remote_path)
return True
except Exception:
return False
finally:
try:
ftp.cwd(current)
except Exception:
pass
def _download_ftp_file(
self,
ftp: FTP,
remote_path: str,
local_path: Path,
result: RemoteDownloadResult,
expected_size: int = 0,
) -> None:
local_path.parent.mkdir(parents=True, exist_ok=True)
self._log(f"下载: {remote_path}")
with local_path.open("wb") as file:
ftp.retrbinary(f"RETR {remote_path}", file.write)
result.file_count += 1
result.total_bytes += expected_size or local_path.stat().st_size
result.remote_files.append(remote_path)
def _download_selected_files(self, destination: Path, remote_files: list[RemoteFileInfo]) -> RemoteDownloadResult:
result = RemoteDownloadResult()
if not remote_files:
return result
if self.config.protocol == "ftp":
with self._ftp_client() as ftp:
self._log(f"已连接 FTP: {self.config.host}:{self.config.port}")
for remote_file in remote_files:
self._download_ftp_file(
ftp,
remote_file.path,
destination / remote_file.relative_path,
result,
remote_file.size,
)
return result
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
self._log(f"已连接 SFTP: {self.config.host}:{self.config.port}")
for remote_file in remote_files:
self._download_sftp_file(
sftp,
remote_file.path,
destination / remote_file.relative_path,
result,
remote_file.size,
)
finally:
ssh.close()
return result
def _list_ftp_zip_files(self, remote_dir: str) -> list[RemoteFileInfo]:
with self._ftp_client() as ftp:
return self._collect_ftp_zip_files(ftp, remote_dir)
def _collect_ftp_zip_files(self, ftp: FTP, remote_dir: str) -> list[RemoteFileInfo]:
files: list[RemoteFileInfo] = []
for name, entry_type, size in self._list_ftp_entries(ftp, remote_dir):
if name in {".", ".."}:
continue
remote_path = self._join_remote_path(remote_dir, name)
if entry_type == "dir" or (entry_type == "unknown" and self._ftp_is_dir(ftp, remote_path)):
files.extend(self._collect_ftp_zip_files(ftp, remote_path))
continue
if name.lower().endswith(".zip"):
files.append(self._remote_file_info(remote_path, size))
return files
def _delete_ftp_source_files(self) -> int:
with self._ftp_client() as ftp:
self._log(f"开始清理 FTP 源文件: {self.config.remote_dir}")
return self._delete_ftp_files(ftp, self.config.remote_dir)
def _delete_ftp_files(self, ftp: FTP, remote_dir: str) -> int:
deleted_count = 0
for name, entry_type, _ in self._list_ftp_entries(ftp, remote_dir):
if name in {".", ".."}:
continue
remote_path = self._join_remote_path(remote_dir, name)
if entry_type == "dir" or (entry_type == "unknown" and self._ftp_is_dir(ftp, remote_path)):
deleted_count += self._delete_ftp_files(ftp, remote_path)
continue
self._log(f"删除远程文件: {remote_path}")
ftp.delete(remote_path)
deleted_count += 1
return deleted_count
def _delete_ftp_file_paths(self, remote_files: list[str]) -> int:
deleted_count = 0
with self._ftp_client() as ftp:
self._log(f"开始清理 FTP 源文件,共 {len(remote_files)} 个")
for remote_path in remote_files:
self._log(f"删除远程文件: {remote_path}")
ftp.delete(remote_path)
deleted_count += 1
return deleted_count
def _sftp_ssh_client(self):
try:
import paramiko
except ImportError as exc:
raise RemoteDownloadError("SFTP 功能需要安装 paramiko,请执行 uv pip install -r requirements.txt") from exc
ssh = paramiko.SSHClient()
ssh.set_missing_host_key_policy(paramiko.AutoAddPolicy())
ssh.connect(
hostname=self.config.host,
port=self.config.port,
username=self.config.user,
password=self.config.passwd,
timeout=self.config.timeout,
banner_timeout=self.config.timeout,
auth_timeout=self.config.timeout,
)
return ssh
def _download_sftp(self, destination: Path) -> RemoteDownloadResult:
result = RemoteDownloadResult()
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
self._log(f"已连接 SFTP: {self.config.host}:{self.config.port}")
self._download_sftp_path(sftp, self.config.remote_dir, destination, result)
finally:
ssh.close()
return result
def _download_sftp_path(
self,
sftp,
remote_path: str,
local_path: Path,
result: RemoteDownloadResult,
) -> None:
attrs = sftp.stat(remote_path)
if stat.S_ISDIR(attrs.st_mode):
local_path.mkdir(parents=True, exist_ok=True)
for item in sftp.listdir_attr(remote_path):
if item.filename in {".", ".."}:
continue
self._download_sftp_path(
sftp,
self._join_remote_path(remote_path, item.filename),
local_path / item.filename,
result,
)
return
self._download_sftp_file(sftp, remote_path, local_path, result, int(getattr(attrs, "st_size", 0) or 0))
def _download_sftp_file(
self,
sftp,
remote_path: str,
local_path: Path,
result: RemoteDownloadResult,
expected_size: int = 0,
) -> None:
local_path.parent.mkdir(parents=True, exist_ok=True)
self._log(f"下载: {remote_path}")
sftp.get(remote_path, str(local_path))
result.file_count += 1
result.total_bytes += expected_size or local_path.stat().st_size
result.remote_files.append(remote_path)
def _list_sftp_zip_files(self, remote_dir: str) -> list[RemoteFileInfo]:
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
return self._collect_sftp_zip_files(sftp, remote_dir)
finally:
ssh.close()
def _collect_sftp_zip_files(self, sftp, remote_path: str) -> list[RemoteFileInfo]:
attrs = sftp.stat(remote_path)
if not stat.S_ISDIR(attrs.st_mode):
name = Path(remote_path.replace("\\", "/")).name
if name.lower().endswith(".zip"):
return [self._remote_file_info(remote_path, int(getattr(attrs, "st_size", 0) or 0))]
return []
files: list[RemoteFileInfo] = []
for item in sftp.listdir_attr(remote_path):
if item.filename in {".", ".."}:
continue
child_path = self._join_remote_path(remote_path, item.filename)
if stat.S_ISDIR(item.st_mode):
files.extend(self._collect_sftp_zip_files(sftp, child_path))
elif item.filename.lower().endswith(".zip"):
files.append(self._remote_file_info(child_path, int(getattr(item, "st_size", 0) or 0)))
return files
def _delete_sftp_source_files(self) -> int:
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
self._log(f"开始清理 SFTP 源文件: {self.config.remote_dir}")
return self._delete_sftp_files(sftp, self.config.remote_dir)
finally:
ssh.close()
def _delete_sftp_files(self, sftp, remote_path: str) -> int:
attrs = sftp.stat(remote_path)
if stat.S_ISDIR(attrs.st_mode):
deleted_count = 0
for item in sftp.listdir_attr(remote_path):
if item.filename in {".", ".."}:
continue
deleted_count += self._delete_sftp_files(
sftp,
self._join_remote_path(remote_path, item.filename),
)
return deleted_count
self._log(f"删除远程文件: {remote_path}")
sftp.remove(remote_path)
return 1
def _delete_sftp_file_paths(self, remote_files: list[str]) -> int:
deleted_count = 0
ssh = self._sftp_ssh_client()
try:
with ssh.open_sftp() as sftp:
self._log(f"开始清理 SFTP 源文件,共 {len(remote_files)} 个")
for remote_path in remote_files:
self._log(f"删除远程文件: {remote_path}")
sftp.remove(remote_path)
deleted_count += 1
finally:
ssh.close()
return deleted_count
@staticmethod
def _join_remote_path(parent: str, child: str) -> str:
parent = parent.replace("\\", "/").rstrip("/")
if not parent:
return child
if parent == "/":
return f"/{child}"
return f"{parent}/{child}"
def _remote_file_info(self, remote_path: str, size: int = 0) -> RemoteFileInfo:
normalized_root = self.config.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,
)