ImageGen-Studio / tests /test_image_inputs.py
BlueSkyXN's picture
Unify image references and guided UI/API workflows
ca9a89c verified
Raw History Blame Contribute Delete
10.3 kB
from __future__ import annotations
import unittest
import itertools
import types
from unittest import mock
from PIL import Image
from core.execution_plan import available_run_modes, build_execution_plan, ExecutionPlanError
from core.reference_inputs import normalize_image_references, reference_choices
from mcp_tools.get_model_features import handle_get_model_features
from mcp_tools.run import handle_run
from tests.test_execution_plan import base_inputs
EDIT = "lightx2v/Qwen-Image-Edit-2511-Lightning"
ANIMA = "circlestone-labs/Anima-Turbo-v1.0"
class ImageInputContractTests(unittest.TestCase):
def test_aliases_are_one_based_and_idempotent(self):
prompt = "将图1图2合并;img1 在左、IMG2在右;image 1 保留主体,参考图2提供背景。"
expected = "将image 1 image 2合并;image 1 在左、image 2在右;image 1 保留主体,image 2提供背景。"
self.assertEqual(normalize_image_references(prompt, 2), expected)
self.assertEqual(normalize_image_references(expected, 2), expected)
def test_filenames_identifiers_and_quoted_literals_are_not_rewritten(self):
prompt = 'img1.png /images/img2 myimg3 img4_suffix "图9" “img8” `image7` 「图6」'
self.assertEqual(normalize_image_references(prompt, 0), prompt)
def test_chinese_compound_words_are_not_image_aliases(self):
prompts = (
"电影感,横向构图16:9,日落风景",
"横向构图 16:9,地图1:10000,草图2张,插图3幅",
"蓝图1份,位图2张,视图3个,绘图4幅,截图5张,贴图6张",
)
for prompt in prompts:
for count in (0, 2):
with self.subTest(prompt=prompt, count=count):
self.assertEqual(normalize_image_references(prompt, count), prompt)
plan = build_execution_plan(base_inputs(positive_prompt=prompt))
self.assertEqual(plan[0].inputs["positive_prompt"], prompt)
def test_explicit_aliases_next_to_chinese_text_and_numeric_details_still_work(self):
prompt = "横向构图16:9;保留图1,图1图2合并;图1:2只猫,参考图2:1只狗,图片1和img2。"
expected = "横向构图16:9;保留image 1,image 1 image 2合并;image 1:2只猫,image 2:1只狗,image 1和image 2。"
self.assertEqual(normalize_image_references(prompt, 2), expected)
self.assertEqual(normalize_image_references(expected, 2), expected)
with self.assertRaisesRegex(ExecutionPlanError, "本次只有 2"):
normalize_image_references("横向构图16:9;保留图3", 2)
def test_missing_and_cross_image_references_fail_early(self):
for prompt in ("图0", "img0", "图3", "IMG3", "image 3"):
with self.subTest(prompt=prompt), self.assertRaisesRegex(ExecutionPlanError, "本次只有 2"):
normalize_image_references(prompt, 2)
with self.assertRaisesRegex(ExecutionPlanError, "当前图片"):
normalize_image_references("保留图1", 1, independent=True)
def test_high_level_list_excludes_stale_slots_without_mutating_them(self):
stale = [Image.new("RGB", (8, 8), "black"), Image.new("RGB", (8, 8), "white")]
pictures = [Image.new("RGB", (8, 8), "red"), Image.new("RGB", (8, 8), "blue")]
values = base_inputs(model_display_name=EDIT, task_type="reference", positive_prompt="图1左、img2右", qwen_image_edit_data=stale)
plan = build_execution_plan(values, images=pictures)
self.assertEqual(values["qwen_image_edit_data"], stale)
self.assertEqual(values["positive_prompt"], "图1左、img2右")
self.assertEqual(plan[0].inputs["qwen_image_edit_data"], pictures)
self.assertEqual(plan[0].inputs["positive_prompt"], "image 1左、image 2右")
self.assertEqual(plan[0].inputs["task_type"], "txt2img")
self.assertIn("2 张参考图", plan[0].caption)
self.assertEqual([item["id"] for item in plan[0].inputs["_image_references"]], ["img1", "img2"])
reordered = build_execution_plan(values, images=list(reversed(pictures)))
self.assertIs(reordered[0].inputs["qwen_image_edit_data"][0], pictures[1])
def test_reference_limit_never_truncates(self):
with self.assertRaisesRegex(ExecutionPlanError, "3 张"):
build_execution_plan(base_inputs(task_type="reference", model_display_name=EDIT), images=["a", "b", "c", "d"])
with self.assertRaisesRegex(ExecutionPlanError, "不会截断"):
build_execution_plan(base_inputs(model_display_name=EDIT, qwen_image_edit_data=["a", "b", "c", "d"]))
def test_supported_native_slots_compact_and_unsupported_slots_do_not_count(self):
values = base_inputs(model_display_name=EDIT, positive_prompt="图1", qwen_image_edit_data=[None, "a", None])
plan = build_execution_plan(values)
self.assertEqual(len(plan[0].inputs["_image_references"]), 1)
plan = build_execution_plan(base_inputs(task_type="img2img", model_display_name=ANIMA, img2img_image="source", positive_prompt="图1", qwen_image_edit_data=["old"]))
self.assertEqual(plan[0].inputs["qwen_image_edit_data"], [])
self.assertEqual(plan[0].inputs["positive_prompt"], "image 1")
def test_discovery_matches_checkpoint_limits_and_modes(self):
edit = handle_get_model_features(EDIT)["image_workflows"]
self.assertEqual(edit["reference_roles"]["auto"]["max_images"], 3)
self.assertNotIn("style", edit["reference_roles"])
regular = handle_get_model_features("Qwen-Image")["image_workflows"]
self.assertFalse(regular["reference_supported"])
self.assertEqual(regular["run_modes_by_task"]["reference"], [])
self.assertNotIn("multi_reference", [mode for _, mode in available_run_modes("img2img", EDIT)])
self.assertNotIn("multi_independent", [mode for _, mode in available_run_modes("inpaint", EDIT)])
self.assertIn("identity", reference_choices("Krea-2-Turbo"))
def test_qwen_injector_wires_ordered_images_to_both_prompt_nodes(self):
from chain_injectors.qwen_image_edit_injector import inject
ids = itertools.count(10)
assembler = types.SimpleNamespace(
node_map={"ksampler": "1", "pos_prompt": "2", "neg_prompt": "3", "vae_loader": "4"},
workflow={
"1": {"class_type": "KSampler", "inputs": {"model": ["5", 0]}},
"2": {"class_type": "TextEncodeQwenImageEditPlus", "inputs": {"prompt": "image 1左,image 2右"}},
"3": {"class_type": "TextEncodeQwenImageEditPlus", "inputs": {"prompt": ""}},
"4": {"class_type": "VAELoader", "inputs": {}},
},
_get_unique_id=lambda: str(next(ids)),
_get_node_template=lambda kind: {"class_type": kind, "inputs": {}, "_meta": {}},
)
inject(assembler, {}, ["blue.png", "red.png"])
for node_id in ("2", "3"):
for i, filename in enumerate(("blue.png", "red.png"), 1):
scale_id = assembler.workflow[node_id]["inputs"][f"image{i}"][0]
load_id = assembler.workflow[scale_id]["inputs"]["image"][0]
self.assertEqual(assembler.workflow[load_id]["inputs"]["image"], filename)
def test_multiple_native_namespaces_cannot_share_numbered_prompt(self):
with self.assertRaisesRegex(ExecutionPlanError, "多个原生参考链"):
build_execution_plan(base_inputs(positive_prompt="图1", krea2_identity_edit_data=["a"], krea2_reference_edit_data=["b"]))
class ApiImageRequestTests(unittest.TestCase):
def request(self, **overrides):
return dict(task_type="reference", model=EDIT, prompt="图1左,img2右", width=512, height=512, async_execution=True, **overrides)
def test_numbered_fields_sort_numerically_and_keep_caller_data(self):
request = self.request(**{"图2": "blue", "IMG1": "red"})
with mock.patch("mcp_tools.run.submit_background") as submit:
result = handle_run(request)
self.assertEqual(result["status"], "queued")
self.assertEqual(submit.call_args.args[2]["images"], ["red", "blue"])
self.assertEqual(submit.call_args.args[2]["run_mode"], "multi_reference")
self.assertEqual(request["图2"], "blue")
self.assertNotIn("images", request)
def test_invalid_image_contract_never_enters_queue(self):
cases = [
{"images": ["a"], "img1": "b"}, {"image": "a", "images": ["b"]},
{"img1": "a", "图1": "b"}, {"img2": "a"}, {"img0": "a"},
{"images": []}, {"images": [None]}, {"images": [1]},
{"images": ["a"]}, {"images": ["a", "b", "c", "d"]},
{"images": ["a", "b"], "reference_role": "style"},
{"images": ["a", "b"], "chain": [{"injector_type": "qwen_image_edit", "image": "c"}]},
{"images": ["a", "b"], "run_mode": "nonsense"},
]
for case in cases:
with self.subTest(case=case), mock.patch("mcp_tools.run.submit_background") as submit:
result = handle_run(self.request(**case))
self.assertEqual(result["error"]["code"], "INVALID_PARAMS")
submit.assert_not_called()
def test_legacy_chain_and_single_image_calls_still_queue(self):
cases = [
{"task_type": "txt2img", "model": EDIT, "prompt": "img1左,图2右", "width": 512, "height": 512,
"chain": [{"injector_type": "qwen_image_edit", "image": "a"}, {"injector_type": "qwen_image_edit", "image": "b"}]},
{"task_type": "img2img", "model": ANIMA, "prompt": "图1改为水彩", "image": "a"},
{"task_type": "img2img", "model": ANIMA, "prompt": "当前图片转水彩", "images": ["a", "b"], "run_mode": "multi_independent"},
]
for params in cases:
with self.subTest(params=params), mock.patch("mcp_tools.run.submit_background") as submit:
result = handle_run({**params, "async_execution": True})
self.assertEqual(result["status"], "queued")
submit.assert_called_once()
if __name__ == "__main__":
unittest.main()