feat(product-suite): add cmhub vision AI writing
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user