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

1#!/usr/bin/env python3 

2 

3import sys 

4import pytest 

5 

6from unittest.mock import patch 

7from unittest.mock import MagicMock 

8from tests.torch_mock import TorchMock 

9from tests.diffusers_mock import DiffusersMock 

10 

11mock_torch = TorchMock() 

12mock_diffusers = DiffusersMock() 

13 

14sys.path.append("wrapper") 

15sys.path.append("wrapper/januspro") 

16 

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

36 

37with patch.dict(sys.modules, mock_modules): 

38 from januspro.wrapper_januspro import JanusProGeneration 

39 

40 

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" 

47 

48 with pytest.raises(ValueError): 

49 await model.generate(64, 48, "test prompt") 

50 

51 model.init() 

52 assert model.status == "ok" 

53 

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 

59 

60 health = model.get_health() 

61 assert health is not None 

62 timestamps = model.get_timestamps() 

63 assert timestamps is not None 

64 

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" 

79 

80 with pytest.raises(ValueError, match="could not broadcast input array from shape"): 

81 await model.warmup() 

82 

83 with pytest.raises(ValueError, match="could not broadcast input array from shape"): 

84 await model.generate(prompt="Test prompt") 

85 

86 del model