from __future__ import annotations import importlib.util import pickle import sys import tempfile import threading import types import unittest from pathlib import Path from unittest import mock from core.task_scheduler import TaskCancelledError ROOT = Path(__file__).resolve().parents[1] class _LockedProgress: def __init__(self, *args, **kwargs): self.lock = threading.Lock() self.updates = [] def __call__(self, value, desc=None): self.updates.append((value, desc)) class _Assembler: def __init__(self, after_assembly=None): self.after_assembly = after_assembly def assemble(self, values): self.ui_values = values if self.after_assembly: self.after_assembly() return {"1": {"class_type": "TestNode", "inputs": {"seed": values["seed"]}}} def _module(name, **attributes): module = types.ModuleType(name) module.__dict__.update(attributes) return module class GpuBoundaryTests(unittest.TestCase): """执行真实 pipeline 调用链,只替换下载、推理依赖和 GPU 传输。""" def setUp(self): output_dir = self.enterContext(tempfile.TemporaryDirectory()) self.payloads = [] self.assemblers = [] self.after_assembly = None self.download = mock.Mock() self.release_models = mock.Mock() self.execute_workflow = mock.Mock( return_value=types.SimpleNamespace(shape=(0,)) ) def gpu_decorator(**_options): def decorate(function): def run(*args, **kwargs): # 真正序列化入参,不能像 UI 烟测那样直接跳过生成调用。 args, kwargs = pickle.loads(pickle.dumps((args, kwargs))) self.payloads.append(kwargs) return function(*args, **kwargs) return run return decorate def create_assembler(*_args, **_kwargs): assembler = _Assembler(self.after_assembly) self.assemblers.append(assembler) return assembler prepared_keys = ( "temp_files_to_clean", "active_loras_for_gpu", "active_loras_for_meta", "active_controlnets", "active_anima_controlnets", "active_diffsynth_controlnets", "active_ipadapters", "active_flux1_ipadapters", "active_sd3_ipadapters", "active_styles", "active_reference_latents", "active_hidream_o1_reference", "active_conditioning", ) base_name = "core.pipelines.base_pipeline" pipeline_name = "core.pipelines.sd_image_pipeline" stubs = { "gradio": _module("gradio", Error=RuntimeError, Progress=_LockedProgress), "spaces": _module("spaces", GPU=gpu_decorator), "torch": _module("torch"), "imageio": _module("imageio"), "numpy": _module("numpy"), "core.model_manager": _module( "core.model_manager", model_manager=types.SimpleNamespace(ensure_models_downloaded=self.download), release_loaded_models=self.release_models, ), "core.workflow_assembler": _module( "core.workflow_assembler", WorkflowAssembler=create_assembler ), "imagegen_utils.app_utils": _module( "imagegen_utils.app_utils", sanitize_prompt=lambda value: value ), "core.pipelines.pipeline_input_processor": _module( "core.pipelines.pipeline_input_processor", process_pipeline_inputs=lambda *_: {key: [] for key in prepared_keys}, ), "core.pipelines.workflow_executor": _module( "core.pipelines.workflow_executor", WorkflowExecutor=types.SimpleNamespace(execute_workflow=self.execute_workflow), ), } specs = [] for name in (base_name, pipeline_name): spec = importlib.util.spec_from_file_location( name, ROOT / "core" / "pipelines" / f"{name.rsplit('.', 1)[1]}.py" ) stubs[name] = importlib.util.module_from_spec(spec) specs.append(spec) self.enterContext(mock.patch.dict(sys.modules, stubs)) for spec in specs: spec.loader.exec_module(stubs[spec.name]) self.module = stubs[pipeline_name] self.module.OUTPUT_DIR = output_dir self.pipeline = self.module.SdImagePipeline() def inputs(self, **overrides): values = { "task_type": "txt2img", "model_display_name": "circlestone-labs/Anima-Turbo-v1.0", "positive_prompt": "测试图片", "negative_prompt": "", "seed": 42, "batch_size": 1, "num_inference_steps": 4, "guidance_scale": 1.0, "sampler": "euler", "scheduler": "simple", "width": 64, "height": 64, "denoise": 1.0, } values.update(overrides) return values def test_repeated_ui_calls_serialize_and_preserve_model_release(self): for release in (True, False): with self.subTest(release=release): cancellation = threading.Event() values = self.inputs( _cancel_event=cancellation, _release_models_after_run=release ) progress = _LockedProgress() self.assertEqual(self.pipeline.run(values, progress), []) payload = self.payloads[-1] self.assertEqual(set(payload), {"ui_inputs", "loras_string", "workflow"}) self.assertNotIn("_cancel_event", payload["ui_inputs"]) self.assertEqual(payload["ui_inputs"]["_release_models_after_run"], release) self.assertIs(values["_cancel_event"], cancellation) self.assertIs(self.assemblers[-1].ui_values["_cancel_event"], cancellation) self.assertTrue(progress.updates) self.assertEqual(len(self.payloads), 2) self.release_models.assert_called_once_with() def test_request_without_cancel_event_serializes(self): self.pipeline.run(self.inputs(), _LockedProgress()) self.assertNotIn("_cancel_event", self.payloads[0]["ui_inputs"]) self.release_models.assert_not_called() def test_unsupported_stale_reference_does_not_block_img2img(self): from PIL import Image self.pipeline.run(self.inputs(task_type="img2img", img2img_image=Image.new("RGB", (64, 64)), qwen_image_edit_data=[Image.new("RGB", (32, 32))], positive_prompt="保留图1"), _LockedProgress()) self.download.assert_called_once() self.assertEqual(self.payloads[0]["ui_inputs"]["qwen_image_edit_data"], []) self.assertEqual(self.payloads[0]["ui_inputs"]["positive_prompt"], "保留image 1") def test_invalid_numbered_reference_rejected_before_download_and_gpu(self): with self.assertRaisesRegex(RuntimeError, "本次只有 0"): self.pipeline.run(self.inputs(positive_prompt="img2"), _LockedProgress()) self.download.assert_not_called() self.assertEqual(self.payloads, []) def test_cancellation_after_download_stays_on_cpu(self): cancellation = threading.Event() self.download.side_effect = lambda *args, **kwargs: cancellation.set() with self.assertRaises(TaskCancelledError): self.pipeline.run(self.inputs(_cancel_event=cancellation), _LockedProgress()) self.assertEqual(self.payloads, []) self.execute_workflow.assert_not_called() def test_cancellation_after_assembly_stays_on_cpu(self): cancellation = threading.Event() self.after_assembly = cancellation.set with self.assertRaises(TaskCancelledError): self.pipeline.run(self.inputs(_cancel_event=cancellation), _LockedProgress()) self.assertEqual(self.payloads, []) self.execute_workflow.assert_not_called() def test_gpu_failure_still_releases_model_state(self): self.execute_workflow.side_effect = RuntimeError("sampling failed") with self.assertRaisesRegex(RuntimeError, "sampling failed"): self.pipeline.run( self.inputs(_cancel_event=threading.Event()), _LockedProgress() ) self.release_models.assert_called_once_with() if __name__ == "__main__": unittest.main()