test(ai-studio): add e2e integration checklist
This commit is contained in:
@@ -0,0 +1,159 @@
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
from _helpers import TempDirMixin
|
||||
|
||||
from app import db, image_studio, image_studio_export, image_studio_generation, image_studio_images
|
||||
|
||||
|
||||
class ImageStudioE2ETests(TempDirMixin, unittest.TestCase):
|
||||
def _png_bytes(self, color=(80, 130, 220)):
|
||||
from PIL import Image
|
||||
|
||||
output = io.BytesIO()
|
||||
Image.new("RGB", (32, 32), color).save(output, format="PNG")
|
||||
return output.getvalue()
|
||||
|
||||
def _config(self, temp_dir):
|
||||
return {
|
||||
"data_dir": temp_dir,
|
||||
"db_path": os.path.join(temp_dir, "cmshopee.db"),
|
||||
"image_dir": os.path.join(temp_dir, "images"),
|
||||
"ai": {
|
||||
"backend": "cmhub",
|
||||
"image_concurrency": 2,
|
||||
"retry": 0,
|
||||
"resolution": "1k",
|
||||
"jpg_quality": 90,
|
||||
"cmhub": {
|
||||
"base_url": "https://cmhub.example.com",
|
||||
"image_alias": "image-hd",
|
||||
"connect_timeout": 3,
|
||||
"download_with_curl": "false",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
def _runtime(self):
|
||||
return {
|
||||
"base_url": "https://cmhub.example.com",
|
||||
"api_key": "sk-test",
|
||||
"alias": "image-hd",
|
||||
"connect_timeout": 3,
|
||||
"use_system_proxy": False,
|
||||
"download_with_curl": "false",
|
||||
}
|
||||
|
||||
def test_mock_ai_studio_project_to_export_flow(self):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
cfg = self._config(temp_dir)
|
||||
db.init_db(cfg["db_path"])
|
||||
project = image_studio.create_or_get_project(
|
||||
account_alias="alias",
|
||||
account_slug="alias_slug",
|
||||
item_id="51100639510",
|
||||
draft_prompt="完整提示词",
|
||||
path=cfg["db_path"],
|
||||
)
|
||||
original = image_studio.sync_original_asset_urls(
|
||||
project.id,
|
||||
[{"index": 1, "src": "https://cdn.example.com/original.png"}],
|
||||
path=cfg["db_path"],
|
||||
)[0]
|
||||
|
||||
remote = image_studio_images.RemoteImage(
|
||||
url=original.remote_url,
|
||||
content=self._png_bytes(),
|
||||
content_type="image/png",
|
||||
final_url=original.remote_url,
|
||||
redirected=False,
|
||||
)
|
||||
with mock.patch("app.image_studio_images.download_remote_image", return_value=remote):
|
||||
source = image_studio_images.download_original_asset(
|
||||
original.id,
|
||||
path=cfg["db_path"],
|
||||
config=cfg,
|
||||
)
|
||||
|
||||
submitted = []
|
||||
|
||||
def fake_submit(method, url, api_key, **kwargs):
|
||||
task_id = f"cmhub-task-{len(submitted) + 1}"
|
||||
submitted.append(task_id)
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": "queued",
|
||||
"call_id": f"call-{len(submitted)}",
|
||||
"points_cost": 2,
|
||||
"points_balance": 100 - len(submitted) * 2,
|
||||
}
|
||||
|
||||
def fake_poll(method, url, api_key, **kwargs):
|
||||
task_id = url.rsplit("/", 1)[-1]
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": "succeeded",
|
||||
"result": {"image_url": f"https://cdn.example.com/{task_id}.png"},
|
||||
}
|
||||
|
||||
with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \
|
||||
mock.patch(
|
||||
"app.image_studio_generation.ai._cmhub_call_with_retry",
|
||||
side_effect=fake_submit,
|
||||
), \
|
||||
mock.patch("app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll), \
|
||||
mock.patch(
|
||||
"app.image_studio_generation.ai._download_cmhub_image_with_retry",
|
||||
return_value=(self._png_bytes(color=(10, 200, 120)), 0.1),
|
||||
):
|
||||
summary = image_studio_generation.generate_image_jobs(
|
||||
project.id,
|
||||
source.id,
|
||||
"完整提示词",
|
||||
2,
|
||||
config=cfg,
|
||||
path=cfg["db_path"],
|
||||
)
|
||||
|
||||
self.assertEqual(2, summary["success"])
|
||||
generated = image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"])
|
||||
self.assertEqual(2, len(generated))
|
||||
image_studio.replace_selections(
|
||||
project.id,
|
||||
"main",
|
||||
[asset.id for asset in generated],
|
||||
path=cfg["db_path"],
|
||||
)
|
||||
image_studio.replace_selections(project.id, "detail", [source.id], path=cfg["db_path"])
|
||||
|
||||
export_parent = os.path.join(temp_dir, "exports")
|
||||
os.makedirs(export_parent)
|
||||
result = image_studio_export.export_project_selection(
|
||||
project.id,
|
||||
export_parent,
|
||||
path=cfg["db_path"],
|
||||
config=cfg,
|
||||
)
|
||||
|
||||
self.assertEqual(2, result.main_count)
|
||||
self.assertEqual(1, result.detail_count)
|
||||
self.assertEqual(
|
||||
[
|
||||
"51100639510_主图_1.jpg",
|
||||
"51100639510_主图_2.jpg",
|
||||
"51100639510_详情图_1.jpg",
|
||||
],
|
||||
[os.path.basename(item.output_path) for item in result.files],
|
||||
)
|
||||
self.assertTrue(all(os.path.isfile(item.output_path) for item in result.files))
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user