Files
cmshoppe/tests/test_workers.py
T

288 lines
9.9 KiB
Python
Raw Normal View History

2026-06-27 10:05:29 +08:00
import os
import tempfile
2026-06-27 10:05:29 +08:00
import unittest
from types import SimpleNamespace
from unittest import mock
2026-06-27 10:05:29 +08:00
os.environ.setdefault("QT_QPA_PLATFORM", "offscreen")
from _helpers import REPO_ROOT # noqa: F401
from app import db, image_studio, image_studio_images, workers
2026-06-27 10:05:29 +08:00
if workers.QT_IMPORT_ERROR is not None:
raise unittest.SkipTest("PySide6 未安装")
from PySide6.QtCore import QEventLoop, QTimer
from PySide6.QtWidgets import QApplication
from app.workers import BaseWorker, run_worker
from app.gui.workers import (
ImageStudioDownloadOriginalWorker,
ImageStudioPullImagesWorker,
ProductSuiteAiWriteWorker,
ProductSuiteGenerateWorker,
)
2026-06-27 10:05:29 +08:00
class DemoWorker(BaseWorker):
def execute(self):
self.log.emit("开始")
self.progress.emit({"done": 1, "total": 1})
self.row_updated.emit(7, {"status": "success"})
return {"ok": True, "done": 1}
class CancelAwareWorker(BaseWorker):
def execute(self):
if self.should_cancel():
return {"done": 0}
self.cancel()
return {"done": 0}
class FailingWorker(BaseWorker):
def execute(self):
raise RuntimeError("模拟失败")
class WorkerTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.app = QApplication.instance() or QApplication([])
def start_and_wait(self, thread, timeout_ms=2000):
loop = QEventLoop()
finished = []
def on_finished():
finished.append(True)
loop.quit()
timer = QTimer()
timer.setSingleShot(True)
timer.timeout.connect(loop.quit)
thread.finished.connect(on_finished)
timer.start(timeout_ms)
thread.start()
loop.exec()
self.assertTrue(finished, "worker thread did not finish before timeout")
def test_worker_emits_common_signals_and_finishes(self):
worker = DemoWorker()
logs = []
progress = []
rows = []
finished = []
worker.log.connect(logs.append)
worker.progress.connect(lambda payload: progress.append(dict(payload)))
worker.row_updated.connect(lambda task_id, fields: rows.append((task_id, dict(fields))))
worker.finished.connect(lambda payload: finished.append(dict(payload)))
thread = run_worker(worker, start=False)
self.start_and_wait(thread)
self.assertEqual(["开始"], logs)
self.assertEqual([{"done": 1, "total": 1}], progress)
self.assertEqual([(7, {"status": "success"})], rows)
self.assertEqual([{"ok": True, "done": 1}], finished)
def test_cancel_flag_emits_cancelled_instead_of_finished(self):
worker = CancelAwareWorker()
cancelled = []
finished = []
worker.cancelled.connect(lambda payload: cancelled.append(dict(payload)))
worker.finished.connect(lambda payload: finished.append(dict(payload)))
thread = run_worker(worker, start=False)
self.start_and_wait(thread)
self.assertEqual([{"done": 0, "cancelled": True}], cancelled)
self.assertEqual([], finished)
def test_uncaught_exception_emits_failed_and_finished_summary(self):
worker = FailingWorker()
failed = []
finished = []
worker.failed.connect(lambda task_id, error: failed.append((task_id, error)))
worker.finished.connect(lambda payload: finished.append(dict(payload)))
thread = run_worker(worker, start=False)
self.start_and_wait(thread)
self.assertEqual([(-1, "模拟失败")], failed)
self.assertEqual([{"ok": False, "error": "模拟失败"}], finished)
def test_run_worker_rejects_plain_object(self):
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"])
def test_image_studio_pull_worker_returns_token_when_cancelled(self):
worker = ImageStudioPullImagesWorker(
"alias-a",
"51100639510",
pull_run_token="pull-token",
)
worker.cancel()
with mock.patch(
"app.gui.workers.image_studio.pull_remote_main_image_urls",
) as pull:
summary = worker.execute()
pull.assert_not_called()
self.assertTrue(summary["cancelled"])
self.assertEqual("pull-token", summary["pull_run_token"])
def test_image_studio_original_download_maps_safe_cancel(self):
worker = ImageStudioDownloadOriginalWorker(12)
with mock.patch(
"app.gui.workers.image_studio_images.download_original_asset",
side_effect=image_studio_images.ImageStudioImageCancelled("停止"),
):
summary = worker.execute()
self.assertEqual({"asset_id": 12, "cancelled": True}, summary)
def test_product_suite_ai_write_worker_uses_image_analysis_not_title_generation(self):
worker = ProductSuiteAiWriteWorker(
"补充要求",
"输出语言:繁体中文",
image_paths=["first.jpg", "second.jpg"],
config={"ai": {"backend": "cmhub"}},
cmhub_config_path="cmhub.json",
)
expected = {
"text": "根据图片整理的商品卖点",
"image_count": 2,
"metadata": {"points_cost": 1, "points_balance": 231},
}
with mock.patch(
"app.gui.workers.ai.analyze_product_images",
return_value=expected,
) as analyze, mock.patch("app.gui.workers.ai.gen_title") as gen_title:
result = worker.execute()
analyze.assert_called_once_with(
"补充要求",
"输出语言:繁体中文",
["first.jpg", "second.jpg"],
config={"ai": {"backend": "cmhub"}},
cmhub_config_path="cmhub.json",
)
gen_title.assert_not_called()
self.assertEqual(expected, result)
def test_product_suite_worker_writes_generation_round_and_stable_slots(self):
with tempfile.TemporaryDirectory() as temp_dir:
db_path = os.path.join(temp_dir, "cmshopee.db")
db.init_db(db_path)
project = image_studio.create_or_get_project(
account_alias="alias-a",
account_slug="alias-a",
item_id="51100639510",
path=db_path,
)
source = image_studio.add_asset(
project.id,
image_studio.ASSET_KIND_ORIGINAL,
path=db_path,
)
worker = ProductSuiteGenerateWorker(
project.id,
[
{
"source_asset_id": source.id,
"job_type": "白底图",
"prompt": "第一张",
"generation_round_key": "worker-round",
"generation_slot_index": 0,
},
{
"source_asset_id": source.id,
"job_type": "场景图",
"prompt": "第二张",
"generation_round_key": "worker-round",
"generation_slot_index": 1,
},
],
generation_round_key="worker-round",
db_path=db_path,
)
with mock.patch(
"app.gui.workers.image_studio_generation.run_jobs",
return_value={"total": 2, "success": 0, "failed": 0, "cancelled": 0},
):
summary = worker.execute()
jobs = [image_studio.get_job(job_id, path=db_path) for job_id in worker.job_ids]
self.assertEqual("worker-round", summary["generation_round_key"])
self.assertEqual(
[("worker-round", 0), ("worker-round", 1)],
[
(job.generation_round_key, job.generation_slot_index)
for job in jobs
],
)
2026-06-27 10:05:29 +08:00
if __name__ == "__main__":
unittest.main()