feat(product-suite): add cmhub vision AI writing

This commit is contained in:
chengma
2026-07-17 09:09:09 +08:00
parent 56d6b59a00
commit 40a9f5fbf2
15 changed files with 575 additions and 40 deletions
+139
View File
@@ -97,6 +97,7 @@ class AITests(TempDirMixin, unittest.TestCase):
"base_url": "https://cmhub.example.com",
"title_alias": "title-standard",
"image_alias": "image-hd",
"vision_alias": "vision-standard",
"connect_timeout": 3,
"download_with_curl": "false",
"check_balance_before_batch": False,
@@ -454,6 +455,144 @@ class AITests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_cmhub_analyze_product_images_uses_vision_alias_and_safe_metadata(self):
with self.make_temp_dir() as temp_dir:
cfg, key_path = self._cmhub_config(temp_dir)
first = os.path.join(temp_dir, "first.png")
second = os.path.join(temp_dir, "second.jpg")
with open(first, "wb") as fh:
fh.write(b"first-image")
with open(second, "wb") as fh:
fh.write(b"second-image")
calls = []
def fake_request(method, url, **kwargs):
calls.append((method, url, kwargs))
return _RequestsResponse(
{
"text": "浅绿色连帽上衣,突出宽松版型与日常穿搭场景。",
"alias": "vision-standard",
"model_used": "vision-provider",
"points_cost": 1,
"points_balance": 231,
"call_id": "vision-call-1",
}
)
with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request):
result = ai.analyze_product_images(
"补充卖点:避免夸大",
"输出语言:繁体中文",
[first, second],
config=cfg,
cmhub_config_path=key_path,
)
self.assertEqual("浅绿色连帽上衣,突出宽松版型与日常穿搭场景。", result["text"])
self.assertEqual(2, result["image_count"])
self.assertEqual(1, result["metadata"]["points_cost"])
self.assertEqual(231, result["metadata"]["points_balance"])
self.assertNotIn("image_base64", result)
self.assertEqual("POST", calls[0][0])
self.assertEqual(
"https://cmhub.example.com/api/v1/analyze/images",
calls[0][1],
)
payload = calls[0][2]["json"]
self.assertEqual("vision-standard", payload["model"])
self.assertEqual(2, len(payload["images"]))
self.assertTrue(payload["images"][0]["image_base64"].startswith("data:image/png;base64,"))
self.assertTrue(payload["images"][1]["image_base64"].startswith("data:image/jpeg;base64,"))
self.assertEqual({"temperature": 0.2}, payload["parameters"])
self.assertEqual((3, ai.CMHUB_VISION_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"])
self.assert_removed(temp_dir)
def test_cmhub_analyze_product_images_requires_vision_alias_and_local_limits(self):
with self.make_temp_dir() as temp_dir:
cfg, key_path = self._cmhub_config(temp_dir)
source = os.path.join(temp_dir, "source.jpg")
with open(source, "wb") as fh:
fh.write(b"source")
cfg["ai"]["cmhub"]["vision_alias"] = ""
with self.assertRaises(ai.CMHubError) as raised:
ai.analyze_product_images(
"提示",
"输出语言:繁体中文",
[source],
config=cfg,
cmhub_config_path=key_path,
)
self.assertEqual("cmhub_not_configured", raised.exception.code)
self.assertIn("图片理解别名", str(raised.exception))
cfg["ai"]["cmhub"]["vision_alias"] = "vision-standard"
paths = []
for index in range(ai.CMHUB_VISION_MAX_IMAGES + 1):
path = os.path.join(temp_dir, "source-%d.jpg" % index)
with open(path, "wb") as fh:
fh.write(b"image")
paths.append(path)
with mock.patch.object(ai._cmhub_session(), "request") as request:
with self.assertRaises(ai.AIError) as too_many:
ai.analyze_product_images(
"提示",
"输出语言:繁体中文",
paths,
config=cfg,
cmhub_config_path=key_path,
)
self.assertIn("最多支持", str(too_many.exception))
request.assert_not_called()
oversized = os.path.join(temp_dir, "oversized.jpg")
with open(oversized, "wb") as fh:
fh.truncate(ai.CMHUB_VISION_MAX_IMAGE_BYTES + 1)
with mock.patch.object(ai._cmhub_session(), "request") as request:
with self.assertRaises(ai.AIError) as too_large:
ai.analyze_product_images(
"提示",
"输出语言:繁体中文",
[oversized],
config=cfg,
cmhub_config_path=key_path,
)
self.assertIn("超过10MiB", str(too_large.exception))
request.assert_not_called()
self.assert_removed(temp_dir)
def test_cmhub_analyze_product_images_read_timeout_is_not_retried(self):
with self.make_temp_dir() as temp_dir:
cfg, key_path = self._cmhub_config(temp_dir)
source = os.path.join(temp_dir, "source.jpg")
with open(source, "wb") as fh:
fh.write(b"source")
calls = []
def fake_request(method, url, **kwargs):
calls.append((method, url, kwargs))
raise ai.requests.exceptions.ReadTimeout("slow")
with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request):
with self.assertRaises(ai.CMHubError) as raised:
ai.analyze_product_images(
"提示",
"输出语言:繁体中文",
[source],
config=cfg,
cmhub_config_path=key_path,
)
self.assertEqual("read_timeout", raised.exception.code)
self.assertIn("结果未确认", str(raised.exception))
self.assertEqual(1, len(calls))
self.assertEqual((3, ai.CMHUB_VISION_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"])
self.assert_removed(temp_dir)
def test_cmhub_gen_cover_downloads_image_url_safely(self):
try:
from PIL import Image
+29
View File
@@ -28,6 +28,7 @@ class AppConfigTests(TempDirMixin, unittest.TestCase):
self.assertEqual(os.path.join(temp_dir, "config", "cmhub.json"), config["cmhub_config_path"])
self.assertEqual(240, appconfig.response_timeout(config))
self.assertFalse(appconfig.ai_config(config)["generate_cover"])
self.assertEqual("vision-standard", appconfig.cmhub_config(config)["vision_alias"])
self.assertEqual("title", appconfig.ai_generate_mode(config))
self.assertEqual("title", appconfig.shopee_update_config(config)["update_mode"])
self.assertNotIn("allow_cover_update", appconfig.shopee_update_config(config))
@@ -153,6 +154,34 @@ class AppConfigTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_cmhub_vision_alias_defaults_for_legacy_config_and_preserves_saved_value(self):
with self.make_temp_dir() as temp_dir:
config_path = os.path.join(temp_dir, "config.json")
with open(config_path, "w", encoding="utf-8") as fh:
json.dump(
{
"ai": {
"backend": "cmhub",
"cmhub": {
"base_url": "https://cmhub.example.com",
"title_alias": "title-standard",
"image_alias": "image-hd",
},
}
},
fh,
ensure_ascii=False,
)
legacy = appconfig.load_config(config_path)
self.assertEqual("vision-standard", appconfig.cmhub_config(legacy)["vision_alias"])
legacy["ai"]["cmhub"]["vision_alias"] = "vision-custom"
saved = appconfig.save_config(legacy, path=config_path)
self.assertEqual("vision-custom", appconfig.cmhub_config(saved)["vision_alias"])
self.assert_removed(temp_dir)
def test_shopee_update_legacy_parallel_config_is_migrated(self):
with self.make_temp_dir() as temp_dir:
config_path = os.path.join(temp_dir, "config.json")
+41
View File
@@ -2798,6 +2798,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assertEqual(QLineEdit.Password, tab.cmhub_api_key_edit.echoMode())
self.assertEqual("title-old", tab.cmhub_title_alias_combo.currentData())
self.assertEqual("image-old", tab.cmhub_image_alias_combo.currentData())
self.assertEqual("vision-standard", tab.cmhub_vision_alias_combo.currentData())
tab.cmhub_base_url_edit.setText("https://cmhub.example.com/api/v1/")
tab.cmhub_api_key_edit.setText("sk-new-secret")
@@ -2819,9 +2820,17 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"pricing_status": "priced",
"prices": [{"resolution": "1K", "points_cost": 5}],
},
{
"alias": "vision-standard",
"operation_type": "vision",
"requires_image": True,
"pricing_status": "priced",
"prices": [{"points_cost": 1}],
},
],
title_selected="title-standard",
image_selected="image-standard",
vision_selected="vision-standard",
)
with mock.patch("app.gui.QMessageBox.warning") as warning, \
@@ -2837,6 +2846,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assertEqual("https://cmhub.example.com", tab.cmhub_base_url_edit.text())
self.assertEqual("title-standard", saved["ai"]["cmhub"]["title_alias"])
self.assertEqual("image-standard", saved["ai"]["cmhub"]["image_alias"])
self.assertEqual("vision-standard", saved["ai"]["cmhub"]["vision_alias"])
self.assertEqual(12, saved["ai"]["cmhub"]["connect_timeout"])
self.assertTrue(saved["ai"]["cmhub"]["check_balance_before_batch"])
self.assertEqual(
@@ -2856,6 +2866,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"base_url": "https://cmhub.example.com",
"title_alias": "title-saved",
"image_alias": "image-saved",
"vision_alias": "vision-saved",
"connect_timeout": 10,
"check_balance_before_batch": False,
}
@@ -2889,6 +2900,27 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"pricing_status": "priced",
"prices": [{"resolution": "1K", "points_cost": 5}],
},
{
"alias": "vision-priced",
"operation_type": "vision",
"requires_image": True,
"pricing_status": "priced",
"prices": [{"points_cost": 1}],
},
{
"alias": "vision-without-image",
"operation_type": "vision",
"requires_image": False,
"pricing_status": "priced",
"prices": [{"points_cost": 1}],
},
{
"alias": "vision-free",
"operation_type": "vision",
"requires_image": True,
"pricing_status": "unpriced",
"prices": [],
},
],
}
)
@@ -2901,6 +2933,10 @@ class GuiTests(TempDirMixin, unittest.TestCase):
tab.cmhub_image_alias_combo.itemText(index)
for index in range(tab.cmhub_image_alias_combo.count())
]
vision_aliases = [
tab.cmhub_vision_alias_combo.itemData(index)
for index in range(tab.cmhub_vision_alias_combo.count())
]
self.assertIn("title-priced", title_aliases)
self.assertIn("title-saved", title_aliases)
self.assertNotIn("title-free", title_aliases)
@@ -2908,7 +2944,12 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assertIn("512:1点", tab.cmhub_title_alias_combo.itemText(0))
self.assertTrue(any("需参考图" in label for label in image_labels))
self.assertTrue(any("默认档" in label for label in image_labels))
self.assertIn("vision-priced", vision_aliases)
self.assertIn("vision-saved", vision_aliases)
self.assertNotIn("vision-without-image", vision_aliases)
self.assertNotIn("vision-free", vision_aliases)
self.assertIn("余额 55", tab.cmhub_result_label.text())
self.assertIn("当前已保存值暂不可用", tab.cmhub_result_label.text())
self.assert_removed(temp_dir)
+66
View File
@@ -929,6 +929,61 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_ai_write_uses_first_eight_originals_in_source_order(self):
with self.make_temp_dir() as temp_dir:
config = self._config(temp_dir)
project, assets = self._create_project_with_assets(temp_dir, config, 9)
tab = ProductSuiteTab(config=config, db_path=config["db_path"])
self.addCleanup(tab.close)
state = tab._displayed_state
state.account_alias = "alias-a"
state.item_id = project.item_id
state.project_id = project.id
state.project_binding_state = project.binding_state
tab._load_state(state)
captured = {}
class _Signal:
def connect(self, callback):
self.callback = callback
class _AiWriteWorker:
def __init__(self, instruction, context, **kwargs):
captured["instruction"] = instruction
captured["context"] = context
captured["image_paths"] = list(kwargs.get("image_paths") or [])
self.finished = _Signal()
self.cancelled = _Signal()
self.failed = _Signal()
def cancel(self):
pass
with mock.patch(
"app.gui.tabs.product_suite.ProductSuiteAiWriteWorker",
_AiWriteWorker,
), mock.patch.object(tab, "_start_thread", return_value=object()), mock.patch.object(
tab, "_status"
) as status:
tab.start_ai_write()
self.assertEqual(
[asset.local_path for asset in assets[:8]],
captured["image_paths"],
)
self.assertTrue(
any(
"已使用前8张商品原图进行理解" in str(call.args[0])
for call in status.call_args_list
)
)
state.ai_worker = None
state.ai_thread = None
state.ai_started_at = None
tab._apply_running_state(state)
self.assert_removed(temp_dir)
def test_prompt_edit_expands_shrinks_and_reflows_without_internal_scrollbars(self):
edit = AutoHeightPlainTextEdit()
self.addCleanup(edit.close)
@@ -2219,6 +2274,7 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
def __init__(self, instruction, context, **kwargs):
captured["instruction"] = instruction
captured["context"] = context
captured["image_paths"] = list(kwargs.get("image_paths") or [])
self.finished = _Signal()
self.cancelled = _Signal()
self.failed = _Signal()
@@ -2233,6 +2289,16 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
tab.start_ai_write()
self.assertIn("未绑定商品", captured["context"])
self.assertNotIn("draft_", captured["context"])
self.assertEqual(
[
image_studio.list_assets(
draft.id,
kind=image_studio.ASSET_KIND_ORIGINAL,
path=config["db_path"],
)[0].local_path
],
captured["image_paths"],
)
state.ai_worker = None
state.ai_thread = None
state.ai_started_at = None
+31
View File
@@ -20,6 +20,7 @@ from app.workers import BaseWorker, run_worker
from app.gui.workers import (
ImageStudioDownloadOriginalWorker,
ImageStudioPullImagesWorker,
ProductSuiteAiWriteWorker,
ProductSuiteGenerateWorker,
)
@@ -198,6 +199,36 @@ class WorkerTests(unittest.TestCase):
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")