Coverage for tests/test_wrapper_bagel.py: 100%

48 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 

11from PIL import Image 

12 

13mock_torch = TorchMock() 

14mock_diffusers = DiffusersMock() 

15 

16sys.path.append("wrapper") 

17sys.path.append("wrapper/bagel") 

18 

19mock_modules = { 

20 'nvidia_smi': MagicMock(), 

21 'imageio': MagicMock(), 

22 'cv2': MagicMock(), 

23 'xfuser': MagicMock(), 

24 'xfuser.config': MagicMock(), 

25 'xfuser.core': MagicMock(), 

26 'xfuser.core.distributed': MagicMock(), 

27 'xfuser.model_executor': MagicMock(), 

28 'xfuser.model_executor.layers': MagicMock(), 

29 'xfuser.model_executor.layers.attention_processor': MagicMock(), 

30 'modeling': MagicMock(), 

31 'modeling.autoencoder': MagicMock(), 

32 'modeling.qwen2': MagicMock(), 

33 'modeling.bagel': MagicMock(), 

34 'modeling.bagel.qwen2_navit': MagicMock(), 

35 'data': MagicMock(), 

36 'data.transforms': MagicMock(), 

37 'data.data_utils': MagicMock(), 

38 'accelerate': MagicMock(), 

39 'safetensors': MagicMock(), 

40 'safetensors.torch': MagicMock(), 

41} 

42mock_modules.update(mock_torch.get_sub_modules()) 

43mock_modules.update(mock_diffusers.get_sub_modules()) 

44 

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

46 from image_utils import img_to_base64 

47 from bagel.wrapper_bagel import BagelGeneration 

48 

49 

50@pytest.mark.asyncio 

51async def test_basic() -> None: 

52 model = BagelGeneration() 

53 assert model is not None 

54 assert model.model_name == "bagel" 

55 assert model.status == "initializing" 

56 

57 img = Image.new("RGB", (40, 30)) 

58 img_base64 = img_to_base64(img) 

59 

60 with pytest.raises(ValueError, match="Model not initialized"): 

61 await model.generate( 

62 imgs=[img], 

63 width=128, 

64 height=80, 

65 prompt="test prompt") 

66 

67 with pytest.raises(ValueError): 

68 model.init() 

69 assert model.status == "failed" 

70 

71 health = model.get_health() 

72 assert health is not None 

73 timestamps = model.get_timestamps() 

74 assert timestamps is not None 

75 

76 with pytest.raises(ValueError): 

77 await model.get_rest_args(None) 

78 with pytest.raises(ValueError): 

79 await model.get_rest_args({}) 

80 with pytest.raises(ValueError): 

81 # Missing prompt 

82 await model.get_rest_args({ 

83 "imgs": img_base64, 

84 }) 

85 

86 # Success case 

87 args = await model.get_rest_args({ 

88 "job_id": "unittest", 

89 "prompt": "Test prompt", 

90 "width": 80, 

91 "height": 60, 

92 "imgs": [img_base64] 

93 }) 

94 assert "args" in args 

95 assert args["task"] == "bagel" 

96 

97 with pytest.raises(ValueError, match="Model not initialized"): 

98 await model.warmup() 

99 

100 with pytest.raises(ValueError, match="Model not initialized"): 

101 await model.generate( 

102 imgs=[img], 

103 width=256, 

104 height=160, 

105 prompt="Test prompt", 

106 ) 

107 

108 del model