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

1#!/usr/bin/env python3 

2 

3import sys 

4import pytest 

5 

6from typing import Any 

7 

8from unittest.mock import patch 

9from unittest.mock import MagicMock 

10from tests.torch_mock import TorchMock 

11from tests.diffusers_mock import DiffusersMock 

12 

13from PIL import Image 

14 

15mock_torch = TorchMock() 

16mock_diffusers = DiffusersMock() 

17 

18sys.path.append("wrapper") 

19sys.path.append("wrapper/fluxkontext") 

20sys.path.append("wrapper/flux") 

21 

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()) 

39 

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'] 

44 

45 

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" 

52 

53 img = Image.new("RGB", (40, 30)) 

54 img_base64 = img_to_base64(img) 

55 

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") 

62 

63 model.init() 

64 assert model.status == "ok" 

65 

66 health = model.get_health() 

67 assert health is not None 

68 timestamps = model.get_timestamps() 

69 assert timestamps is not None 

70 

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 }) 

80 

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 

90 

91 await model.warmup() 

92 

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) 

102 

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 ) 

111 

112 del model 

113 

114 

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" 

121 

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) 

130 

131 pipeline_instance = model.pipeline 

132 

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 

142 

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) 

151 

152 model.world_size = 2 

153 with patch.object(_fluxkontext_module, 'parallelize_transformer'): 

154 model.init_model_parallelism() 

155 

156 model.torch_compile = False 

157 model.model_compile() 

158 

159 del model