import base64 import io import json import os import sys import unittest from types import SimpleNamespace from unittest import mock sys.path.insert(0, os.path.dirname(__file__)) from _helpers import TempDirMixin from app import ai, appconfig, db class _Response: def __init__(self, payload): self.payload = payload def __enter__(self): return self def __exit__(self, exc_type, exc, tb): return False def read(self, size=-1): return json.dumps(self.payload).encode("utf-8") class AITests(TempDirMixin, unittest.TestCase): def _write_models(self, path, text=None, image=None): text = text or { "name": "Text", "category": "text", "enabled": True, "url": "https://example.invalid/v1/chat/completions", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {"temperature": 0}, } image = image or { "name": "Image", "category": "image", "enabled": True, "url": "https://example.invalid/v1/chat/completions", "model": "image-model", "api_key": "sk-image-secret", "api_type": "auto", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, } appconfig.save_ai_models_config({"models": [text, image]}, path=path) def _config(self): cfg = appconfig.default_config() cfg["ai"]["default_text_model"] = "Text" cfg["ai"]["default_image_model"] = "Image" cfg["ai"]["retry"] = 1 cfg["ai"]["resolution"] = "512" cfg["ai"]["jpg_quality"] = 80 return cfg def _collected_tasks(self, temp_dir, cfg, titles=None): titles = titles or ["旧标题A", "旧标题B"] db.init_db(cfg["db_path"]) batch_id = db.create_batch([os.path.join(temp_dir, "input.xlsx")], path=cfg["db_path"]) db.insert_tasks( batch_id, [ { "source_file_abs": os.path.join(temp_dir, "input.xlsx"), "source_sheet": "商品", "source_row": index + 2, "account_name": "Excel主店", "alias": "alias-a", "item_id": "5110063951%s" % index, } for index in range(len(titles)) ], path=cfg["db_path"], ) tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) for task, title in zip(tasks, titles): db.set_collected( task.id, title, os.path.join(temp_dir, "%s_old.jpg" % task.item_id), path=cfg["db_path"], ) return batch_id, db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) def test_gen_title_uses_configured_model_and_retries(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models(models_path) calls = [] def fake_urlopen(request, timeout=None): calls.append((request, timeout)) if len(calls) == 1: raise ai.urllib.error.URLError("temporary") return _Response({"choices": [{"message": {"content": " 新标题 "}}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): title = ai.gen_title( "优化标题", "旧标题", config=self._config(), models_path=models_path, ) self.assertEqual("新标题", title) self.assertEqual(2, len(calls)) body = json.loads(calls[-1][0].data.decode("utf-8")) self.assertEqual("text-model", body["model"]) self.assertEqual(0, body["temperature"]) self.assertNotIn("sk-text-secret", body["messages"][1]["content"]) self.assert_removed(temp_dir) def test_gen_title_accepts_openai_compatible_base_url(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models( models_path, text={ "name": "Text", "category": "text", "enabled": True, "url": "https://example.invalid/v1", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, }, ) calls = [] def fake_urlopen(request, timeout=None): calls.append(request) return _Response({"choices": [{"message": {"content": "新标题"}}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): title = ai.gen_title( "优化标题", "旧标题", config=self._config(), models_path=models_path, ) self.assertEqual("新标题", title) self.assertEqual( "https://example.invalid/v1/chat/completions", calls[0].full_url, ) self.assert_removed(temp_dir) def test_missing_model_fields_raise_clear_error_without_secret(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models( models_path, text={ "name": "Text", "category": "text", "enabled": True, "url": "", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, }, ) with self.assertRaises(ai.AIError) as raised: ai.gen_title("prompt", "old", config=self._config(), models_path=models_path) message = str(raised.exception) self.assertIn("缺少字段: url", message) self.assertNotIn("sk-text-secret", message) self.assert_removed(temp_dir) def test_gen_cover_saves_jpeg_with_resolution_and_quality(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models(models_path) old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") generated = io.BytesIO() Image.new("RGB", (8, 8), (200, 120, 80)).save(generated, "PNG") b64_image = base64.b64encode(generated.getvalue()).decode("ascii") def fake_urlopen(request, timeout=None): body = json.loads(request.data.decode("utf-8")) self.assertEqual("image-model", body["model"]) self.assertIn("目标分辨率:512", body["messages"][0]["content"][0]["text"]) return _Response({"data": [{"b64_json": b64_image}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): result = ai.gen_cover( "生成封面", old_cover, output, config=self._config(), models_path=models_path, ) self.assertEqual(os.path.abspath(output), result) with Image.open(output) as saved: self.assertEqual((512, 512), saved.size) self.assertEqual("JPEG", saved.format) self.assert_removed(temp_dir) def test_generate_batch_persists_titles_and_covers_per_task(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["title_concurrency"] = 2 cfg["ai"]["image_concurrency"] = 2 batch_id, tasks = self._collected_tasks(temp_dir, cfg) cover_prompts = [] progress = [] def fake_title(title_prompt, old_title, **kwargs): self.assertEqual("标题提示", title_prompt) return "新" + old_title def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): cover_prompts.append(cover_prompt) os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title", side_effect=fake_title), \ mock.patch("app.ai.gen_cover", side_effect=fake_cover): summary = ai.generate_batch( tasks, { "title": "标题提示", "cover": "封面 {新标题} {店铺} {商品id}", }, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "account_by_alias": { "alias-a": SimpleNamespace(account_name="主店", slug="main") }, }, on_progress=progress.append, ) self.assertTrue(summary["ok"]) self.assertEqual(2, summary["title_done"]) self.assertEqual(2, summary["cover_done"]) self.assertEqual(0, summary["failed"]) self.assertTrue(progress) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(task.stage == "generated" for task in updated)) self.assertEqual({"新旧标题A", "新旧标题B"}, {task.new_title for task in updated}) expected_cover_paths = { os.path.abspath( os.path.join( cfg["image_dir"], str(task.batch_id), "main", f"{task.id}_{task.item_id}_new.jpg", ) ) for task in updated } self.assertEqual(expected_cover_paths, {task.new_cover_path for task in updated}) self.assertTrue(all(os.path.exists(task.new_cover_path) for task in updated)) self.assertIn("主店", "\n".join(cover_prompts)) self.assertIn("新旧标题A", "\n".join(cover_prompts)) self.assert_removed(temp_dir) def test_generate_batch_marks_failed_task_without_blocking_others(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["好标题", "坏标题"]) def fake_title(title_prompt, old_title, **kwargs): if old_title == "坏标题": raise ai.AIError("标题生成失败") return "新" + old_title def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title", side_effect=fake_title), \ mock.patch("app.ai.gen_cover", side_effect=fake_cover): summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertFalse(summary["ok"]) self.assertEqual(1, summary["title_done"]) self.assertEqual(1, summary["cover_done"]) self.assertEqual(1, summary["failed"]) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) by_title = {task.old_title: task for task in updated} self.assertEqual("generated", by_title["好标题"].stage) self.assertEqual("success", by_title["好标题"].status) self.assertEqual("collected", by_title["坏标题"].stage) self.assertEqual("failed", by_title["坏标题"].status) self.assertIn("标题生成失败", by_title["坏标题"].last_error) self.assert_removed(temp_dir) def test_generate_batch_stop_before_scheduling_keeps_tasks_collected(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg) with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, should_stop=lambda: True, ) self.assertFalse(summary["ok"]) self.assertTrue(summary["cancelled"]) self.assertEqual(0, summary["title_done"]) self.assertEqual(0, summary["cover_done"]) self.assertEqual(0, summary["failed"]) gen_title.assert_not_called() gen_cover.assert_not_called() updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(task.stage == "collected" for task in updated)) self.assert_removed(temp_dir) if __name__ == "__main__": unittest.main()