Coverage for tests/test_wrapper_model.py: 100%

132 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 PIL import Image 

7from unittest.mock import patch 

8from unittest.mock import MagicMock 

9 

10from typing import override 

11from typing import Dict 

12from typing import Any 

13from typing import Optional 

14 

15sys.path.append("wrapper") 

16 

17from wrapper_model import ModelGeneration 

18 

19 

20class MockModelGeneration(ModelGeneration): 

21 def __init__( 

22 self, 

23 output_type: str = "str", 

24 ) -> None: 

25 super().__init__("test_model") 

26 self.output_type = output_type 

27 

28 async def get_rest_args(self, data_json: Dict[str, str]) -> Dict[str, Any]: 

29 if data_json is None or not isinstance(data_json, dict): 

30 raise ValueError("Missing JSON body") 

31 job_id = data_json.get("job_id", "default_job") 

32 ret = { 

33 "task": self.model_name, 

34 "args": { 

35 "job_id": job_id, 

36 } 

37 } 

38 return ret 

39 

40 async def warmup(self) -> None: 

41 return await self.generate() 

42 

43 @override 

44 async def generate( 

45 self, 

46 job_id: Optional[str] = None, 

47 ) -> Any: 

48 gen_timer = self._new_gen_timer(job_id) 

49 try: 

50 self._assert_model_init() 

51 if self.output_type == "str": 

52 return "reply" 

53 if self.output_type == "bytes": 

54 return b"reply" 

55 if self.output_type == "list_str": 

56 return ["reply1", "reply2"] 

57 if self.output_type == "list_bytes": 

58 return [b"reply1", b"reply2"] 

59 if self.output_type == "pillow": 

60 return Image.new('RGB', (64, 64), color="red") 

61 if self.output_type == "list_pillow": 

62 return [ 

63 Image.new('RGB', (64, 64), color="red"), 

64 Image.new('RGB', (64, 64), color="blue") 

65 ] 

66 return "reply" 

67 finally: 

68 gen_timer.end() 

69 

70 

71def test_abstract_wrapper() -> None: 

72 with pytest.raises(TypeError): 

73 ModelGeneration() 

74 

75 

76@pytest.mark.asyncio 

77async def test_wrapper_model() -> None: 

78 model = MockModelGeneration() 

79 assert model is not None 

80 assert model.model_name == "test_model" 

81 

82 model.init() 

83 assert model.status == "ok" 

84 

85 health = model.get_health() 

86 assert health is not None 

87 timestamps = model.get_timestamps() 

88 assert timestamps is not None 

89 

90 gen_args = await model.get_rest_args({}) 

91 assert gen_args is not None 

92 

93 with pytest.raises(ValueError, match="Missing JSON body"): 

94 await model.get_rest_args(None) # type: ignore[arg-type] 

95 

96 await model.warmup() 

97 

98 resp = await model.generate() 

99 assert resp == "reply" 

100 

101 resp = await model.generate("job0") 

102 assert resp == "reply" 

103 

104 resp = await model.generate("job1") 

105 assert resp == "reply" 

106 

107 # Data types 

108 model.output_type = "str" 

109 resp = await model.generate("job2") 

110 assert resp == "reply" 

111 

112 model.output_type = "bytes" 

113 resp = await model.generate("job3") 

114 assert resp == b"reply" 

115 

116 model.output_type = "list_str" 

117 resp = await model.generate("job4") 

118 assert resp == ["reply1", "reply2"] 

119 

120 model.output_type = "pillow" 

121 resp = await model.generate("job5") 

122 assert resp.size == (64, 64) 

123 

124 model.output_type = "list_pillow" 

125 resp = await model.generate("job6") 

126 assert len(resp) == 2 

127 assert resp[0].size == (64, 64) 

128 assert resp[1].size == (64, 64) 

129 

130 model.output_type = "list_bytes" 

131 resp = await model.generate("job7") 

132 assert len(resp) == 2 

133 assert resp[0] == b"reply1" 

134 assert resp[1] == b"reply2" 

135 

136 model.output_type = "unknown" 

137 resp = await model.generate("job8") 

138 assert resp == "reply" 

139 

140 # Other methods 

141 health = model.get_health() 

142 assert health is not None 

143 timestamps = model.get_timestamps() 

144 assert timestamps is not None 

145 

146 model.interrupt() 

147 

148 model.get_gpu_info() 

149 

150 

151def test_get_gpu_info_mig() -> None: 

152 """Test that get_gpu_info handles MIG (Not Supported) errors gracefully.""" 

153 import nvidia_smi 

154 import torch 

155 from nvidia_smi import NVMLError 

156 

157 model = MockModelGeneration() 

158 model.init() 

159 

160 mock_handle = MagicMock() 

161 mock_mem_info = MagicMock() 

162 mock_mem_info.used = 4 * 1024 ** 3 # 4 GiB 

163 mock_mem_info.total = 40 * 1024 ** 3 # 40 GiB 

164 

165 # Simulate a MIG instance: memory info works, but utilization / power / 

166 # temperature / clock queries all raise NVMLError("Not Supported"). 

167 with patch.object(nvidia_smi, 'nvmlInit'), \ 

168 patch.object(nvidia_smi, 'nvmlShutdown'), \ 

169 patch.object(nvidia_smi, 'nvmlDeviceGetCount', return_value=1), \ 

170 patch.object(nvidia_smi, 'nvmlDeviceGetHandleByIndex', return_value=mock_handle), \ 

171 patch.object(nvidia_smi, 'nvmlDeviceGetName', return_value="MIG 1g.5gb"), \ 

172 patch.object(nvidia_smi, 'nvmlDeviceGetMemoryInfo', return_value=mock_mem_info), \ 

173 patch.object(nvidia_smi, 'nvmlDeviceGetUtilizationRates', side_effect=NVMLError("Not Supported")), \ 

174 patch.object(nvidia_smi, 'nvmlDeviceGetTemperature', side_effect=NVMLError("Not Supported")), \ 

175 patch.object(nvidia_smi, 'nvmlDeviceGetPowerUsage', side_effect=NVMLError("Not Supported")), \ 

176 patch.object(nvidia_smi, 'nvmlDeviceGetEnforcedPowerLimit', side_effect=NVMLError("Not Supported")), \ 

177 patch.object(nvidia_smi, 'nvmlDeviceGetClockInfo', side_effect=NVMLError("Not Supported")), \ 

178 patch.object(torch.cuda, 'current_device', return_value=0): 

179 

180 gpu_info = model.get_gpu_info() 

181 

182 assert gpu_info is not None, "get_gpu_info should return results even when some NVML calls are not supported" 

183 assert len(gpu_info) == 1 

184 

185 g = gpu_info[0] 

186 assert g["name"] == "MIG 1g.5gb" 

187 assert g["mem_gib_used"] == pytest.approx(4.0) 

188 assert g["mem_gib_total"] == pytest.approx(40.0) 

189 

190 # Fields not supported by MIG instances should be None (not raise an exception) 

191 assert g["sm_util"] is None 

192 assert g["mem_util"] is None 

193 assert g["temp"] is None 

194 assert g["power_draw_watts"] is None 

195 assert g["power_limit_watts"] is None 

196 assert g["graphics_clock"] is None 

197 assert g["sm_clock"] is None 

198 assert g["mem_clock"] is None 

199 

200 # gpu_setup must remain True so future calls are not suppressed 

201 assert model.gpu_setup is True 

202 

203 

204@pytest.mark.asyncio 

205async def test_wrapper_health() -> None: 

206 model = MockModelGeneration() 

207 assert model is not None 

208 assert model.model_name == "test_model" 

209 

210 health = model.get_health() 

211 assert health is not None