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
+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()