Coverage for tests/test_wrapper_januspro.py: 100%
45 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 unittest.mock import patch
7from unittest.mock import MagicMock
8from tests.torch_mock import TorchMock
9from tests.diffusers_mock import DiffusersMock
11mock_torch = TorchMock()
12mock_diffusers = DiffusersMock()
14sys.path.append("wrapper")
15sys.path.append("wrapper/januspro")
17mock_modules = {
18 'nvidia_smi': MagicMock(),
19 'imageio': MagicMock(),
20 'cv2': MagicMock(),
21 'torch': mock_torch,
22 'torch.amp': MagicMock(),
23 'torch.distributed': MagicMock(),
24 'xfuser': MagicMock(),
25 'xfuser.config': MagicMock(),
26 'xfuser.core': MagicMock(),
27 'xfuser.core.distributed': MagicMock(),
28 'xfuser.model_executor': MagicMock(),
29 'xfuser.model_executor.layers': MagicMock(),
30 'xfuser.model_executor.layers.attention_processor': MagicMock(),
31 'transformers': MagicMock(),
32 'janus': MagicMock(),
33 'janus.models': MagicMock(),
34}
35mock_modules.update(mock_diffusers.get_sub_modules())
37with patch.dict(sys.modules, mock_modules):
38 from januspro.wrapper_januspro import JanusProGeneration
41@pytest.mark.asyncio
42async def test_wrapper_januspro() -> None:
43 model = JanusProGeneration()
44 assert model is not None
45 assert model.model_name == "januspro"
46 assert model.status == "initializing"
48 with pytest.raises(ValueError):
49 await model.generate(64, 48, "test prompt")
51 model.init()
52 assert model.status == "ok"
54 # Mock pipeline return object
55 mock_output = MagicMock()
56 mock_output.images = ["image"]
57 model.pipeline = MagicMock(return_value=mock_output)
58 model.pipeline.vae_scale_factor = 8
60 health = model.get_health()
61 assert health is not None
62 timestamps = model.get_timestamps()
63 assert timestamps is not None
65 with pytest.raises(ValueError):
66 await model.get_rest_args(None)
67 with pytest.raises(ValueError):
68 await model.get_rest_args({})
69 args = await model.get_rest_args({
70 "job_id": "unittest",
71 "prompt": "Test prompt",
72 "width": 80,
73 "height": 60,
74 "seed": 7,
75 })
76 assert args is not None
77 assert "args" in args
78 assert args["args"]["prompt"] == "Test prompt"
80 with pytest.raises(ValueError, match="could not broadcast input array from shape"):
81 await model.warmup()
83 with pytest.raises(ValueError, match="could not broadcast input array from shape"):
84 await model.generate(prompt="Test prompt")
86 del model