feat: 导入 CapacityReport 初始源码
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
"""Service helpers."""
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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)}
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user