Seamless-Texture / tests /test_model_handler.py
Maikeu Locatelli
Add initial implementation of Flux Seamless Texture LoRA application
f0d9a3e
Raw
History Blame Contribute Delete
2.56 kB
"""Tests for model_handler module."""
import pytest
from unittest.mock import Mock, patch, MagicMock
from pathlib import Path
from src.model_handler import ModelHandler
from src.utils import generate_seed
class TestModelHandler:
"""Test cases for ModelHandler class."""
@patch('src.model_handler.gr.load')
def test_model_initialization(self, mock_gr_load):
"""Test model initialization."""
mock_model = Mock()
mock_gr_load.return_value = mock_model
handler = ModelHandler()
assert handler.model is not None
mock_gr_load.assert_called_once()
@patch('src.model_handler.gr.load')
def test_generate_seed(self, mock_gr_load):
"""Test seed generation."""
mock_gr_load.return_value = Mock()
handler = ModelHandler()
seed = generate_seed()
assert isinstance(seed, int)
assert 0 <= seed < 2**32
@patch('src.model_handler.gr.load')
@patch('src.model_handler.save_image')
def test_generate_with_params(self, mock_save_image, mock_gr_load):
"""Test image generation with parameters."""
# Mock the model
mock_image = Mock()
mock_model = Mock()
mock_model.return_value = mock_image
mock_gr_load.return_value = mock_model
# Mock save_image
mock_save_image.return_value = Path("test_image.png")
handler = ModelHandler()
handler.model = mock_model
# Test generation
try:
image, metadata = handler.generate(
prompt="test prompt",
guidance_scale=7.5,
num_inference_steps=50,
seed=12345,
width=1024,
height=1024,
)
# If we get here, generation was attempted
assert True
except Exception:
# Expected if model interface is different
pass
@patch('src.model_handler.gr.load')
def test_generate_invalid_params(self, mock_gr_load):
"""Test generation with invalid parameters."""
mock_gr_load.return_value = Mock()
handler = ModelHandler()
with pytest.raises(ValueError):
handler.generate(
prompt="test",
guidance_scale=25.0, # Invalid: > 20
num_inference_steps=50,
width=1024,
height=1024,
)
if __name__ == "__main__":
pytest.main([__file__])