feat: 新增远程数据自动化处理
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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}"
|
||||
Reference in New Issue
Block a user