feat: 新增远程数据自动化处理

This commit is contained in:
2026-05-18 18:21:20 +08:00
parent 3750160891
commit f2bef1f2f7
12 changed files with 698 additions and 5 deletions
+11 -1
View File
@@ -6,7 +6,7 @@ from fastapi import APIRouter, Body, HTTPException, UploadFile, File
from fastapi.responses import FileResponse
from app import state
from app.config import CONFIG_FILE
from app.config import CONFIG_FILE, RemoteDataConfig
router = APIRouter(tags=["config"])
@@ -39,6 +39,13 @@ async def update_mysql_config(
return {"success": True, "message": "数据库配置已更新", "update": state.config.update}
@router.post("/api/config/remote")
async def update_remote_config(config: dict[str, Any] = Body(...)):
state.config.remote_data = RemoteDataConfig.from_dict(config)
state.config.save()
return {"success": True, "message": "远程数据配置已更新", "update": state.config.update}
@router.post("/api/config/sheet-filter")
async def update_sheet_filter(filters: list[str] = Body(...)):
state.config.sheet_filter = filters
@@ -97,3 +104,6 @@ def _apply_config_data(data: dict[str, Any]) -> None:
if "ExtractField" in data:
state.config.extract_fields = data["ExtractField"] if isinstance(data["ExtractField"], list) else []
remote_data = data.get("RemoteData")
if isinstance(remote_data, dict):
state.config.remote_data = RemoteDataConfig.from_dict(remote_data)
+129
View File
@@ -0,0 +1,129 @@
from datetime import datetime
from pathlib import Path
from threading import Thread
from typing import Any
from fastapi import APIRouter, Body, HTTPException
from app import state
from app.config import CACHE_DIR, RemoteDataConfig
from app.processor import DataProcessor, ProcessLogger
from app.services.remote_download import RemoteDataDownloader
router = APIRouter(tags=["remote"])
@router.post("/api/remote/test")
async def test_remote_connection(config: dict[str, Any] | None = Body(None)):
remote_config = RemoteDataConfig.from_dict(config) if config else state.config.remote_data
try:
RemoteDataDownloader(remote_config).test_connection()
return {"success": True, "message": "远程服务器连接成功"}
except Exception as exc:
return {"success": False, "message": f"远程服务器连接失败: {exc}"}
@router.post("/api/remote/start")
async def start_remote_processing():
if state.global_task_lock["locked"]:
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
remote_config = state.config.remote_data.normalized()
if not remote_config.enabled:
raise HTTPException(status_code=400, detail="请先启用远程数据配置")
task_id = datetime.now().strftime("%Y%m%d_%H%M%S")
work_dir = CACHE_DIR / task_id
work_dir.mkdir(parents=True, exist_ok=True)
logs: list[str] = []
def log_callback(message: str) -> None:
logs.append(message)
state.processing_tasks[task_id] = {"logs": logs.copy(), "status": "processing"}
logger = ProcessLogger(log_file=work_dir / "log.txt", callback=log_callback)
state.history_manager.create(work_dir, 0, record_id=task_id)
state.history_manager.update(task_id, status="processing")
state.processing_tasks[task_id] = {"logs": [], "status": "processing"}
state.global_task_lock.update(
{
"locked": True,
"task_id": task_id,
"stage": "downloading",
"started_at": datetime.now().isoformat(),
}
)
thread = Thread(
target=_run_remote_processing,
args=(task_id, work_dir, remote_config, logger),
daemon=True,
)
thread.start()
return {
"success": True,
"message": "远程下载处理任务已启动",
"task_id": task_id,
"stage": "downloading",
}
def _run_remote_processing(
task_id: str,
work_dir: Path,
remote_config: RemoteDataConfig,
logger: ProcessLogger,
) -> None:
try:
logger.info(
f"开始远程下载,协议: {remote_config.protocol.upper()},"
f"服务器: {remote_config.host}:{remote_config.port},目录: {remote_config.remote_dir}"
)
downloader = RemoteDataDownloader(remote_config, logger.info)
download_result = downloader.download_to(work_dir)
logger.success(
f"远程下载完成,共 {download_result.file_count} 个文件,"
f"{_format_bytes(download_result.total_bytes)}"
)
if download_result.file_count == 0:
raise RuntimeError("远程目录中未下载到任何文件")
state.history_manager.update(task_id, file_count=download_result.file_count)
state.global_task_lock["stage"] = "processing"
processor = DataProcessor(state.config, work_dir, logger)
result = processor.process()
status = "completed" if result.get("success") else "failed"
state.history_manager.update(
task_id,
status=status,
elapsed_time=result.get("elapsed_time", 0),
error=result.get("error"),
result_tables=["4G_结果表", "5G_结果表"],
)
state.processing_tasks[task_id] = {
"logs": state.history_manager.get_logs(task_id),
"status": status,
}
except Exception as exc:
logger.error(f"远程自动化任务失败: {exc}")
state.history_manager.update(task_id, status="failed", error=str(exc))
state.processing_tasks[task_id] = {
"logs": state.history_manager.get_logs(task_id),
"status": "failed",
}
finally:
state.reset_task_lock()
def _format_bytes(size: int) -> str:
value = float(size)
for unit in ("B", "KB", "MB", "GB"):
if value < 1024:
return f"{value:.1f} {unit}"
value /= 1024
return f"{value:.1f} TB"
+77
View File
@@ -23,10 +23,82 @@ class MySQLConfig:
dbname: str = "CapacityReport"
@dataclass
class RemoteDataConfig:
enabled: bool = False
protocol: str = "sftp"
host: str = ""
port: int = 22
user: str = ""
passwd: str = ""
remote_dir: str = "/"
passive: bool = True
timeout: int = 30
def normalized(self) -> "RemoteDataConfig":
protocol = self.protocol.lower().strip()
if protocol not in {"ftp", "sftp"}:
protocol = "sftp"
port = self.port or (22 if protocol == "sftp" else 21)
return RemoteDataConfig(
enabled=bool(self.enabled),
protocol=protocol,
host=self.host.strip(),
port=port,
user=self.user.strip(),
passwd=self.passwd,
remote_dir=(self.remote_dir or "/").strip() or "/",
passive=bool(self.passive),
timeout=max(int(self.timeout or 30), 1),
)
def to_dict(self, include_password: bool = False) -> Dict[str, Any]:
data = {
"enabled": self.enabled,
"protocol": self.protocol,
"host": self.host,
"port": self.port,
"user": self.user,
"remote_dir": self.remote_dir,
"passive": self.passive,
"timeout": self.timeout,
}
if include_password:
data["passwd"] = self.passwd
return data
@classmethod
def from_dict(cls, data: Dict[str, Any] | None) -> "RemoteDataConfig":
data = data or {}
protocol = str(data.get("protocol", "sftp")).lower()
default_port = 22 if protocol == "sftp" else 21
try:
port = int(data.get("port") or default_port)
except (TypeError, ValueError):
port = default_port
try:
timeout = int(data.get("timeout") or 30)
except (TypeError, ValueError):
timeout = 30
return cls(
enabled=bool(data.get("enabled", False)),
protocol=protocol,
host=str(data.get("host", "")),
port=port,
user=str(data.get("user", "")),
passwd=str(data.get("passwd", "")),
remote_dir=str(data.get("remote_dir", "/")),
passive=bool(data.get("passive", True)),
timeout=timeout,
).normalized()
@dataclass
class AppConfig:
update: str = ""
mysql: MySQLConfig = field(default_factory=MySQLConfig)
remote_data: RemoteDataConfig = field(default_factory=RemoteDataConfig)
sheet_filter: List[str] = field(default_factory=list)
extract_fields: List[Dict[str, Any]] = field(default_factory=list)
@@ -47,10 +119,12 @@ class AppConfig:
passwd=mysql_data.get("passwd", ""),
dbname=mysql_data.get("dbname", "CapacityReport")
)
remote_config = RemoteDataConfig.from_dict(data.get("RemoteData"))
return cls(
update=data.get("Update", ""),
mysql=mysql_config,
remote_data=remote_config,
sheet_filter=data.get("SheetFilter", []),
extract_fields=data.get("ExtractField", [])
)
@@ -68,6 +142,7 @@ class AppConfig:
"passwd": self.mysql.passwd,
"dbname": self.mysql.dbname
},
"RemoteData": self.remote_data.normalized().to_dict(include_password=True),
"SheetFilter": self.sheet_filter,
"ExtractField": self.extract_fields
}
@@ -85,6 +160,7 @@ class AppConfig:
"user": self.mysql.user,
"dbname": self.mysql.dbname
},
"remote_data": self.remote_data.normalized().to_dict(),
"sheet_filter": self.sheet_filter,
"extract_fields": self.extract_fields
}
@@ -100,6 +176,7 @@ class AppConfig:
"passwd": self.mysql.passwd,
"dbname": self.mysql.dbname
},
"remote_data": self.remote_data.normalized().to_dict(include_password=True),
"sheet_filter": self.sheet_filter,
"extract_fields": self.extract_fields
}
+2 -1
View File
@@ -8,7 +8,7 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, JSONResponse
from app import state
from app.api.routers import auth, cache, config, database, health, history, script, service, tasks, upload
from app.api.routers import auth, cache, config, database, health, history, remote, script, service, tasks, upload
from app.auth import verify_jwt_token
from app.config import BASE_DIR
@@ -65,6 +65,7 @@ def register_routes(app: FastAPI) -> None:
auth.router,
health.router,
upload.router,
remote.router,
tasks.router,
history.router,
database.router,
+217
View File
@@ -0,0 +1,217 @@
from __future__ import annotations
import stat
from dataclasses import dataclass
from ftplib import FTP
from pathlib import Path
from typing import Callable
from app.config import RemoteDataConfig
LogFn = Callable[[str], None]
@dataclass
class RemoteDownloadResult:
file_count: int = 0
total_bytes: 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) -> RemoteDownloadResult:
self._validate_config()
destination.mkdir(parents=True, exist_ok=True)
if self.config.protocol == "ftp":
return self._download_ftp(destination)
return self._download_sftp(destination)
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
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
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 += int(getattr(attrs, "st_size", 0) or local_path.stat().st_size)
@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}"