Coverage for tests/test_wrapper_fluxkontext.py: 100%
80 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
3import sys
4import pytest
6from typing import Any
8from unittest.mock import patch
9from unittest.mock import MagicMock
10from tests.torch_mock import TorchMock
11from tests.diffusers_mock import DiffusersMock
13from PIL import Image
15mock_torch = TorchMock()
16mock_diffusers = DiffusersMock()
18sys.path.append("wrapper")
19sys.path.append("wrapper/fluxkontext")
20sys.path.append("wrapper/flux")
22mock_modules = {
23 'nvidia_smi': MagicMock(),
24 'imageio': MagicMock(),
25 'cv2': MagicMock(),
26 'torch': mock_torch,
27 'xfuser': MagicMock(),
28 'xfuser.config': MagicMock(),
29 'xfuser.core': MagicMock(),
30 'xfuser.core.distributed': MagicMock(),
31 'xfuser.model_executor': MagicMock(),
32 'xfuser.model_executor.models': MagicMock(),
33 'xfuser.model_executor.models.transformers.transformer_flux': MagicMock(),
34 'xfuser.model_executor.layers': MagicMock(),
35 'xfuser.model_executor.layers.attention_processor': MagicMock(),
36}
37mock_modules.update(mock_torch.get_sub_modules())
38mock_modules.update(mock_diffusers.get_sub_modules())
40with patch.dict(sys.modules, mock_modules):
41 from image_utils import img_to_base64
42 from fluxkontext.wrapper_fluxkontext import FluxKontextGeneration
43 _fluxkontext_module = sys.modules['fluxkontext.wrapper_fluxkontext']
46@pytest.mark.asyncio
47async def test_wrapper_fluxkontext() -> None:
48 model = FluxKontextGeneration()
49 assert model is not None
50 assert model.model_name == "fluxkontext"
51 assert model.status == "initializing"
53 img = Image.new("RGB", (40, 30))
54 img_base64 = img_to_base64(img)
56 with pytest.raises(ValueError, match="Model not initialized"):
57 await model.generate(
58 img=img,
59 width=128,
60 height=80,
61 prompt="test prompt")
63 model.init()
64 assert model.status == "ok"
66 health = model.get_health()
67 assert health is not None
68 timestamps = model.get_timestamps()
69 assert timestamps is not None
71 with pytest.raises(ValueError, match="Missing JSON body"):
72 await model.get_rest_args(None)
73 with pytest.raises(ValueError, match="Missing 'img' parameter"):
74 await model.get_rest_args({})
75 with pytest.raises(ValueError, match="Missing 'prompt' parameter"):
76 # Missing prompt
77 await model.get_rest_args({
78 "img": img_base64,
79 })
81 # Success case
82 args = await model.get_rest_args({
83 "job_id": "unittest",
84 "prompt": "Test prompt",
85 "width": 80,
86 "height": 60,
87 "img": img_base64
88 })
89 assert "args" in args
91 await model.warmup()
93 image = await model.generate(
94 img=img,
95 width=256,
96 height=160,
97 prompt="Test prompt",
98 )
99 assert image is not None
100 assert isinstance(image, Image.Image)
101 assert image.size == (256, 160)
103 model.world_size = 4
104 with pytest.raises(ValueError, match="48x48 not supported for 4 GPUs"):
105 await model.generate(
106 img=img,
107 width=48,
108 height=48,
109 prompt="Test prompt",
110 )
112 del model
115@pytest.mark.asyncio
116async def test_additional_coverage() -> None:
117 """Cover seed path, step callbacks, parallelism init, and compile-disabled path."""
118 model = FluxKontextGeneration()
119 model.init()
120 assert model.status == "ok"
122 img = Image.new("RGB", (256, 160))
123 image = await model.generate(
124 img=img,
125 width=256,
126 height=160,
127 prompt="Seed coverage test",
128 seed=42)
129 assert isinstance(image, Image.Image)
131 pipeline_instance = model.pipeline
133 def _pipeline_with_callback(*args: Any, **kwargs: Any) -> Any:
134 n_steps = kwargs.get("num_inference_steps", 2)
135 callback = kwargs.get("callback_on_step_end")
136 if callback:
137 for step in range(n_steps):
138 callback(pipeline_instance, step, 0, {})
139 out = MagicMock()
140 out.images = [Image.new("RGB", (kwargs.get("width", 64), kwargs.get("height", 64)))]
141 return out
143 pipeline_instance.side_effect = _pipeline_with_callback
144 image = await model.generate(
145 img=img,
146 width=256,
147 height=160,
148 prompt="Callback coverage test",
149 sampling_steps=2)
150 assert isinstance(image, Image.Image)
152 model.world_size = 2
153 with patch.object(_fluxkontext_module, 'parallelize_transformer'):
154 model.init_model_parallelism()
156 model.torch_compile = False
157 model.model_compile()
159 del model