183 lines
6.3 KiB
Python
183 lines
6.3 KiB
Python
"""
|
|
图片下载器
|
|
支持多线程批量下载
|
|
"""
|
|
import os
|
|
import requests
|
|
from datetime import datetime
|
|
from typing import List, Dict, Callable, Optional
|
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
|
import threading
|
|
import sys
|
|
|
|
# 添加父目录到路径
|
|
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
from config import DOWNLOAD_THREADS, DOWNLOAD_TIMEOUT, CREATE_SUBFOLDERS
|
|
|
|
|
|
class ImageDownloader:
|
|
"""图片下载器"""
|
|
|
|
def __init__(self, sqlite_manager, max_workers: int = DOWNLOAD_THREADS):
|
|
"""
|
|
初始化下载器
|
|
|
|
Args:
|
|
sqlite_manager: SQLite管理器实例
|
|
max_workers: 最大并发数
|
|
"""
|
|
self.sqlite_manager = sqlite_manager
|
|
self.max_workers = max_workers
|
|
self.is_cancelled = False
|
|
self._lock = threading.Lock()
|
|
|
|
def cancel(self):
|
|
"""取消下载"""
|
|
self.is_cancelled = True
|
|
|
|
def download_images(
|
|
self,
|
|
image_data_list: List[Dict],
|
|
save_directory: str,
|
|
progress_callback: Optional[Callable[[int, int, str], None]] = None,
|
|
complete_callback: Optional[Callable[[int, int, List[str]], None]] = None
|
|
):
|
|
"""
|
|
批量下载图片
|
|
|
|
Args:
|
|
image_data_list: 图片数据列表,每项包含:
|
|
- id: 数据库ID
|
|
- goods_id: 物品ID
|
|
- goods_name: 物品名称
|
|
- image_url: 图片URL
|
|
- datasource: 数据源
|
|
- create_time: 创建时间
|
|
save_directory: 保存目录
|
|
progress_callback: 进度回调函数 (当前数, 总数, 当前文件名)
|
|
complete_callback: 完成回调函数 (成功数, 失败数, 错误列表)
|
|
"""
|
|
self.is_cancelled = False
|
|
total = len(image_data_list)
|
|
success_count = 0
|
|
failed_count = 0
|
|
error_messages = []
|
|
|
|
# 创建保存目录
|
|
os.makedirs(save_directory, exist_ok=True)
|
|
|
|
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
|
|
# 提交所有下载任务
|
|
future_to_data = {
|
|
executor.submit(
|
|
self._download_single_image,
|
|
data,
|
|
save_directory
|
|
): data for data in image_data_list
|
|
}
|
|
|
|
# 处理完成的任务
|
|
completed = 0
|
|
for future in as_completed(future_to_data):
|
|
if self.is_cancelled:
|
|
# 取消所有未完成的任务
|
|
for f in future_to_data:
|
|
f.cancel()
|
|
break
|
|
|
|
data = future_to_data[future]
|
|
completed += 1
|
|
|
|
try:
|
|
success, error_msg = future.result()
|
|
if success:
|
|
success_count += 1
|
|
else:
|
|
failed_count += 1
|
|
if error_msg:
|
|
error_messages.append(error_msg)
|
|
|
|
# 调用进度回调
|
|
if progress_callback:
|
|
progress_callback(completed, total, data['goods_name'])
|
|
|
|
except Exception as e:
|
|
failed_count += 1
|
|
error_messages.append(f"{data['goods_name']}: {str(e)}")
|
|
|
|
# 调用完成回调
|
|
if complete_callback and not self.is_cancelled:
|
|
complete_callback(success_count, failed_count, error_messages)
|
|
|
|
def _download_single_image(
|
|
self,
|
|
data: Dict,
|
|
save_directory: str
|
|
) -> tuple[bool, Optional[str]]:
|
|
"""
|
|
下载单张图片
|
|
|
|
Returns:
|
|
(是否成功, 错误信息)
|
|
"""
|
|
try:
|
|
# 检查是否已下载
|
|
if self.sqlite_manager.is_downloaded(data['datasource'], data['image_url']):
|
|
return True, None
|
|
|
|
# 构建文件名: {分类}_{时间戳}_{ID}.jpg
|
|
# 示例: 食材_20231127143022_1001.jpg
|
|
create_time = data.get('create_time', '')
|
|
if create_time:
|
|
# 移除时间字符串中的特殊字符
|
|
timestamp = create_time.replace('-', '').replace(':', '').replace(' ', '')
|
|
else:
|
|
timestamp = datetime.now().strftime('%Y%m%d%H%M%S')
|
|
|
|
filename = f"{data['datasource']}_{timestamp}_{data['id']}.jpg"
|
|
|
|
# 确定保存路径
|
|
if CREATE_SUBFOLDERS:
|
|
# 按分类创建子文件夹
|
|
subfolder = os.path.join(save_directory, data['datasource'])
|
|
os.makedirs(subfolder, exist_ok=True)
|
|
file_path = os.path.join(subfolder, filename)
|
|
else:
|
|
file_path = os.path.join(save_directory, filename)
|
|
|
|
# 下载图片
|
|
response = requests.get(
|
|
data['image_url'],
|
|
timeout=DOWNLOAD_TIMEOUT,
|
|
stream=True
|
|
)
|
|
response.raise_for_status()
|
|
|
|
# 保存图片
|
|
with open(file_path, 'wb') as f:
|
|
for chunk in response.iter_content(chunk_size=8192):
|
|
if chunk:
|
|
f.write(chunk)
|
|
|
|
# 获取文件大小
|
|
file_size = os.path.getsize(file_path)
|
|
|
|
# 记录下载历史
|
|
self.sqlite_manager.add_download_record(
|
|
datasource=data['datasource'],
|
|
goods_id=data.get('goods_id'),
|
|
goods_name=data['goods_name'],
|
|
image_url=data['image_url'],
|
|
local_path=file_path,
|
|
file_size=file_size
|
|
)
|
|
|
|
return True, None
|
|
|
|
except requests.exceptions.RequestException as e:
|
|
error_msg = f"{data['goods_name']} (ID:{data['id']}): 网络错误 - {str(e)}"
|
|
return False, error_msg
|
|
except Exception as e:
|
|
error_msg = f"{data['goods_name']} (ID:{data['id']}): {str(e)}"
|
|
return False, error_msg
|