Files
FoodClassifier/data_management/utils/image_downloader.py
T

185 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"img_{data['id']}.jpg"
# filename = f"img_{timestamp}_{data['id']}.jpg"
# 确定保存路径
if CREATE_SUBFOLDERS:
# 按分类创建子文件夹
# subfolder = os.path.join(save_directory, data['datasource'])
subfolder = os.path.join(save_directory)
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