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
« 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
11from PIL import Image
13mock_torch = TorchMock()
14mock_diffusers = DiffusersMock()
16sys.path.append("wrapper")
17sys.path.append("wrapper/bagel")
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())
45with patch.dict(sys.modules, mock_modules):
46 from image_utils import img_to_base64
47 from bagel.wrapper_bagel import BagelGeneration
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"
57 img = Image.new("RGB", (40, 30))
58 img_base64 = img_to_base64(img)
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")
67 with pytest.raises(ValueError):
68 model.init()
69 assert model.status == "failed"
71 health = model.get_health()
72 assert health is not None
73 timestamps = model.get_timestamps()
74 assert timestamps is not None
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 })
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"
97 with pytest.raises(ValueError, match="Model not initialized"):
98 await model.warmup()
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 )
108 del model