- 修复 remove_sync_source 引入的 stop_sync KeyError 回归(容错 get) - stop_sync 不再删除存活线程条目(防双 worker);FTP 子目录递归传真名(不再 KeyError) - 远程文件名统一 _safe_remote_name 校验(FTP/SFTP/HTTP 防路径穿越) - temp 任务名加随机后缀防碰撞;temp 状态/线程完成后清理(防 sync_state.json 膨胀) - save_sync_state 快照+原子写(临时文件+os.replace),持锁调用不死锁 - URL 打印脱敏(user:pass@ -> ***@) - _need_sync_http size=0 不再全量重下(仅按存在性) - cron 同一分钟去重;PyPI 进度只更新当前源 - v2 start_sync 不再把 bool 当 task_id - UserRecord.to_dict 脱敏(不返回 password_hash/token),新增 to_dict_private/get_user_with_password - config_hotreload 单配置源(set+persist 与热重载一致);server.py 用 get_all() - 会话创建时顺带清理过期项;JSON stats/历史读改写加锁 - 解压目标目录先校验;FTP RETR 命令注入防护;rsync --delete 目标保护 - git/urlopen 补超时;cleanup_completed_tasks 删旧留新
1375 lines
53 KiB
Python
1375 lines
53 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
|
||
"""镜像同步管理模块"""
|
||
|
||
import os
|
||
import json
|
||
import time
|
||
import ftplib
|
||
import shutil
|
||
# import importlib # 未使用
|
||
from datetime import datetime
|
||
from threading import Thread, Lock
|
||
from concurrent.futures import ThreadPoolExecutor
|
||
|
||
try:
|
||
import paramiko
|
||
except ImportError:
|
||
paramiko = None
|
||
|
||
|
||
class MirrorSyncManager:
|
||
"""镜像同步管理器"""
|
||
|
||
def __init__(self, config):
|
||
self.config = config
|
||
# 从settings.json加载同步源配置
|
||
self.sync_sources = self._load_sync_sources()
|
||
self.sync_threads = {}
|
||
self.sync_status = {}
|
||
self.sync_lock = Lock()
|
||
self.running = False
|
||
self.executor = ThreadPoolExecutor(max_workers=3)
|
||
|
||
# 定时同步调度器
|
||
self.scheduler_thread = None
|
||
self.scheduler_running = False
|
||
|
||
# 加载同步状态
|
||
self.sync_state_file = config.get('sync_state_file', 'sync_state.json')
|
||
self.load_sync_state()
|
||
|
||
# 合并状态和源
|
||
for name in list(self.sync_sources.keys()):
|
||
if name not in self.sync_status:
|
||
self.sync_status[name] = {
|
||
'last_sync': None,
|
||
'status': 'stopped',
|
||
'files_synced': 0,
|
||
'total_files': 0,
|
||
'error': None,
|
||
'schedule': self.sync_sources[name].get('schedule', {}),
|
||
'next_sync': None
|
||
}
|
||
|
||
# 启动定时同步调度器
|
||
self._start_scheduler()
|
||
|
||
def _load_sync_sources(self):
|
||
"""加载同步源配置 - 从settings.json读取"""
|
||
sources = self.config.get('sync_sources', {})
|
||
if sources:
|
||
print(f"[Sync] 从settings.json加载了 {len(sources)} 个同步源")
|
||
return sources
|
||
|
||
def _save_sync_sources(self):
|
||
"""保存同步源配置到settings.json"""
|
||
try:
|
||
# 读取现有settings.json
|
||
settings_file = 'settings.json'
|
||
if os.path.exists(settings_file):
|
||
with open(settings_file, 'r', encoding='utf-8') as f:
|
||
settings = json.load(f)
|
||
|
||
# 更新sync_sources
|
||
settings['sync_sources'] = self.sync_sources
|
||
|
||
# 写回settings.json
|
||
with open(settings_file, 'w', encoding='utf-8') as f:
|
||
json.dump(settings, f, ensure_ascii=False, indent=2)
|
||
|
||
print(f"[Sync] 保存了 {len(self.sync_sources)} 个同步源到settings.json")
|
||
|
||
# 更新内存中的配置
|
||
self.config['sync_sources'] = self.sync_sources
|
||
except Exception as e:
|
||
print(f"保存同步源到settings.json失败: {e}")
|
||
|
||
def load_sync_state(self):
|
||
"""加载同步状态"""
|
||
try:
|
||
if os.path.exists(self.sync_state_file):
|
||
with open(self.sync_state_file, 'r', encoding='utf-8') as f:
|
||
self.sync_status = json.load(f)
|
||
except Exception as e:
|
||
print(f"加载同步状态失败: {e}")
|
||
self.sync_status = {}
|
||
|
||
def save_sync_state(self):
|
||
"""保存同步状态(快照 + 原子写,可在持锁状态下调用,不会死锁)"""
|
||
try:
|
||
# dict() 浅拷贝是单条 C 操作,并发修改安全
|
||
snapshot = dict(self.sync_status)
|
||
tmp_file = self.sync_state_file + '.tmp'
|
||
with open(tmp_file, 'w', encoding='utf-8') as f:
|
||
json.dump(snapshot, f, ensure_ascii=False, indent=2)
|
||
os.replace(tmp_file, self.sync_state_file)
|
||
except Exception as e:
|
||
print(f"保存同步状态失败: {e}")
|
||
|
||
def get_sync_history(self, limit=100):
|
||
"""获取同步历史记录
|
||
|
||
Args:
|
||
limit: 返回记录数量限制
|
||
|
||
Returns:
|
||
list: 同步历史记录列表
|
||
"""
|
||
history = []
|
||
|
||
# 从 sync_status 中提取历史记录
|
||
with self.sync_lock:
|
||
for name, status in self.sync_status.items():
|
||
# 跳过临时同步任务
|
||
if status.get('is_temp_sync'):
|
||
continue
|
||
|
||
# 只记录已完成或错误的同步
|
||
if status.get('status') in ('completed', 'error') and status.get('last_sync'):
|
||
history.append({
|
||
'source_name': name,
|
||
'status': status.get('status'),
|
||
'last_sync': status.get('last_sync'),
|
||
'files_synced': status.get('files_synced', 0),
|
||
'total_files': status.get('total_files', 0),
|
||
'error': status.get('error')
|
||
})
|
||
|
||
# 按时间倒序排列
|
||
history.sort(key=lambda x: x.get('last_sync', ''), reverse=True)
|
||
|
||
return history[:limit]
|
||
|
||
def add_sync_source(self, name, source_config):
|
||
"""添加同步源"""
|
||
with self.sync_lock:
|
||
self.sync_sources[name] = source_config
|
||
self.sync_status[name] = {
|
||
'last_sync': None,
|
||
'status': 'stopped',
|
||
'files_synced': 0,
|
||
'total_files': 0,
|
||
'error': None,
|
||
'schedule': source_config.get('schedule', {}), # 定时同步配置
|
||
'next_sync': None # 下次同步时间
|
||
}
|
||
self._calculate_next_sync(name)
|
||
self.save_sync_state()
|
||
self._save_sync_sources()
|
||
|
||
def _calculate_next_sync(self, name):
|
||
"""计算下次同步时间"""
|
||
try:
|
||
import croniter
|
||
except ImportError:
|
||
print("[WARN] croniter not installed")
|
||
return
|
||
|
||
if name not in self.sync_status:
|
||
return
|
||
|
||
schedule = self.sync_status[name].get('schedule', {})
|
||
if not schedule.get('enabled'):
|
||
self.sync_status[name]['next_sync'] = None
|
||
return
|
||
|
||
cron_str = schedule.get('cron')
|
||
if cron_str:
|
||
try:
|
||
cron = croniter.croniter(cron_str)
|
||
self.sync_status[name]['next_sync'] = cron.get_next(datetime).isoformat()
|
||
except Exception as e:
|
||
print(f"[WARN] croniter error: {e}")
|
||
self.sync_status[name]['next_sync'] = None
|
||
|
||
def _start_scheduler(self):
|
||
"""启动定时同步调度器"""
|
||
if self.scheduler_running:
|
||
return
|
||
self.scheduler_running = True
|
||
self.scheduler_thread = Thread(target=self._scheduler_loop, daemon=True)
|
||
self.scheduler_thread.start()
|
||
|
||
def _scheduler_loop(self):
|
||
"""定时同步调度循环"""
|
||
while self.scheduler_running:
|
||
try:
|
||
now = datetime.now()
|
||
due_sources = []
|
||
with self.sync_lock:
|
||
for name, status in self.sync_status.items():
|
||
source = self.sync_sources.get(name, {})
|
||
schedule = status.get('schedule', {}) or source.get('schedule', {})
|
||
if not schedule.get('enabled'):
|
||
continue
|
||
|
||
next_sync = status.get('next_sync')
|
||
if next_sync:
|
||
try:
|
||
next_time = datetime.fromisoformat(next_sync)
|
||
except ValueError:
|
||
continue
|
||
if now >= next_time:
|
||
due_sources.append(name)
|
||
|
||
# 锁外触发同步,避免与 start_sync 内的 sync_lock 形成死锁
|
||
for name in due_sources:
|
||
try:
|
||
print(f"[定时同步] 触发同步: {name}")
|
||
self.start_sync(name)
|
||
self._calculate_next_sync(name)
|
||
except Exception as e:
|
||
print(f"[定时同步] 触发 {name} 失败: {e}")
|
||
except Exception as e:
|
||
print(f"定时同步调度错误: {e}")
|
||
time.sleep(60) # 每分钟检查一次
|
||
|
||
def stop_scheduler(self):
|
||
"""停止定时同步调度器"""
|
||
self.scheduler_running = False
|
||
if self.scheduler_thread:
|
||
self.scheduler_thread.join(timeout=5)
|
||
|
||
def remove_sync_source(self, name):
|
||
"""移除同步源"""
|
||
with self.sync_lock:
|
||
if name in self.sync_sources:
|
||
del self.sync_sources[name]
|
||
if name in self.sync_status:
|
||
del self.sync_status[name]
|
||
|
||
# 锁外停止线程(stop_sync 内部会再取 sync_lock,避免不可重入锁死锁)
|
||
if name in self.sync_threads:
|
||
self.stop_sync(name)
|
||
|
||
with self.sync_lock:
|
||
self.save_sync_state()
|
||
self._save_sync_sources()
|
||
|
||
def start_sync(self, name):
|
||
"""开始同步指定源"""
|
||
if name not in self.sync_sources:
|
||
return False
|
||
|
||
with self.sync_lock:
|
||
existing = self.sync_threads.get(name)
|
||
if existing and existing.is_alive():
|
||
return True # 已经在运行(检查+记录在同一把锁内,防双 worker 竞态)
|
||
|
||
# 设置运行标志,确保同步循环可以执行
|
||
self.running = True
|
||
self.sync_status[name]['status'] = 'syncing'
|
||
self.sync_status[name]['error'] = None
|
||
|
||
thread = Thread(target=self._sync_worker, args=(name,), daemon=True)
|
||
self.sync_threads[name] = thread
|
||
thread.start()
|
||
return True
|
||
|
||
def sync_packages(self, source_name, packages):
|
||
"""临时单次同步指定源的特定包
|
||
|
||
Args:
|
||
source_name: 同步源名称(如 'pypi-mirrord')
|
||
packages: 包名列表(如 ['requests', 'numpy'])
|
||
|
||
Returns:
|
||
dict: 包含 success 状态和 task_id
|
||
"""
|
||
if source_name not in self.sync_sources:
|
||
return {"success": False, "error": f"同步源 '{source_name}' 不存在"}
|
||
|
||
if not packages or not isinstance(packages, list):
|
||
return {"success": False, "error": "请提供有效的包名列表"}
|
||
|
||
# 生成临时任务ID(随机后缀防同秒碰撞)
|
||
import time
|
||
import uuid
|
||
_ts = int(time.time())
|
||
task_id = f"temp_sync_{_ts}_{uuid.uuid4().hex[:6]}"
|
||
|
||
# 使用临时任务名进行同步
|
||
temp_name = f"{source_name}_temp_{_ts}_{uuid.uuid4().hex[:6]}"
|
||
|
||
# 设置状态
|
||
with self.sync_lock:
|
||
self.sync_status[temp_name] = {
|
||
'last_sync': time.strftime('%Y-%m-%dT%H:%M:%S'),
|
||
'status': 'syncing',
|
||
'files_synced': 0,
|
||
'total_files': 0,
|
||
'error': None,
|
||
'schedule': {},
|
||
'next_sync': None,
|
||
'is_temp_sync': True,
|
||
'source_name': source_name,
|
||
'packages': packages
|
||
}
|
||
|
||
# 设置运行标志
|
||
self.running = True
|
||
|
||
# 在后台线程执行临时同步
|
||
thread = Thread(
|
||
target=self._sync_worker_temp,
|
||
args=(temp_name, source_name, packages),
|
||
daemon=True
|
||
)
|
||
self.sync_threads[temp_name] = thread
|
||
thread.start()
|
||
|
||
return {"success": True, "task_id": task_id, "source": source_name, "packages": packages}
|
||
|
||
def _sync_worker_temp(self, temp_name, source_name, packages):
|
||
"""临时同步工作线程"""
|
||
print(f"[Temp Sync] 开始临时同步: {source_name}, 包: {', '.join(packages[:3])}{'...' if len(packages) > 3 else ''}")
|
||
|
||
config = self.sync_sources[source_name]
|
||
try:
|
||
self._sync_http(temp_name, config, specific_packages=packages)
|
||
with self.sync_lock:
|
||
if temp_name in self.sync_status:
|
||
self.sync_status[temp_name]['status'] = 'completed'
|
||
print(f"[Temp Sync] 临时同步完成: {source_name}")
|
||
except Exception as e:
|
||
with self.sync_lock:
|
||
if temp_name in self.sync_status:
|
||
self.sync_status[temp_name]['status'] = 'error'
|
||
self.sync_status[temp_name]['error'] = str(e)
|
||
print(f"[Temp Sync] 临时同步失败: {e}")
|
||
finally:
|
||
# 清理临时条目,防止 sync_state.json 无限膨胀
|
||
with self.sync_lock:
|
||
self.sync_status.pop(temp_name, None)
|
||
self.sync_threads.pop(temp_name, None)
|
||
self.save_sync_state()
|
||
|
||
def stop_sync(self, name):
|
||
"""停止同步指定源 - 请求停止
|
||
|
||
注意: 不删除 sync_threads 条目(删除后 start_sync 会再起一个 worker,
|
||
与仍在运行的旧线程形成双 worker)。线程结束后 is_alive() 为 False,
|
||
start_sync 会自然复用该条目。
|
||
"""
|
||
# 设置状态为停止(源可能已从 sync_status 移除,容错)
|
||
with self.sync_lock:
|
||
status = self.sync_status.get(name)
|
||
if status is not None:
|
||
status['status'] = 'stopped'
|
||
# 设置停止标志,让线程提前退出
|
||
self.running = False
|
||
# 短暂等待后重置运行标志
|
||
import time
|
||
time.sleep(0.5)
|
||
self.running = True
|
||
|
||
def start_all_sync(self):
|
||
"""开始所有同步源"""
|
||
for name in self.sync_sources:
|
||
self.start_sync(name)
|
||
|
||
def stop_all_sync(self):
|
||
"""停止所有同步源"""
|
||
for name in list(self.sync_threads.keys()):
|
||
self.stop_sync(name)
|
||
|
||
def _sync_worker(self, name):
|
||
"""同步工作线程"""
|
||
print(f"[SYNC] 开始同步: {name}")
|
||
source_config = self.sync_sources[name]
|
||
sync_type = source_config.get('type', 'http')
|
||
print(f"[SYNC] 类型: {sync_type}, URL: {self._redact_url(source_config.get('url', 'N/A'))}")
|
||
|
||
try:
|
||
if sync_type in ('http', 'https'):
|
||
self._sync_http(name, source_config)
|
||
elif sync_type == 'ftp':
|
||
self._sync_ftp(name, source_config)
|
||
elif sync_type == 'sftp':
|
||
self._sync_sftp(name, source_config)
|
||
elif sync_type == 'local':
|
||
self._sync_local(name, source_config)
|
||
# 新增同步类型
|
||
elif sync_type == 'rsync':
|
||
self._sync_rsync(name, source_config)
|
||
elif sync_type == 'git':
|
||
self._sync_git(name, source_config)
|
||
elif sync_type == 'aws' or sync_type == 's3':
|
||
self._sync_s3(name, source_config)
|
||
elif sync_type == 'oss':
|
||
self._sync_oss(name, source_config)
|
||
elif sync_type == 'cos':
|
||
self._sync_cos(name, source_config)
|
||
elif sync_type == 'webdav':
|
||
self._sync_webdav(name, source_config)
|
||
elif sync_type == 'rsync':
|
||
self._sync_rsync(name, source_config)
|
||
else:
|
||
raise ValueError(f"不支持的同步类型: {sync_type}")
|
||
|
||
with self.sync_lock:
|
||
if not self.running:
|
||
# 被 stop_sync 中断:标记为停止而非完成
|
||
self.sync_status[name]['status'] = 'stopped'
|
||
else:
|
||
self.sync_status[name]['status'] = 'completed'
|
||
self.sync_status[name]['last_sync'] = datetime.now().isoformat()
|
||
|
||
except Exception as e:
|
||
with self.sync_lock:
|
||
self.sync_status[name]['status'] = 'error'
|
||
self.sync_status[name]['error'] = str(e)
|
||
|
||
finally:
|
||
# 清理线程条目(start_sync 通过 is_alive 判断, 清理后允许重新 start)
|
||
with self.sync_lock:
|
||
self.sync_threads.pop(name, None)
|
||
self.save_sync_state()
|
||
|
||
def _sync_http(self, name, config, specific_packages=None):
|
||
"""HTTP/HTTPS同步 - 完整版
|
||
|
||
Args:
|
||
name: 同步源名称
|
||
config: 同步源配置
|
||
specific_packages: 可选的特定包列表,用于临时单次同步
|
||
"""
|
||
print(f"[HTTP Sync] 开始同步: {name}")
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
print(f"[HTTP Sync] 目标目录: {target_dir}")
|
||
|
||
# 获取文件列表
|
||
source_url = config.get('url', '')
|
||
print(f"[HTTP Sync] 源URL: {self._redact_url(source_url)}")
|
||
file_list = self._get_http_file_list(source_url, config, specific_packages)
|
||
|
||
total_files = len(file_list)
|
||
print(f"[HTTP Sync] 获取到 {total_files} 个文件")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['total_files'] = total_files
|
||
|
||
if total_files == 0:
|
||
print(f"[HTTP Sync] 没有文件需要同步")
|
||
return
|
||
|
||
synced_count = 0
|
||
failed_count = 0
|
||
|
||
for i, file_info in enumerate(file_list):
|
||
if not self.running:
|
||
print(f"[HTTP Sync] 同步被中断")
|
||
break
|
||
|
||
# 更新进度
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
|
||
# 获取本地路径(先做路径穿越校验)
|
||
filename = file_info.get('name', '')
|
||
if not filename:
|
||
continue
|
||
if self._safe_remote_name(filename) is None:
|
||
print(f"[HTTP Sync] 跳过不安全的文件名: {filename!r}")
|
||
continue
|
||
|
||
# 处理子目录
|
||
local_path = os.path.join(target_dir, filename)
|
||
local_dir = os.path.dirname(local_path)
|
||
if local_dir and not os.path.exists(local_dir):
|
||
os.makedirs(local_dir, exist_ok=True)
|
||
|
||
# 检查是否需要同步
|
||
if self._need_sync_http(file_info, local_path):
|
||
try:
|
||
if self._download_file_http(file_info['url'], local_path, config):
|
||
synced_count += 1
|
||
else:
|
||
failed_count += 1
|
||
except Exception as e:
|
||
print(f"[HTTP Sync] 下载失败 {filename}: {e}")
|
||
failed_count += 1
|
||
|
||
# 每10个文件输出一次进度
|
||
if (i + 1) % 10 == 0:
|
||
print(f"[HTTP Sync] 进度: {i+1}/{total_files}, 已同步: {synced_count}")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
|
||
print(f"[HTTP Sync] 完成: 成功 {synced_count}, 失败 {failed_count}")
|
||
|
||
@staticmethod
|
||
def _safe_remote_name(filename):
|
||
"""校验远程文件名/相对路径,拒绝路径穿越(.. 段、绝对路径、反斜杠)"
|
||
|
||
允许 '/' 分隔的多级子路径(HTTP 列表常见),但拒绝任何 '..' 段。
|
||
"""
|
||
if not filename or filename in ('.', '..'):
|
||
return None
|
||
if filename.startswith('/') or '\\' in filename:
|
||
return None
|
||
for seg in filename.split('/'):
|
||
if seg in ('', '.', '..'):
|
||
return None
|
||
return filename
|
||
|
||
@staticmethod
|
||
def _redact_url(url):
|
||
"""脱敏 URL 中的 user:pass@ 段(防止凭证进日志/状态文件)"""
|
||
if not url:
|
||
return url
|
||
import re as _re
|
||
return _re.sub(r'(://)([^/@\s]+)@', r'\1***@', str(url))
|
||
|
||
def _sync_ftp(self, name, config):
|
||
"""FTP同步实现"""
|
||
try:
|
||
host = config['host']
|
||
port = config.get('port', 21)
|
||
username = config.get('username', 'anonymous')
|
||
password = config.get('password', 'anonymous@')
|
||
remote_path = config.get('remote_path', '/')
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
print(f"开始FTP同步 {name} -> {target_dir}")
|
||
|
||
ftp = ftplib.FTP()
|
||
ftp.connect(host, port)
|
||
ftp.login(username, password)
|
||
|
||
if remote_path:
|
||
ftp.cwd(remote_path)
|
||
|
||
file_list = []
|
||
ftp.retrlines('LIST', file_list.append)
|
||
|
||
remote_files = {}
|
||
for line in file_list:
|
||
parts = line.split()
|
||
if len(parts) < 9:
|
||
continue
|
||
filename = ' '.join(parts[8:])
|
||
if filename in ['.', '..']:
|
||
continue
|
||
is_dir = parts[0].startswith('d')
|
||
size = int(parts[4]) if not is_dir else 0
|
||
remote_files[filename] = {
|
||
'name': filename,
|
||
'size': size,
|
||
'is_dir': is_dir
|
||
}
|
||
|
||
synced_count = 0
|
||
for filename, file_info in remote_files.items():
|
||
if not self.running:
|
||
break
|
||
safe_name = self._safe_remote_name(filename)
|
||
if safe_name is None:
|
||
continue
|
||
local_path = os.path.join(target_dir, safe_name)
|
||
if file_info['is_dir']:
|
||
sub_config = config.copy()
|
||
sub_config['remote_path'] = os.path.join(remote_path, safe_name).replace('\\', '/')
|
||
sub_config['target'] = os.path.join(config.get('target', name), safe_name)
|
||
# 递归沿用顶层 name,避免子目录用假名导致 sync_status KeyError
|
||
self._sync_ftp(name, sub_config)
|
||
else:
|
||
if self._need_sync_ftp(ftp, safe_name, local_path, file_info):
|
||
if self._download_ftp_file(ftp, safe_name, local_path, file_info):
|
||
synced_count += 1
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
self.sync_status[name]['total_files'] = len([f for f in remote_files.values() if not f['is_dir']])
|
||
|
||
except Exception as e:
|
||
# 确保连接被关闭
|
||
try:
|
||
ftp.quit()
|
||
except Exception:
|
||
pass
|
||
raise Exception(f"FTP同步失败: {str(e)}")
|
||
|
||
finally:
|
||
try:
|
||
ftp.quit()
|
||
except Exception:
|
||
pass
|
||
|
||
def _sync_sftp(self, name, config):
|
||
"""SFTP同步实现"""
|
||
if not paramiko:
|
||
raise ImportError("paramiko未安装,无法使用SFTP同步")
|
||
|
||
host = config['host']
|
||
port = config.get('port', 22)
|
||
username = config.get('username', 'anonymous')
|
||
password = config.get('password')
|
||
private_key = config.get('private_key')
|
||
remote_path = config.get('remote_path', '/')
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
print(f"开始SFTP同步 {name} -> {target_dir}")
|
||
|
||
ssh = paramiko.SSHClient()
|
||
ssh.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||
|
||
connect_kwargs = {'hostname': host, 'port': port, 'username': username}
|
||
if private_key:
|
||
key = paramiko.RSAKey.from_private_key_file(private_key)
|
||
connect_kwargs['pkey'] = key
|
||
else:
|
||
connect_kwargs['password'] = password
|
||
|
||
ssh.connect(**connect_kwargs)
|
||
sftp = ssh.open_sftp()
|
||
|
||
try:
|
||
synced_count = self._sync_sftp_directory(sftp, remote_path, target_dir, name)
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
|
||
finally:
|
||
try:
|
||
sftp.close()
|
||
except Exception:
|
||
pass
|
||
try:
|
||
ssh.close()
|
||
except Exception:
|
||
pass
|
||
|
||
def _sync_sftp_directory(self, sftp, remote_path, local_path, sync_name):
|
||
"""递归同步SFTP目录"""
|
||
synced_count = 0
|
||
os.makedirs(local_path, exist_ok=True)
|
||
|
||
file_list = sftp.listdir_attr(remote_path)
|
||
|
||
for file_attr in file_list:
|
||
if not self.running:
|
||
break
|
||
remote_filename = self._safe_remote_name(file_attr.filename)
|
||
if remote_filename is None:
|
||
continue
|
||
remote_filepath = os.path.join(remote_path, remote_filename).replace('\\', '/')
|
||
local_filepath = os.path.join(local_path, remote_filename)
|
||
|
||
if file_attr.st_mode & 0o40000:
|
||
sub_synced = self._sync_sftp_directory(sftp, remote_filepath, local_filepath, sync_name)
|
||
synced_count += sub_synced
|
||
else:
|
||
if self._need_sync_sftp(file_attr, local_filepath):
|
||
if self._download_sftp_file(sftp, remote_filepath, local_filepath):
|
||
synced_count += 1
|
||
|
||
return synced_count
|
||
|
||
def _sync_local(self, name, config):
|
||
"""本地目录同步"""
|
||
print(f"[Local Sync] 开始同步: {name}")
|
||
source_dir = config['path']
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
|
||
if not os.path.exists(source_dir):
|
||
raise ValueError(f"源目录不存在: {source_dir}")
|
||
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
print(f"[Local Sync] 源: {source_dir} -> 目标: {target_dir}")
|
||
|
||
# 统计文件总数
|
||
all_files = []
|
||
for root, dirs, files in os.walk(source_dir):
|
||
for file in files:
|
||
all_files.append(os.path.join(root, file))
|
||
|
||
total_files = len(all_files)
|
||
print(f"[Local Sync] 共 {total_files} 个文件")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['total_files'] = total_files
|
||
|
||
synced_count = 0
|
||
for i, source_file in enumerate(all_files):
|
||
if not self.running:
|
||
break
|
||
|
||
relative_path = os.path.relpath(source_file, source_dir)
|
||
target_file = os.path.join(target_dir, relative_path)
|
||
|
||
if self._need_sync_local(source_file, target_file):
|
||
# 确保目标目录存在
|
||
target_file_dir = os.path.dirname(target_file)
|
||
if not os.path.exists(target_file_dir):
|
||
os.makedirs(target_file_dir, exist_ok=True)
|
||
|
||
try:
|
||
shutil.copy2(source_file, target_file)
|
||
synced_count += 1
|
||
except Exception as e:
|
||
print(f"[Local Sync] 复制失败 {relative_path}: {e}")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
|
||
if (i + 1) % 100 == 0:
|
||
print(f"[Local Sync] 进度: {i+1}/{total_files}")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
print(f"[Local Sync] 完成: 成功同步 {synced_count} 个文件")
|
||
|
||
def _sync_rsync(self, name, config):
|
||
"""Rsync同步 - 需要系统安装rsync命令"""
|
||
import subprocess
|
||
|
||
source = config.get('source', config.get('url', ''))
|
||
target_rel = config.get('target', name) or name
|
||
if target_rel in ('.', './', '') or os.path.isabs(target_rel):
|
||
raise Exception(f"Rsync 同步目标不安全: {target_rel!r} (禁止同步到 '.' 或绝对路径, 防止 --delete 清空目录)")
|
||
target_dir = os.path.join(self.config['base_dir'], target_rel)
|
||
# 目标目录不得等于 base_dir 本身
|
||
if os.path.realpath(target_dir) == os.path.realpath(self.config['base_dir']):
|
||
raise Exception("Rsync 同步目标不允许是下载根目录本身 (--delete 会清空全部文件)")
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
|
||
print(f"开始Rsync同步 {name} -> {target_dir}")
|
||
|
||
# rsync选项
|
||
options = ['-avz', '--progress', '--delete']
|
||
if config.get('exclude'):
|
||
for pattern in config['exclude']:
|
||
options.extend(['--exclude', pattern])
|
||
|
||
# 添加SSH选项(如果需要)
|
||
if config.get('ssh'):
|
||
options.extend(['-e', 'ssh'])
|
||
|
||
rsync_cmd = ['rsync'] + options + [source, target_dir]
|
||
|
||
try:
|
||
result = subprocess.run(
|
||
rsync_cmd,
|
||
capture_output=True,
|
||
text=True,
|
||
timeout=config.get('timeout', 3600)
|
||
)
|
||
if result.returncode != 0:
|
||
raise Exception(f"rsync失败: {result.stderr}")
|
||
except FileNotFoundError:
|
||
raise Exception("rsync命令未安装,请运行: apt install rsync")
|
||
|
||
def _sync_git(self, name, config):
|
||
"""Git仓库同步 - 克隆或更新Git仓库"""
|
||
import subprocess
|
||
|
||
repo_url = config.get('url')
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
branch = config.get('branch', 'main')
|
||
depth = config.get('depth', 1) # 浅克隆
|
||
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
repo_path = os.path.join(target_dir, '.repo')
|
||
|
||
print(f"开始Git同步 {name} -> {target_dir}")
|
||
|
||
if os.path.exists(os.path.join(target_dir, '.git')):
|
||
# 已存在,执行git pull
|
||
try:
|
||
git_timeout = config.get('git_timeout', 1800)
|
||
subprocess.run(['git', 'fetch', '--all'], cwd=target_dir, check=True,
|
||
capture_output=True, timeout=git_timeout)
|
||
subprocess.run(['git', 'reset', '--hard', f'origin/{branch}'], cwd=target_dir,
|
||
check=True, capture_output=True, timeout=git_timeout)
|
||
except subprocess.CalledProcessError as e:
|
||
raise Exception(f"Git pull失败: {e}")
|
||
except subprocess.TimeoutExpired:
|
||
raise Exception(f"Git pull超时({git_timeout}s)")
|
||
else:
|
||
# 克隆新仓库
|
||
cmd = ['git', 'clone']
|
||
if depth:
|
||
cmd.extend(['--depth', str(depth)])
|
||
if branch != 'main':
|
||
cmd.extend(['-b', branch])
|
||
cmd.extend([repo_url, target_dir])
|
||
|
||
try:
|
||
git_timeout = config.get('git_timeout', 1800)
|
||
subprocess.run(cmd, check=True, capture_output=True, timeout=git_timeout)
|
||
except subprocess.CalledProcessError as e:
|
||
raise Exception(f"Git clone失败: {e}")
|
||
except subprocess.TimeoutExpired:
|
||
raise Exception(f"Git clone超时({git_timeout}s)")
|
||
|
||
def _sync_s3(self, name, config):
|
||
"""AWS S3兼容存储同步 (S3/OSS/COS/MinIO等)"""
|
||
try:
|
||
import boto3
|
||
except ImportError:
|
||
raise Exception("请安装boto3库: pip install boto3")
|
||
|
||
print(f"[S3 Sync] 开始同步: {name}")
|
||
endpoint = config.get('endpoint', config.get('url'))
|
||
bucket = config.get('bucket')
|
||
access_key = config.get('access_key', config.get('aws_access_key_id'))
|
||
secret_key = config.get('secret_key', config.get('aws_secret_access_key'))
|
||
region = config.get('region', 'us-east-1')
|
||
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
print(f"[S3 Sync] 目标目录: {target_dir}")
|
||
|
||
# 创建S3客户端
|
||
s3_client = boto3.client(
|
||
's3',
|
||
endpoint_url=endpoint,
|
||
aws_access_key_id=access_key,
|
||
aws_secret_access_key=secret_key,
|
||
region_name=region
|
||
)
|
||
|
||
prefix = config.get('prefix', '')
|
||
|
||
# 首先统计文件数量
|
||
total_files = 0
|
||
try:
|
||
paginator = s3_client.get_paginator('list_objects_v2')
|
||
for page in paginator.paginate(Bucket=bucket, Prefix=prefix):
|
||
if 'Contents' in page:
|
||
total_files += len([o for o in page['Contents'] if not o['Key'].endswith('/')])
|
||
except Exception as e:
|
||
print(f"[S3 Sync] 统计文件数失败: {e}")
|
||
|
||
print(f"[S3 Sync] 共 {total_files} 个文件")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['total_files'] = total_files
|
||
|
||
synced_count = 0
|
||
|
||
try:
|
||
paginator = s3_client.get_paginator('list_objects_v2')
|
||
for page in paginator.paginate(Bucket=bucket, Prefix=prefix):
|
||
if 'Contents' not in page:
|
||
continue
|
||
|
||
for obj in page['Contents']:
|
||
if not self.running:
|
||
break
|
||
|
||
key = obj['Key']
|
||
if key.endswith('/'):
|
||
continue
|
||
|
||
local_path = os.path.join(target_dir, key[len(prefix):].lstrip('/'))
|
||
os.makedirs(os.path.dirname(local_path), exist_ok=True)
|
||
|
||
# 检查是否需要下载
|
||
need_download = True
|
||
if os.path.exists(local_path):
|
||
local_size = os.path.getsize(local_path)
|
||
if local_size == obj['Size']:
|
||
need_download = False
|
||
|
||
if need_download:
|
||
try:
|
||
s3_client.download_file(bucket, key, local_path)
|
||
synced_count += 1
|
||
except Exception as e:
|
||
print(f"[S3 Sync] 下载失败 {key}: {e}")
|
||
|
||
with self.sync_lock:
|
||
self.sync_status[name]['files_synced'] = synced_count
|
||
|
||
except Exception as e:
|
||
raise Exception(f"S3同步失败: {e}")
|
||
|
||
print(f"[S3 Sync] 完成: 成功同步 {synced_count} 个文件")
|
||
|
||
def _sync_oss(self, name, config):
|
||
"""阿里云OSS同步"""
|
||
config['type'] = 's3' # 复用S3逻辑
|
||
config['endpoint'] = config.get('endpoint', f"https://{config.get('bucket')}.oss-{config.get('region', 'cn-hangzhou')}.aliyuncs.com")
|
||
self._sync_s3(name, config)
|
||
|
||
def _sync_cos(self, name, config):
|
||
"""腾讯云COS同步"""
|
||
config['type'] = 's3' # 复用S3逻辑
|
||
config['endpoint'] = config.get('endpoint', f"https://{config.get('bucket')}.cos.{config.get('region', 'ap-guangzhou')}.myqcloud.com")
|
||
self._sync_s3(name, config)
|
||
|
||
def _sync_webdav(self, name, config):
|
||
"""WebDAV同步"""
|
||
import urllib.request
|
||
|
||
url = config.get('url', config.get('server'))
|
||
username = config.get('username')
|
||
password = config.get('password')
|
||
remote_path = config.get('remote_path', '/')
|
||
|
||
target_dir = os.path.join(self.config['base_dir'], config.get('target', name))
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
|
||
print(f"开始WebDAV同步 {name} -> {target_dir}")
|
||
|
||
# 构建认证
|
||
if username and password:
|
||
import base64
|
||
auth_string = f"{username}:{password}"
|
||
encoded_auth = base64.b64encode(auth_string.encode()).decode()
|
||
|
||
# PROPFIND获取文件列表
|
||
req = urllib.request.Request(
|
||
f"{url}{remote_path}",
|
||
method='PROPFIND',
|
||
headers={'Authorization': f'Basic {encoded_auth}'}
|
||
)
|
||
|
||
# 简化实现 - 实际需要解析XML
|
||
raise Exception("WebDAV同步需要完整实现,目前仅支持基础配置")
|
||
|
||
def _get_http_file_list(self, base_url, config, specific_packages=None):
|
||
"""获取HTTP文件列表 - 支持PyPI等复杂索引
|
||
|
||
Args:
|
||
base_url: 源URL
|
||
config: 配置信息
|
||
specific_packages: 可选的特定包列表
|
||
"""
|
||
import urllib.request
|
||
import re
|
||
|
||
file_list = []
|
||
|
||
# 检测是否为PyPI simple索引
|
||
is_pypi = 'pypi' in base_url.lower() or config.get('is_pypi', False)
|
||
|
||
if is_pypi:
|
||
print(f"[HTTP] 检测到PyPI索引,使用PyPI专用解析")
|
||
return self._get_pypi_file_list(base_url, config, specific_packages, sync_name=name)
|
||
|
||
if 'api_url' in config:
|
||
try:
|
||
req = urllib.request.Request(config['api_url'])
|
||
if 'username' in config and 'password' in config:
|
||
auth_string = f"{config['username']}:{config['password']}"
|
||
encoded_auth = base64.b64encode(auth_string.encode()).decode()
|
||
req.add_header('Authorization', f'Basic {encoded_auth}')
|
||
|
||
with urllib.request.urlopen(req, timeout=60) as response:
|
||
api_data = json.loads(response.read().decode())
|
||
if 'files' in api_data:
|
||
for file_info in api_data['files']:
|
||
file_list.append({
|
||
'name': file_info.get('name', ''),
|
||
'url': file_info.get('url', ''),
|
||
'size': file_info.get('size', 0)
|
||
})
|
||
return file_list
|
||
except Exception as e:
|
||
print(f"API获取文件列表失败: {e}")
|
||
|
||
if config.get('parse_html', True):
|
||
try:
|
||
file_list = self._fetch_http_directory(base_url, config, max_depth=3)
|
||
return file_list
|
||
except Exception as e:
|
||
print(f"HTML解析文件列表失败: {e}")
|
||
|
||
return file_list
|
||
|
||
def _get_pypi_file_list(self, base_url, config, specific_packages=None, sync_name=None):
|
||
"""获取PyPI文件列表
|
||
|
||
Args:
|
||
base_url: PyPI镜像源URL
|
||
config: 配置信息
|
||
specific_packages: 可选的特定包列表,如果提供则只获取这些包的文件
|
||
"""
|
||
import urllib.request
|
||
import urllib.parse
|
||
import re
|
||
|
||
print(f"[PyPI] 获取PyPI包列表: {base_url}")
|
||
|
||
# 获取包名列表
|
||
package_index_url = base_url.rstrip('/') + '/'
|
||
if not package_index_url.endswith('/simple/') and not package_index_url.endswith('/simple'):
|
||
if package_index_url.endswith('/'):
|
||
package_index_url += 'simple/'
|
||
else:
|
||
package_index_url += '/simple/'
|
||
|
||
# 获取包列表
|
||
if specific_packages:
|
||
# 指定了特定包,直接使用
|
||
packages = list(specific_packages)
|
||
print(f"[PyPI] 将只同步指定 {len(packages)} 个包: {', '.join(packages[:5])}{'...' if len(packages) > 5 else ''}")
|
||
else:
|
||
# 获取所有包
|
||
try:
|
||
# 添加User-Agent避免被拒绝
|
||
req = urllib.request.Request(package_index_url, headers={'User-Agent': 'MirrorSync/1.0'})
|
||
with urllib.request.urlopen(req, timeout=30) as response:
|
||
html_content = response.read().decode('utf-8', errors='ignore')
|
||
|
||
# 提取所有包名 - 兼容不同格式
|
||
# 格式1: <a href="/simple/package/">package</a>
|
||
package_pattern = r'<a[^>]+href="/simple/([^/"]+)/"[^>]*>'
|
||
packages = re.findall(package_pattern, html_content)
|
||
|
||
if not packages:
|
||
# 格式2: <a href="package/">package</a> (清华源格式)
|
||
package_pattern = r'<a[^>]+href="([^/"]+)/"[^>]*>'
|
||
packages = re.findall(package_pattern, html_content)
|
||
|
||
# 过滤掉非包名
|
||
packages = [p for p in packages if p and not p.startswith('..')]
|
||
|
||
print(f"[PyPI] 发现 {len(packages)} 个包")
|
||
|
||
# 限制包数量
|
||
max_packages = config.get('max_packages', 50)
|
||
packages = packages[:max_packages]
|
||
print(f"[PyPI] 将获取前 {max_packages} 个包的文件")
|
||
except Exception as e:
|
||
print(f"[PyPI] 获取包列表失败: {e}")
|
||
return []
|
||
|
||
file_list = []
|
||
success_count = 0
|
||
fail_count = 0
|
||
|
||
for i, package_name in enumerate(packages):
|
||
if not self.running:
|
||
break
|
||
|
||
# 更新进度(只更新当前源,避免污染其他源状态)
|
||
with self.sync_lock:
|
||
if sync_name and sync_name in self.sync_status:
|
||
self.sync_status[sync_name]['files_synced'] = i
|
||
|
||
# 获取每个包的文件列表
|
||
package_url = f"{package_index_url}{package_name}/"
|
||
try:
|
||
req = urllib.request.Request(package_url, headers={'User-Agent': 'MirrorSync/1.0'})
|
||
with urllib.request.urlopen(req, timeout=30) as response:
|
||
pkg_html = response.read().decode('utf-8', errors='ignore')
|
||
|
||
# 提取whl文件 - 允许前面有路径,后面可能有 #sha256= 锚点
|
||
file_pattern = r'<a[^>]+href="([^"]*\.whl)(?:#[^"]*)?"[^>]*>'
|
||
whl_files = re.findall(file_pattern, pkg_html)
|
||
|
||
# 也提取tar.gz文件 - 允许前面有路径,后面可能有 #sha256= 锚点
|
||
file_pattern = r'<a[^>]+href="([^"]*\.tar\.gz)(?:#[^"]*)?"[^>]*>'
|
||
tar_files = re.findall(file_pattern, pkg_html)
|
||
whl_files.extend(tar_files)
|
||
|
||
# 也提取zip文件 - 允许前面有路径,后面可能有 #sha256= 锚点
|
||
file_pattern = r'<a[^>]+href="([^"]*\.zip)(?:#[^"]*)?"[^>]*>'
|
||
zip_files = re.findall(file_pattern, pkg_html)
|
||
whl_files.extend(zip_files)
|
||
|
||
if whl_files:
|
||
success_count += 1
|
||
for filename in whl_files:
|
||
# 使用urljoin正确处理相对路径
|
||
# PyPI页面中的href可能是 "../../packages/.../file.tar.gz" 格式
|
||
file_url = urllib.parse.urljoin(package_url, filename)
|
||
file_list.append({
|
||
'name': f"{package_name}/{filename.split('/')[-1]}",
|
||
'url': file_url,
|
||
'size': 0,
|
||
'is_package_file': True
|
||
})
|
||
else:
|
||
fail_count += 1
|
||
|
||
except Exception as e:
|
||
fail_count += 1
|
||
if fail_count <= 3: # 只显示前几个错误
|
||
print(f"[PyPI] 获取包 {package_name} 文件列表失败: {e}")
|
||
elif fail_count == 4:
|
||
print(f"[PyPI] 更多错误不再显示...")
|
||
|
||
print(f"[PyPI] 获取完成: 成功 {success_count} 个包, 失败 {fail_count} 个包, 共 {len(file_list)} 个文件")
|
||
return file_list
|
||
|
||
def _fetch_http_directory(self, base_url, config, max_depth=3, current_depth=0):
|
||
"""递归获取HTTP目录内容"""
|
||
import urllib.request
|
||
import re
|
||
|
||
if current_depth > max_depth:
|
||
return []
|
||
|
||
file_list = []
|
||
|
||
try:
|
||
req = urllib.request.Request(base_url)
|
||
with urllib.request.urlopen(req, timeout=30) as response:
|
||
html_content = response.read().decode('utf-8', errors='ignore')
|
||
|
||
# 更健壮的链接提取
|
||
link_pattern = r'<a[^>]+href=["\']([^"\']+)["\'][^>]*>([^<]*)</a>'
|
||
matches = re.findall(link_pattern, html_content)
|
||
|
||
for href, text in matches:
|
||
href = href.strip()
|
||
if not href or href in ['../', './', '/', ''] or href.startswith('?'):
|
||
continue
|
||
if href.startswith('#'):
|
||
continue
|
||
|
||
# 构建完整URL
|
||
if href.startswith('http'):
|
||
file_url = href
|
||
else:
|
||
file_url = base_url.rstrip('/') + '/' + href.lstrip('/')
|
||
|
||
# 获取文件名
|
||
filename = href.split('/')[-1] if '/' in href else href
|
||
|
||
# 判断是目录还是文件
|
||
is_directory = href.endswith('/')
|
||
|
||
if is_directory:
|
||
# 递归获取子目录
|
||
sub_files = self._fetch_http_directory(file_url, config, max_depth, current_depth + 1)
|
||
file_list.extend(sub_files)
|
||
else:
|
||
# 添加文件
|
||
file_list.append({
|
||
'name': filename,
|
||
'url': file_url,
|
||
'size': 0
|
||
})
|
||
|
||
except Exception as e:
|
||
print(f"获取目录 {base_url} 失败: {e}")
|
||
|
||
return file_list
|
||
|
||
def _need_sync_http(self, file_info, local_path):
|
||
"""检查HTTP文件是否需要同步"""
|
||
if not os.path.exists(local_path):
|
||
return True
|
||
remote_size = file_info.get('size', 0)
|
||
if remote_size <= 0:
|
||
# 远端大小未知(PyPI/HTML 列表恒为 0):仅按存在性判断,避免每次全量重下
|
||
return False
|
||
local_size = os.path.getsize(local_path)
|
||
return local_size != remote_size
|
||
|
||
def _need_sync_ftp(self, ftp, remote_filename, local_path, file_info):
|
||
"""检查FTP文件是否需要同步"""
|
||
if not os.path.exists(local_path):
|
||
return True
|
||
local_size = os.path.getsize(local_path)
|
||
remote_size = file_info['size']
|
||
return local_size != remote_size
|
||
|
||
def _need_sync_sftp(self, file_attr, local_path):
|
||
"""检查SFTP文件是否需要同步"""
|
||
if not os.path.exists(local_path):
|
||
return True
|
||
local_size = os.path.getsize(local_path)
|
||
remote_size = file_attr.st_size
|
||
if local_size != remote_size:
|
||
return True
|
||
local_mtime = os.path.getmtime(local_path)
|
||
remote_mtime = file_attr.st_mtime
|
||
return abs(local_mtime - remote_mtime) > 1
|
||
|
||
def _need_sync_local(self, source_path, target_path):
|
||
"""检查本地文件是否需要同步"""
|
||
if not os.path.exists(target_path):
|
||
return True
|
||
source_mtime = os.path.getmtime(source_path)
|
||
target_mtime = os.path.getmtime(target_path)
|
||
return source_mtime > target_mtime or os.path.getsize(source_path) != os.path.getsize(target_path)
|
||
|
||
def _download_file_http(self, url, local_path, config):
|
||
"""下载HTTP文件(支持断点续传,续传数据追加到已有前缀,避免损坏)"""
|
||
import urllib.request
|
||
import base64
|
||
import shutil
|
||
from email.utils import parsedate
|
||
|
||
temp_path = local_path + '.tmp' # 提前初始化,避免 except 分支 NameError
|
||
|
||
try:
|
||
req = urllib.request.Request(url)
|
||
headers = config.get('headers', {})
|
||
for key, value in headers.items():
|
||
req.add_header(key, value)
|
||
|
||
if 'username' in config and 'password' in config:
|
||
auth_string = f"{config['username']}:{config['password']}"
|
||
encoded_auth = base64.b64encode(auth_string.encode()).decode()
|
||
req.add_header('Authorization', f'Basic {encoded_auth}')
|
||
|
||
timeout = config.get('timeout', 30)
|
||
|
||
start_byte = 0
|
||
if os.path.exists(local_path):
|
||
start_byte = os.path.getsize(local_path)
|
||
if start_byte > 0:
|
||
req.add_header('Range', f'bytes={start_byte}-')
|
||
|
||
|
||
with urllib.request.urlopen(req, timeout=timeout) as response:
|
||
status_code = response.getcode()
|
||
if status_code not in [200, 206]:
|
||
raise Exception(f"HTTP错误: {status_code}")
|
||
|
||
if start_byte > 0 and status_code != 206:
|
||
# 服务器忽略了 Range,返回完整内容:从头开始写
|
||
if os.path.exists(temp_path):
|
||
os.remove(temp_path)
|
||
start_byte = 0
|
||
|
||
mode = 'ab' if start_byte > 0 else 'wb'
|
||
|
||
if start_byte > 0:
|
||
# 续传: 先把已有前缀复制进临时文件,再追加尾部数据,
|
||
# 避免用"仅尾部"覆盖完整文件导致数据损坏
|
||
shutil.copyfile(local_path, temp_path)
|
||
|
||
with open(temp_path, mode) as f:
|
||
while True:
|
||
if not self.running:
|
||
break
|
||
chunk = response.read(8192)
|
||
if not chunk:
|
||
break
|
||
f.write(chunk)
|
||
|
||
if self.running:
|
||
# 校验续传结果: 续传后文件应不小于原有大小
|
||
if start_byte > 0:
|
||
try:
|
||
final_size = os.path.getsize(temp_path)
|
||
if final_size < start_byte:
|
||
raise Exception(f"续传后文件异常变小: {final_size} < {start_byte}")
|
||
except OSError:
|
||
raise
|
||
if os.path.exists(local_path):
|
||
os.remove(local_path)
|
||
os.rename(temp_path, local_path)
|
||
|
||
last_modified = response.headers.get('Last-Modified')
|
||
if last_modified:
|
||
try:
|
||
timestamp = time.mktime(parsedate(last_modified))
|
||
os.utime(local_path, (timestamp, timestamp))
|
||
except:
|
||
pass
|
||
return True
|
||
else:
|
||
if os.path.exists(temp_path):
|
||
os.remove(temp_path)
|
||
return False
|
||
|
||
except Exception as e:
|
||
print(f"下载HTTP文件失败 {self._redact_url(url)}: {e}")
|
||
if os.path.exists(temp_path):
|
||
os.remove(temp_path)
|
||
return False
|
||
|
||
def _download_ftp_file(self, ftp, remote_filename, local_path, file_info):
|
||
"""下载FTP文件"""
|
||
try:
|
||
mode = 'wb'
|
||
start_pos = 0
|
||
if os.path.exists(local_path):
|
||
start_pos = os.path.getsize(local_path)
|
||
if start_pos < file_info['size']:
|
||
mode = 'ab'
|
||
else:
|
||
return True
|
||
|
||
# 防 FTP 命令注入: 文件名含 CR/LF 时拒绝
|
||
if any(c in remote_filename for c in ('\r', '\n')):
|
||
raise Exception(f"文件名含非法字符: {remote_filename!r}")
|
||
|
||
with open(local_path, mode) as f:
|
||
if start_pos > 0:
|
||
ftp.voidcmd('TYPE I')
|
||
ftp.retrbinary(f'RETR {remote_filename}', f.write, rest=start_pos)
|
||
else:
|
||
ftp.retrbinary(f'RETR {remote_filename}', f.write)
|
||
|
||
return True
|
||
except Exception as e:
|
||
print(f"下载FTP文件失败 {remote_filename}: {e}")
|
||
return False
|
||
|
||
def _download_sftp_file(self, sftp, remote_path, local_path):
|
||
"""下载SFTP文件"""
|
||
try:
|
||
start_pos = 0
|
||
if os.path.exists(local_path):
|
||
start_pos = os.path.getsize(local_path)
|
||
|
||
file_attr = sftp.stat(remote_path)
|
||
remote_size = file_attr.st_size
|
||
|
||
if start_pos >= remote_size:
|
||
return True
|
||
|
||
with sftp.open(remote_path, 'rb') as remote_file:
|
||
if start_pos > 0:
|
||
remote_file.seek(start_pos)
|
||
|
||
with open(local_path, 'ab' if start_pos > 0 else 'wb') as local_file:
|
||
while True:
|
||
if not self.running:
|
||
break
|
||
chunk = remote_file.read(8192)
|
||
if not chunk:
|
||
break
|
||
local_file.write(chunk)
|
||
|
||
os.utime(local_path, (file_attr.st_atime, file_attr.st_mtime))
|
||
return True
|
||
except Exception as e:
|
||
print(f"下载SFTP文件失败 {remote_path}: {e}")
|
||
return False
|
||
|
||
def get_sync_status(self):
|
||
"""获取同步状态"""
|
||
with self.sync_lock:
|
||
return self.sync_status.copy()
|
||
|
||
def get_source_status(self, name):
|
||
"""获取指定同步源状态"""
|
||
with self.sync_lock:
|
||
return dict(self.sync_status.get(name, {}))
|
||
|
||
def stop_all_tasks_for_source(self, name):
|
||
"""停止指定同步源的所有任务(v2 API 兼容)"""
|
||
return self.stop_sync(name)
|
||
|
||
def start(self):
|
||
"""启动同步管理器"""
|
||
self.running = True
|
||
auto_sync_sources = [name for name, config in self.sync_sources.items()
|
||
if config.get('auto_sync', False)]
|
||
for name in auto_sync_sources:
|
||
self.start_sync(name)
|
||
|
||
def stop(self):
|
||
"""停止同步管理器"""
|
||
self.running = False
|
||
self.stop_all_sync()
|
||
self.executor.shutdown(wait=False)
|