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
« 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 PIL import Image
7from unittest.mock import patch
8from unittest.mock import MagicMock
10from typing import override
11from typing import Dict
12from typing import Any
13from typing import Optional
15sys.path.append("wrapper")
17from wrapper_model import ModelGeneration
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
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
40 async def warmup(self) -> None:
41 return await self.generate()
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()
71def test_abstract_wrapper() -> None:
72 with pytest.raises(TypeError):
73 ModelGeneration()
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"
82 model.init()
83 assert model.status == "ok"
85 health = model.get_health()
86 assert health is not None
87 timestamps = model.get_timestamps()
88 assert timestamps is not None
90 gen_args = await model.get_rest_args({})
91 assert gen_args is not None
93 with pytest.raises(ValueError, match="Missing JSON body"):
94 await model.get_rest_args(None) # type: ignore[arg-type]
96 await model.warmup()
98 resp = await model.generate()
99 assert resp == "reply"
101 resp = await model.generate("job0")
102 assert resp == "reply"
104 resp = await model.generate("job1")
105 assert resp == "reply"
107 # Data types
108 model.output_type = "str"
109 resp = await model.generate("job2")
110 assert resp == "reply"
112 model.output_type = "bytes"
113 resp = await model.generate("job3")
114 assert resp == b"reply"
116 model.output_type = "list_str"
117 resp = await model.generate("job4")
118 assert resp == ["reply1", "reply2"]
120 model.output_type = "pillow"
121 resp = await model.generate("job5")
122 assert resp.size == (64, 64)
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)
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"
136 model.output_type = "unknown"
137 resp = await model.generate("job8")
138 assert resp == "reply"
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
146 model.interrupt()
148 model.get_gpu_info()
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
157 model = MockModelGeneration()
158 model.init()
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
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):
180 gpu_info = model.get_gpu_info()
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
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)
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
200 # gpu_setup must remain True so future calls are not suppressed
201 assert model.gpu_setup is True
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"
210 health = model.get_health()
211 assert health is not None