Coverage for tests/test_wrapper_usp.py: 100%
52 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 gc
5import pytest
7from typing import override
8from typing import Dict
9from typing import Optional
10from typing import Union
11from typing import Any
13from unittest.mock import patch
14from unittest.mock import MagicMock
15from tests.torch_mock import TorchMock
17mock_torch = TorchMock()
19mock_modules = {
20 'nvidia_smi': MagicMock(),
21 'torch': mock_torch,
22 'xfuser': MagicMock(),
23 'xfuser.config': MagicMock(),
24 'xfuser.core': MagicMock(),
25 'xfuser.core.distributed': MagicMock(),
26}
27mock_modules.update(mock_torch.get_sub_modules())
29with patch.dict(sys.modules, mock_modules):
30 from wrapper_usp import USPGeneration
33class MockUSPGeneration(USPGeneration):
34 def __init__(self) -> None:
35 super().__init__("test_usp")
37 async def get_rest_args(
38 self,
39 data_json: Dict[str, Union[str, int, float]]
40 ) -> Dict[str, Any]:
41 return data_json
43 async def warmup(self) -> None:
44 return await self.generate()
46 @override
47 async def generate(
48 self,
49 job_id: Optional[str] = None,
50 ) -> Any:
51 gen_timer = self._new_gen_timer(job_id)
52 try:
53 self._assert_model_init()
54 return await super().generate()
55 finally:
56 gen_timer.end()
59@pytest.mark.asyncio
60async def test_wrapper_usp() -> None:
61 model = MockUSPGeneration()
62 assert model is not None
63 assert model.model_name == "test_usp"
65 model.init()
66 assert model.status == "ok"
68 health = model.get_health()
69 assert health is not None
70 timestamps = model.get_timestamps()
71 assert timestamps is not None
73 await model.get_rest_args({})
75 with pytest.raises(NotImplementedError):
76 await model.warmup()
77 with pytest.raises(NotImplementedError):
78 await model.generate()
79 with pytest.raises(NotImplementedError):
80 await model.generate("job0")
82 health = model.get_health()
83 assert health is not None
85 del model
86 gc.collect()