feat(ai-studio): make original downloads nonblocking

This commit is contained in:
chengma
2026-07-11 17:28:37 +08:00
parent 4e552ea19b
commit 879a3fe5cc
5 changed files with 726 additions and 45 deletions
+364 -33
View File
@@ -20,6 +20,8 @@ from ..workers import ImageStudioResumeJobsWorker as _RealImageStudioResumeJobsW
ASSET_MIME_TYPE = "application/x-cmshopee-image-studio-asset"
ORIGINAL_DOWNLOAD_CONCURRENCY = 2
_ORIGINAL_DOWNLOAD_THREAD_REFS = {}
def ImageStudioPullImagesWorker(*args, **kwargs):
@@ -298,6 +300,10 @@ class ImageStudioTab(QWidget):
self._thumbnail_pending = set()
self._thumbnail_errors = set()
self._thumbnail_project_id = None
self._original_downloads = {}
self._original_download_queue = []
self._original_download_states = {}
self._original_downloads_closing = False
self._thumbnail_signals = _ThumbnailSignals(self)
self._thumbnail_signals.loaded.connect(self._on_thumbnail_loaded)
self._thumbnail_signals.failed.connect(self._on_thumbnail_failed)
@@ -316,15 +322,16 @@ class ImageStudioTab(QWidget):
root.addWidget(self._build_top_bar(), 0)
splitter = QSplitter(Qt.Horizontal)
splitter.setObjectName("imageStudioMainSplitter")
splitter.addWidget(self._build_project_panel())
splitter.addWidget(self._build_pool_panel())
splitter.addWidget(self._build_generation_panel())
splitter.setStretchFactor(0, 1)
splitter.setStretchFactor(1, 3)
splitter.setStretchFactor(2, 2)
root.addWidget(splitter, 1)
self.main_splitter = QSplitter(Qt.Horizontal)
self.main_splitter.setObjectName("imageStudioMainSplitter")
self.main_splitter.addWidget(self._build_project_panel())
self.main_splitter.addWidget(self._build_pool_panel())
self.main_splitter.addWidget(self._build_generation_panel())
self.main_splitter.setStretchFactor(0, 0)
self.main_splitter.setStretchFactor(1, 4)
self.main_splitter.setStretchFactor(2, 2)
self.main_splitter.setSizes([150, 640, 360])
root.addWidget(self.main_splitter, 1)
root.addWidget(self._build_final_panel(), 0)
self._apply_workspace_style()
self._update_generation_action_text()
@@ -392,7 +399,7 @@ class ImageStudioTab(QWidget):
self.project_table.setEditTriggers(QAbstractItemView.NoEditTriggers)
self.project_table.horizontalHeader().setStretchLastSection(True)
self.project_table.verticalHeader().setVisible(False)
self.project_table.setMinimumWidth(220)
self.project_table.setMinimumWidth(150)
layout.addWidget(self.project_table, 1)
return panel
@@ -412,6 +419,10 @@ class ImageStudioTab(QWidget):
original_title = QLabel("蝦皮原主图")
original_title.setObjectName("imageStudioSectionTitle")
original_header.addWidget(original_title)
self.original_download_label = QLabel("")
self.original_download_label.setObjectName("imageStudioOriginalDownloadLabel")
self.original_download_label.setMinimumWidth(156)
original_header.addWidget(self.original_download_label)
original_header.addStretch(1)
self.original_hint_label = QLabel("单击下载并加入照片池,双击查看大图")
self.original_hint_label.setObjectName("imageStudioOriginalHintLabel")
@@ -420,10 +431,10 @@ class ImageStudioTab(QWidget):
self.original_grid = ImageStudioThumbnailGrid(parent=self)
self.original_grid.setObjectName("imageStudioOriginalGrid")
self.original_grid.setIconSize(QSize(62, 62))
self.original_grid.setGridSize(QSize(78, 92))
self.original_grid.setMinimumHeight(174)
self.original_grid.setMaximumHeight(202)
self.original_grid.setIconSize(QSize(48, 48))
self.original_grid.setGridSize(QSize(58, 78))
self.original_grid.setMinimumHeight(94)
self.original_grid.setMaximumHeight(110)
self.original_grid.setContextMenuPolicy(Qt.CustomContextMenu)
layout.addWidget(self.original_grid, 0)
@@ -645,6 +656,9 @@ class ImageStudioTab(QWidget):
color: #24292f;
font-weight: 600;
}
#imageStudioOriginalDownloadLabel {
color: #0969da;
}
#imageStudioSourcePreview {
border: 1px dashed #bfc8d3;
border-radius: 6px;
@@ -751,6 +765,26 @@ class ImageStudioTab(QWidget):
if not alias or not item_id:
self._message("信息未填写完整", "请先选择店铺并输入商品ID,再拉取蝦皮主图。")
return
project = image_studio.get_project_by_account_item(
alias,
item_id,
path=self.db_path,
include_deleted=True,
)
if project is not None and self._has_original_download_activity(project.id):
self._message(
"原图下载尚未完成",
"当前商品还有蝦皮原主图正在下载或等待下载,请等待完成后再重新拉取。",
)
return
local_count = self._local_original_count(project)
if local_count and not self._confirm(
"重新拉取蝦皮主图",
f"本地已保存 {local_count} 张蝦皮原主图。是否重新拉取线上主图?\n"
"重新拉取会刷新当前线上主图列表;不会删除本地图片、生成记录或终选记录。",
):
self._status("已取消重新拉取蝦皮主图", "muted")
return
worker = ImageStudioPullImagesWorker(
alias,
item_id,
@@ -763,6 +797,135 @@ class ImageStudioTab(QWidget):
self._start_worker(worker, "AI工场拉取蝦皮主图")
self._append_log("[AI工场] 拉取蝦皮主图开始")
def _local_original_count(self, project):
if project is None:
return 0
try:
originals = image_studio.list_assets(
project.id,
kind=image_studio.ASSET_KIND_ORIGINAL,
path=self.db_path,
)
except Exception:
return 0
return sum(
1
for asset in originals
if str(asset.local_path or "").strip() and os.path.isfile(asset.local_path)
)
def _has_original_download_activity(self, project_id):
"""Return whether a project still owns queued or live original downloads."""
if project_id is None:
return False
target_project_id = int(project_id)
if any(
int(request.get("project_id") or 0) == target_project_id
for request in self._original_download_queue
):
return True
return any(
int(record.get("project_id") or 0) == target_project_id
for record in self._original_downloads.values()
)
def _original_download_count(self):
return len(self._original_downloads)
def _original_asset_order(self, asset_id):
for index, asset in enumerate(self.assets, start=1):
if int(asset.id) == int(asset_id):
return int(asset.source_order or index)
asset = image_studio.get_asset(asset_id, path=self.db_path)
if asset is not None:
return int(asset.source_order or 1)
return 1
def _set_original_download_state(self, asset_id, project_id, state, **fields):
payload = {
"project_id": int(project_id),
"state": str(state),
}
payload.update(fields)
self._original_download_states[int(asset_id)] = payload
def _clear_original_download_state(self, asset_id, *, keep_failed=False):
state = self._original_download_states.get(int(asset_id))
if keep_failed and state and state.get("state") == "failed":
return
self._original_download_states.pop(int(asset_id), None)
def _original_card_state(self, asset):
asset_id = int(asset.id)
state = self._original_download_states.get(asset_id, {})
state_name = state.get("state")
if state_name == "queued":
return "等待下载"
if state_name == "downloading":
return "下载中"
if state_name == "retrying":
return f"重试 {state.get('retry', 1)}/{state.get('max_retries', 2)}"
if state_name == "failed":
return "下载失败"
return self._original_thumbnail_state(asset)
def _refresh_original_download_label(self):
project = self.current_project
if project is None:
self.original_download_label.setText("")
return
project_id = int(project.id)
states = [
(asset_id, state)
for asset_id, state in self._original_download_states.items()
if int(state.get("project_id") or 0) == project_id
]
retrying = [item for item in states if item[1].get("state") == "retrying"]
active = [
item
for item in states
if item[1].get("state") in {"downloading", "retrying"}
]
queued = [item for item in states if item[1].get("state") == "queued"]
failed = [item for item in states if item[1].get("state") == "failed"]
if retrying:
asset_id, state = retrying[0]
self.original_download_label.setText(
"蝦皮原主图 #%s 下载失败,正在重试 %s/%s"
% (
self._original_asset_order(asset_id),
state.get("retry", 1),
state.get("max_retries", 2),
)
)
return
if active:
orders = "、".join(
"#%s" % self._original_asset_order(asset_id)
for asset_id, _state in active
)
text = (
f"正在下载蝦皮原主图 {orders}({len(active)}/{ORIGINAL_DOWNLOAD_CONCURRENCY})"
)
if failed:
failed_asset_id, _state = failed[0]
text += (
f";主图 #{self._original_asset_order(failed_asset_id)} 下载失败,可再次点击图片重试"
)
self.original_download_label.setText(text)
return
if queued:
asset_id, _state = queued[0]
self.original_download_label.setText(
f"蝦皮原主图 #{self._original_asset_order(asset_id)} 等待下载"
)
return
if failed:
self.original_download_label.setText("蝦皮原主图下载失败,可再次点击图片重试")
return
self.original_download_label.setText("")
def _on_pull_finished(self, summary):
if self._handle_finished_error(summary, "拉取蝦皮主图失败"):
return
@@ -796,6 +959,12 @@ class ImageStudioTab(QWidget):
if project is None:
self._message("未选择商品", "请先从商品列表选择要删除的项目。")
return
if self._has_original_download_activity(project.id):
self._message(
"暂不能删除项目",
"该商品还有蝦皮原主图正在下载或等待下载,请等待完成后再删除。",
)
return
active_jobs = [
job
for job in self._list_project_jobs(project.id)
@@ -894,18 +1063,20 @@ class ImageStudioTab(QWidget):
self.original_grid.clear()
for row, asset in enumerate(originals):
order = asset.source_order or row + 1
item = QListWidgetItem(f"主图 {order}\n{self._original_thumbnail_state(asset)}")
item = QListWidgetItem(f"主图 {order}\n{self._original_card_state(asset)}")
item.setData(Qt.UserRole, {"type": "asset", "asset_id": int(asset.id)})
item.setIcon(
_asset_icon(
asset,
"原",
size=QSize(62, 62),
size=QSize(48, 48),
cached_pixmap=self._thumbnail_pixmaps.get(int(asset.id)),
)
)
item.setSizeHint(QSize(78, 92))
item.setSizeHint(QSize(58, 78))
tooltip = "单击下载并加入照片池,双击查看大图。缩略图仅用于预览,不保存原图。"
if asset.status == image_studio.ASSET_STATUS_MISSING and str(asset.local_path or "").strip():
tooltip += "\n该图片不在当前蝦皮主图列表中,本地历史文件仍已保留。"
if int(asset.id) in self._thumbnail_errors:
tooltip += "\n缩略图加载失败,可右键重新加载,或直接单击下载原图。"
item.setToolTip(tooltip)
@@ -915,7 +1086,8 @@ class ImageStudioTab(QWidget):
item = self.original_grid.item(row)
if item is not None:
order = asset.source_order or row + 1
item.setText(f"主图 {order}\n{self._original_thumbnail_state(asset)}")
item.setText(f"主图 {order}\n{self._original_card_state(asset)}")
self._refresh_original_download_label()
def _fill_pool_grid(self):
rows = []
@@ -1067,23 +1239,142 @@ class ImageStudioTab(QWidget):
if open_after:
self._open_preview(asset)
return
worker = ImageStudioDownloadOriginalWorker(
asset.id,
db_path=self.db_path,
config=self.config,
open_after=open_after,
)
worker.log.connect(self._append_log)
worker.finished.connect(self._on_download_finished)
worker.failed.connect(self._on_worker_failed)
self._start_worker(worker, "AI工场下载原图")
def _on_download_finished(self, summary):
if self._handle_finished_error(summary, "下载原图失败"):
self._enqueue_original_download(asset, open_after=open_after)
def _enqueue_original_download(self, asset, *, open_after=False):
"""Queue one full original download without locking the workspace."""
asset_id = int(asset.id)
project_id = int(asset.project_id)
active = self._original_downloads.get(asset_id)
if active is not None:
if open_after:
active["open_after"] = True
worker = active.get("worker")
if worker is not None:
worker.open_after = True
self._status(
f"蝦皮原主图 #{self._original_asset_order(asset_id)} 正在下载,请稍候",
"info",
)
return
for request in self._original_download_queue:
if int(request.get("asset_id") or 0) != asset_id:
continue
request["open_after"] = bool(request.get("open_after") or open_after)
self._status(
f"蝦皮原主图 #{self._original_asset_order(asset_id)} 已在下载队列中",
"info",
)
return
self._set_original_download_state(asset_id, project_id, "queued")
self._original_download_queue.append(
{
"asset_id": asset_id,
"project_id": project_id,
"open_after": bool(open_after),
}
)
self._refresh_original_download_ui(project_id)
self._start_queued_original_downloads()
def _start_queued_original_downloads(self):
while (
self._original_download_queue
and self._original_download_count() < ORIGINAL_DOWNLOAD_CONCURRENCY
):
request = self._original_download_queue.pop(0)
asset_id = int(request["asset_id"])
project_id = int(request["project_id"])
try:
asset = image_studio.get_asset(asset_id, path=self.db_path)
except Exception:
asset = None
if asset is None or int(asset.project_id) != project_id:
self._clear_original_download_state(asset_id)
continue
worker = ImageStudioDownloadOriginalWorker(
asset_id,
db_path=self.db_path,
config=self.config,
open_after=bool(request.get("open_after")),
max_retries=2,
)
worker.log.connect(self._append_log)
worker.progress.connect(
lambda payload, aid=asset_id, pid=project_id: self._on_original_download_progress(
aid, pid, payload
)
)
worker.finished.connect(
lambda summary, aid=asset_id, pid=project_id: self._on_original_download_finished(
aid, pid, summary
)
)
worker.cancelled.connect(
lambda summary, aid=asset_id, pid=project_id: self._on_original_download_cancelled(
aid, pid, summary
)
)
thread = run_worker(worker, thread_name="AI工场下载蝦皮原主图", start=False)
thread.finished.connect(
lambda aid=asset_id, active_thread=thread: self._forget_original_download(
aid, active_thread
)
)
_ORIGINAL_DOWNLOAD_THREAD_REFS[id(thread)] = (thread, worker)
self._original_downloads[asset_id] = {
"worker": worker,
"thread": thread,
"project_id": project_id,
"open_after": bool(request.get("open_after")),
}
self._set_original_download_state(asset_id, project_id, "downloading")
self._refresh_original_download_ui(project_id)
thread.start()
def _on_original_download_progress(self, asset_id, project_id, payload):
if self._original_downloads_closing:
return
if not isinstance(payload, dict):
return
state = str(payload.get("state") or "")
if state == "retry":
self._set_original_download_state(
asset_id,
project_id,
"retrying",
retry=int(payload.get("retry") or 1),
max_retries=int(payload.get("max_retries") or 2),
)
elif state in {"start", "success"}:
self._set_original_download_state(asset_id, project_id, "downloading")
elif state == "failed":
self._set_original_download_state(asset_id, project_id, "failed")
self._refresh_original_download_ui(project_id)
def _on_original_download_finished(self, asset_id, project_id, summary):
if self._original_downloads_closing:
self._clear_original_download_state(asset_id)
return
summary = dict(summary or {})
if summary.get("ok") is False:
self._set_original_download_state(asset_id, project_id, "failed")
self._refresh_original_download_ui(project_id)
if self.current_project is not None and int(self.current_project.id) == project_id:
self._status(
f"蝦皮原主图 #{self._original_asset_order(asset_id)} 下载失败,可再次点击图片重试",
"warning",
)
return
self._clear_original_download_state(asset_id)
if self.current_project is None or int(self.current_project.id) != project_id:
return
asset = summary.get("asset")
self._finish_worker()
self.refresh_project_assets()
asset = summary.get("asset")
if asset is not None:
refreshed = self._asset_by_id(asset.id) or asset
self._select_source_asset(refreshed)
@@ -1091,6 +1382,39 @@ class ImageStudioTab(QWidget):
self._open_preview(refreshed)
self._status("原图已加入照片池", "success")
def _on_original_download_cancelled(self, asset_id, project_id, summary):
self._clear_original_download_state(asset_id)
if self._original_downloads_closing:
return
self._refresh_original_download_ui(project_id)
def _forget_original_download(self, asset_id, thread):
_ORIGINAL_DOWNLOAD_THREAD_REFS.pop(id(thread), None)
record = self._original_downloads.get(int(asset_id))
if record is None or record.get("thread") is not thread:
return
project_id = int(record.get("project_id") or 0)
self._original_downloads.pop(int(asset_id), None)
if self._original_downloads_closing:
return
self._refresh_original_download_ui(project_id)
self._start_queued_original_downloads()
def _refresh_original_download_ui(self, project_id):
if self.current_project is None or int(self.current_project.id) != int(project_id):
return
self._fill_original_grid()
self._update_project_action_buttons()
def _cancel_original_downloads(self):
for request in self._original_download_queue:
self._clear_original_download_state(request.get("asset_id"))
self._original_download_queue = []
for record in list(self._original_downloads.values()):
worker = record.get("worker")
if worker is not None:
worker.cancel()
def _select_source_asset(self, asset):
self.selected_source_asset_id = int(asset.id)
self._refresh_source_label()
@@ -1145,7 +1469,12 @@ class ImageStudioTab(QWidget):
def _update_project_action_buttons(self):
enabled = self.current_project is not None and not self._operation_running
self.open_folder_button.setEnabled(enabled)
self.delete_project_button.setEnabled(enabled)
self.delete_project_button.setEnabled(
enabled
and not self._has_original_download_activity(
getattr(self.current_project, "id", None)
)
)
def _thumbnail_key(self, project_id, asset_id):
return f"{int(project_id)}:{int(asset_id)}"
@@ -1746,6 +2075,8 @@ class ImageStudioTab(QWidget):
_emit_status(self.status_callback, message, level=level)
def closeEvent(self, event):
self._original_downloads_closing = True
self._cancel_original_downloads()
try:
self._thumbnail_loader.close()
except Exception:
+91 -9
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import datetime as _dt
import re
import threading
import time
from .. import image_studio, image_studio_export, image_studio_generation, image_studio_images
from .widgets import *
@@ -110,22 +111,103 @@ class ImageStudioPullImagesWorker(BaseWorker):
class ImageStudioDownloadOriginalWorker(BaseWorker):
"""Download one remote original image into the project pool."""
def __init__(self, asset_id, *, db_path=None, config=None, open_after=False):
def __init__(
self,
asset_id,
*,
db_path=None,
config=None,
open_after=False,
max_retries=2,
retry_delays=(1, 2),
):
super().__init__()
self.asset_id = int(asset_id)
self.db_path = db_path
self.config = config
self.open_after = bool(open_after)
self.max_retries = max(0, int(max_retries or 0))
self.retry_delays = tuple(float(delay) for delay in (retry_delays or ()))
def execute(self):
self.log.emit("[AI工场] 下载原图:开始")
asset = image_studio_images.download_original_asset(
self.asset_id,
path=self.db_path,
config=self.config,
)
self.log.emit("[AI工场] 下载原图:成功")
return {"asset": asset, "open_after": self.open_after}
attempts = self.max_retries + 1
for attempt in range(1, attempts + 1):
if self.should_cancel():
return {"asset_id": self.asset_id, "cancelled": True}
self.progress.emit(
{
"asset_id": self.asset_id,
"state": "start",
"attempt": attempt,
"attempts": attempts,
}
)
self.log.emit(f"[AI工场] 下载蝦皮原主图 #{self.asset_id}:开始")
try:
asset = image_studio_images.download_original_asset(
self.asset_id,
path=self.db_path,
config=self.config,
)
except Exception:
retry = attempt
if attempt >= attempts:
self.progress.emit(
{
"asset_id": self.asset_id,
"state": "failed",
"attempt": attempt,
"attempts": attempts,
}
)
self.log.emit(f"[AI工场] 下载蝦皮原主图 #{self.asset_id}:最终失败")
return {
"ok": False,
"asset_id": self.asset_id,
"open_after": self.open_after,
"error": "蝦皮原主图下载失败,请稍后再次点击图片重试。",
}
delay = self._retry_delay(retry)
self.progress.emit(
{
"asset_id": self.asset_id,
"state": "retry",
"retry": retry,
"max_retries": self.max_retries,
"delay_seconds": delay,
}
)
self.log.emit(
f"[AI工场] 下载蝦皮原主图 #{self.asset_id}:失败,准备重试 {retry}/{self.max_retries}"
)
if not self._wait_for_retry(delay):
return {"asset_id": self.asset_id, "cancelled": True}
continue
self.progress.emit(
{
"asset_id": self.asset_id,
"state": "success",
"attempt": attempt,
"attempts": attempts,
}
)
self.log.emit(f"[AI工场] 下载蝦皮原主图 #{self.asset_id}:成功")
return {"asset": asset, "asset_id": self.asset_id, "open_after": self.open_after}
return {"asset_id": self.asset_id, "cancelled": True}
def _retry_delay(self, retry):
if retry <= 0:
return 0.0
index = min(retry - 1, len(self.retry_delays) - 1)
return self.retry_delays[index] if index >= 0 else 0.0
def _wait_for_retry(self, delay_seconds):
deadline = time.monotonic() + max(0.0, float(delay_seconds or 0))
while time.monotonic() < deadline:
if self.should_cancel():
return False
time.sleep(min(0.1, max(0.0, deadline - time.monotonic())))
return not self.should_cancel()
class ImageStudioGenerateJobsWorker(BaseWorker):
+5 -2
View File
@@ -3,7 +3,7 @@ id: T-608
title: AI工场主图重新拉取确认、非阻塞原图下载与区域比例优化
phase: 7
deps: [T-607]
status: TODO
status: DONE
created: 2026-07-11
---
@@ -120,4 +120,7 @@ git diff --check
## 执行记录
- 待执行。
- 2026-07-11:在 `ImageStudioTab` 增加本地原图计数与重新拉取确认;同一项目存在原图下载/排队时,拦截重新拉取和软删除。
- 2026-07-11:将单张原图完整下载从通用全局 worker 拆为专用内存队列;同图去重、最多并发 2、失败自动重试 2 次(最多 3 次尝试),仅刷新原图卡片和原图区状态,不再禁用整个 AI工场。
- 2026-07-11:增加下载中/等待/重试/失败状态、项目切换隔离、关闭时取消未开始队列和已启动 worker 的取消请求;压缩原图区与商品列表尺寸,将空间让给照片池。
- 2026-07-11:新增/更新 GUI 与 worker 回归测试。干净 worktree 验证通过:`python -m ruff check app tests main.py`、`py -3.10 -m compileall app main.py`、`py -3.10 -m unittest discover -s tests`(386 项)、`git diff --check`。主工作区另有未提交默认提示词文件变更,导致其全量测试中 3 项默认提示词断言失败,未纳入本任务修改。
+213 -1
View File
@@ -693,7 +693,10 @@ class GuiTests(TempDirMixin, unittest.TestCase):
)
self.assertEqual(9, len(thumbnail_loader.submissions))
self.assertEqual(0, tab.pool_grid.count())
self.assertEqual(QSize(78, 92), tab.original_grid.gridSize())
self.assertEqual(QSize(58, 78), tab.original_grid.gridSize())
self.assertEqual(94, tab.original_grid.minimumHeight())
self.assertEqual(110, tab.original_grid.maximumHeight())
self.assertEqual(150, tab.project_table.minimumWidth())
self.assertEqual(QSize(108, 124), tab.pool_grid.gridSize())
self.assertGreaterEqual(tab.prompt_edit.minimumHeight(), 210)
@@ -1029,6 +1032,215 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_image_studio_repull_confirms_when_local_originals_exist(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
project = image_studio.create_or_get_project(
account_alias="alias-a",
account_slug="alias_a",
item_id="51100639510",
path=cfg["db_path"],
)
original = image_studio.sync_original_asset_urls(
project.id,
[{"index": 1, "src": "https://susercontent.com/main-1.jpg"}],
path=cfg["db_path"],
)[0]
image_studio.update_asset_local_path(
original.id,
self.write_test_image(os.path.join(temp_dir, "main-1.jpg")),
path=cfg["db_path"],
)
tab = ImageStudioTab(config=cfg, db_path=cfg["db_path"])
self.addCleanup(tab.close)
tab._select_project(project.id)
confirmations = []
tab._confirm = lambda title, text: confirmations.append((title, text)) or False
with mock.patch("app.gui.tabs.image_studio.ImageStudioPullImagesWorker") as worker_factory:
tab.pull_main_images()
worker_factory.assert_not_called()
self.assertEqual("重新拉取蝦皮主图", confirmations[0][0])
self.assertIn("本地已保存 1 张蝦皮原主图", confirmations[0][1])
self.assertIn("不会删除本地图片、生成记录或终选记录", confirmations[0][1])
self.assert_removed(temp_dir)
def test_image_studio_original_download_queue_is_local_and_nonblocking(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
db.init_db(cfg["db_path"])
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
project = image_studio.create_or_get_project(
account_alias="alias-a",
account_slug="alias_a",
item_id="51100639510",
path=cfg["db_path"],
)
originals = image_studio.sync_original_asset_urls(
project.id,
[
{"index": index, "src": f"https://susercontent.com/main-{index}.jpg"}
for index in range(1, 4)
],
path=cfg["db_path"],
)
tab = ImageStudioTab(
config=cfg,
db_path=cfg["db_path"],
thumbnail_loader=FakeThumbnailLoader(),
)
self.addCleanup(tab.close)
tab._select_project(project.id)
messages = []
tab._message = lambda title, text: messages.append((title, text))
class FakeDownloadWorker:
instances = []
def __init__(self, asset_id, **kwargs):
self.asset_id = asset_id
self.open_after = kwargs.get("open_after", False)
self.progress = DummySignal()
self.log = DummySignal()
self.finished = DummySignal()
self.cancelled = DummySignal()
self.cancel_requested = False
FakeDownloadWorker.instances.append(self)
def cancel(self):
self.cancel_requested = True
threads = []
def fake_run_worker(*args, **kwargs):
thread = FakeThread()
threads.append(thread)
return thread
with mock.patch(
"app.gui.tabs.image_studio.ImageStudioDownloadOriginalWorker",
FakeDownloadWorker,
), mock.patch("app.gui.tabs.image_studio.run_worker", side_effect=fake_run_worker):
for original in originals:
tab._ensure_original_in_pool(original)
self.assertEqual(2, len(FakeDownloadWorker.instances))
self.assertEqual(1, len(tab._original_download_queue))
self.assertFalse(tab._operation_running)
self.assertTrue(tab.prompt_edit.isEnabled())
self.assertTrue(tab.pool_grid.isEnabled())
self.assertFalse(tab.delete_project_button.isEnabled())
self.assertIn("等待下载", tab.original_grid.item(2).text())
self.assertIn("正在下载蝦皮原主图 #1、#2(2/2)", tab.original_download_label.text())
tab.delete_current_project()
tab.pull_main_images()
self.assertEqual("暂不能删除项目", messages[0][0])
self.assertEqual("原图下载尚未完成", messages[1][0])
first_worker = FakeDownloadWorker.instances[0]
first_worker.progress.emit(
{"state": "retry", "retry": 1, "max_retries": 2}
)
self.assertIn("正在重试 1/2", tab.original_download_label.text())
self.assertIn("重试 1/2", tab.original_grid.item(0).text())
local_path = self.write_test_image(os.path.join(temp_dir, "downloaded.jpg"))
downloaded = image_studio.update_asset_local_path(
originals[0].id,
local_path,
path=cfg["db_path"],
)
first_worker.finished.emit({"asset": downloaded, "open_after": False})
self.assertEqual(originals[0].id, tab.selected_source_asset_id)
self.assertEqual(1, tab.pool_grid.count())
threads[0].finished.emit()
self.assertEqual(3, len(FakeDownloadWorker.instances))
self.assertEqual(0, len(tab._original_download_queue))
self.assertEqual(2, len(tab._original_downloads))
third_worker = FakeDownloadWorker.instances[2]
third_worker.finished.emit({"ok": False, "error": "英文网络错误"})
self.assertIn("下载失败", tab.original_grid.item(2).text())
self.assertIn("可再次点击图片重试", tab.original_download_label.text())
self.assertFalse(any("英文网络错误" in text for _title, text in messages))
tab.close()
self.assertTrue(FakeDownloadWorker.instances[1].cancel_requested)
self.assertTrue(FakeDownloadWorker.instances[2].cancel_requested)
for thread in threads[1:]:
thread.finished.emit()
self.assert_removed(temp_dir)
def test_image_studio_original_download_result_does_not_refresh_other_project(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
db.init_db(cfg["db_path"])
first = image_studio.create_or_get_project(
account_alias="alias-a",
account_slug="alias_a",
item_id="51100639510",
path=cfg["db_path"],
)
second = image_studio.create_or_get_project(
account_alias="alias-a",
account_slug="alias_a",
item_id="51100639511",
path=cfg["db_path"],
)
original = image_studio.sync_original_asset_urls(
first.id,
[{"index": 1, "src": "https://susercontent.com/main-1.jpg"}],
path=cfg["db_path"],
)[0]
tab = ImageStudioTab(
config=cfg,
db_path=cfg["db_path"],
thumbnail_loader=FakeThumbnailLoader(),
)
self.addCleanup(tab.close)
tab._select_project(first.id)
class FakeDownloadWorker:
def __init__(self, asset_id, **kwargs):
self.asset_id = asset_id
self.open_after = kwargs.get("open_after", False)
self.progress = DummySignal()
self.log = DummySignal()
self.finished = DummySignal()
self.cancelled = DummySignal()
def cancel(self):
return None
thread = FakeThread()
with mock.patch(
"app.gui.tabs.image_studio.ImageStudioDownloadOriginalWorker",
FakeDownloadWorker,
), mock.patch("app.gui.tabs.image_studio.run_worker", return_value=thread):
tab._ensure_original_in_pool(original)
tab._select_project(second.id)
downloaded = image_studio.update_asset_local_path(
original.id,
self.write_test_image(os.path.join(temp_dir, "downloaded.jpg")),
path=cfg["db_path"],
)
next(iter(tab._original_downloads.values()))["worker"].finished.emit(
{"asset": downloaded, "open_after": False}
)
self.assertEqual(second.id, tab.current_project.id)
self.assertEqual(0, tab.pool_grid.count())
self.assertEqual("", tab.original_download_label.text())
thread.finished.emit()
self.assert_removed(temp_dir)
def test_image_studio_event_log_hides_provider_urls(self):
message = gui_workers._format_image_studio_event(
{
+53
View File
@@ -1,5 +1,7 @@
import os
import unittest
from types import SimpleNamespace
from unittest import mock
os.environ.setdefault("QT_QPA_PLATFORM", "offscreen")
@@ -14,6 +16,7 @@ from PySide6.QtCore import QEventLoop, QTimer
from PySide6.QtWidgets import QApplication
from app.workers import BaseWorker, run_worker
from app.gui.workers import ImageStudioDownloadOriginalWorker
class DemoWorker(BaseWorker):
@@ -112,6 +115,56 @@ class WorkerTests(unittest.TestCase):
with self.assertRaises(TypeError):
run_worker(object(), start=False)
def test_image_studio_original_download_retries_twice_without_raw_error(self):
worker = ImageStudioDownloadOriginalWorker(
12,
max_retries=2,
retry_delays=(0, 0),
)
progress = []
logs = []
worker.progress.connect(lambda payload: progress.append(dict(payload)))
worker.log.connect(logs.append)
successful_asset = SimpleNamespace(id=12)
with mock.patch(
"app.gui.workers.image_studio_images.download_original_asset",
side_effect=[RuntimeError("first"), RuntimeError("second"), successful_asset],
) as download:
summary = worker.execute()
self.assertEqual(3, download.call_count)
self.assertIs(successful_asset, summary["asset"])
retries = [payload for payload in progress if payload.get("state") == "retry"]
self.assertEqual(
[
{"asset_id": 12, "state": "retry", "retry": 1, "max_retries": 2, "delay_seconds": 0.0},
{"asset_id": 12, "state": "retry", "retry": 2, "max_retries": 2, "delay_seconds": 0.0},
],
retries,
)
self.assertTrue(any("重试 1/2" in message for message in logs))
self.assertTrue(any("重试 2/2" in message for message in logs))
self.assertFalse(any("first" in message or "second" in message for message in logs))
def test_image_studio_original_download_returns_chinese_final_failure(self):
worker = ImageStudioDownloadOriginalWorker(
12,
max_retries=2,
retry_delays=(0, 0),
)
with mock.patch(
"app.gui.workers.image_studio_images.download_original_asset",
side_effect=RuntimeError("https://example.invalid/private"),
) as download:
summary = worker.execute()
self.assertEqual(3, download.call_count)
self.assertFalse(summary["ok"])
self.assertEqual("蝦皮原主图下载失败,请稍后再次点击图片重试。", summary["error"])
self.assertNotIn("https://", summary["error"])
if __name__ == "__main__":
unittest.main()