Code import - branch release/SIENTIAPDE-1645
This commit is contained in:
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
0
tests/activities/__init__.py
Normal file
0
tests/activities/__init__.py
Normal file
144
tests/activities/test_activities.py
Normal file
144
tests/activities/test_activities.py
Normal file
@@ -0,0 +1,144 @@
|
||||
"""Unit tests for Activities orchestrator (constructor, shutdown, destructor)."""
|
||||
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
|
||||
def _postgres():
|
||||
return {
|
||||
'host': 'h',
|
||||
'port': 5432,
|
||||
'user': 'u',
|
||||
'password': 'p',
|
||||
'dbname': 'db',
|
||||
'min_connections': 1,
|
||||
'max_connections': 2,
|
||||
}
|
||||
|
||||
|
||||
def _mlflow():
|
||||
return {'url': 'http://mlflow:5000', 'username': 'u', 'password': 'p'}
|
||||
|
||||
|
||||
def _minio(endpoint_url: str):
|
||||
return {
|
||||
'endpoint_url': endpoint_url,
|
||||
'access_key': 'a',
|
||||
'secret_key': 's',
|
||||
'region': 'r',
|
||||
'use_ssl': True,
|
||||
'default_bucket': 'test-bucket',
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
'endpoint,expected_endpoint',
|
||||
[
|
||||
('http://minio:9000', 'minio:9000'),
|
||||
('https://minio:9000', 'minio:9000'),
|
||||
('minio:9000', 'minio:9000'),
|
||||
],
|
||||
)
|
||||
def test_activities_strips_minio_endpoint_scheme(endpoint, expected_endpoint):
|
||||
with (
|
||||
patch(
|
||||
'model_manager.activities.activities.ExperimentTracking.__init__',
|
||||
Mock(return_value=None),
|
||||
),
|
||||
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.SientiaMLflowRepository') as m_mlflow,
|
||||
patch('model_manager.activities.activities.MinioRepository') as m_minio,
|
||||
):
|
||||
Activities(
|
||||
postgres_config=_postgres(),
|
||||
mlflow_config=_mlflow(),
|
||||
minio_config=_minio(endpoint),
|
||||
plugin_store=MagicMock(),
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
m_minio.assert_called_once()
|
||||
assert m_minio.call_args.kwargs['endpoint'] == expected_endpoint
|
||||
assert m_minio.call_args.kwargs['bucket'] == 'test-bucket'
|
||||
m_mlflow.assert_called_once()
|
||||
|
||||
|
||||
def test_activities_shutdown_calls_parents():
|
||||
with (
|
||||
patch(
|
||||
'model_manager.activities.activities.ExperimentTracking.__init__',
|
||||
Mock(return_value=None),
|
||||
),
|
||||
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.SientiaMLflowRepository'),
|
||||
patch('model_manager.activities.activities.MinioRepository'),
|
||||
patch('model_manager.activities.activities.ExperimentTracking.close') as m_close,
|
||||
patch('model_manager.activities.activities.SientiaMonitoring.shutdown') as m_mon,
|
||||
):
|
||||
a = Activities(
|
||||
postgres_config=_postgres(),
|
||||
mlflow_config=_mlflow(),
|
||||
minio_config=_minio('http://x:9000'),
|
||||
plugin_store=MagicMock(),
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
with patch.object(SientiaMonitoring, 'info', Mock()):
|
||||
a.shutdown()
|
||||
m_close.assert_called_once()
|
||||
m_mon.assert_called_once()
|
||||
|
||||
|
||||
def test_activities_del_with_engine_runs_without_error():
|
||||
with (
|
||||
patch(
|
||||
'model_manager.activities.activities.ExperimentTracking.__init__',
|
||||
Mock(return_value=None),
|
||||
),
|
||||
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.SientiaMLflowRepository'),
|
||||
patch('model_manager.activities.activities.MinioRepository'),
|
||||
):
|
||||
a = Activities(
|
||||
postgres_config=_postgres(),
|
||||
mlflow_config=_mlflow(),
|
||||
minio_config=_minio('http://x:9000'),
|
||||
plugin_store=MagicMock(),
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
a.engine = MagicMock()
|
||||
Activities.__del__(a)
|
||||
|
||||
|
||||
def test_activities_del_without_engine_runs_without_error():
|
||||
with (
|
||||
patch(
|
||||
'model_manager.activities.activities.ExperimentTracking.__init__',
|
||||
Mock(return_value=None),
|
||||
),
|
||||
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.SientiaMLflowRepository'),
|
||||
patch('model_manager.activities.activities.MinioRepository'),
|
||||
):
|
||||
a = Activities(
|
||||
postgres_config=_postgres(),
|
||||
mlflow_config=_mlflow(),
|
||||
minio_config=_minio('http://x:9000'),
|
||||
plugin_store=MagicMock(),
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
Activities.__del__(a)
|
||||
310
tests/activities/test_cleanup.py
Normal file
310
tests/activities/test_cleanup.py
Normal file
@@ -0,0 +1,310 @@
|
||||
"""Unit tests for the Cleanup activity, ensuring 100% code coverage."""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
from datetime import datetime, timedelta
|
||||
from importlib import reload
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
# Define mocks at the top level to be accessible by all tests
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Fixture for a mock logger."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_notification_handler():
|
||||
"""Fixture for a mock notification handler."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_metrics_controller():
|
||||
"""Fixture for a mock metrics controller with async methods."""
|
||||
controller = MagicMock()
|
||||
controller.shutdown = AsyncMock()
|
||||
controller.emit = AsyncMock()
|
||||
return controller
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def temp_dir():
|
||||
"""Fixture to create and clean up a temporary directory."""
|
||||
path = tempfile.mkdtemp()
|
||||
yield path
|
||||
shutil.rmtree(path)
|
||||
|
||||
|
||||
# --- Initialization Tests ---
|
||||
|
||||
|
||||
@patch.dict(
|
||||
'model_manager.activities.cleanup.os.environ',
|
||||
{
|
||||
'CLEANUP_RETENTION_HOURS': '24',
|
||||
'CLEANUP_DRY_RUN': 'false',
|
||||
},
|
||||
)
|
||||
def test_cleanup_init_default_values(
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test Cleanup initialization uses default environment values."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
assert cleanup.retention_hours == 24
|
||||
assert cleanup.dry_run is False
|
||||
|
||||
|
||||
def test_cleanup_init_custom_env_values(
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test Cleanup initialization with custom environment values."""
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
'CLEANUP_RETENTION_HOURS': '48',
|
||||
'CLEANUP_DRY_RUN': 'true',
|
||||
},
|
||||
):
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
assert cleanup.retention_hours == 48
|
||||
assert cleanup.dry_run is True
|
||||
|
||||
|
||||
@patch.dict(os.environ, {'CLEANUP_RETENTION_HOURS': 'invalid'})
|
||||
def test_cleanup_init_invalid_env_value_raises_error():
|
||||
"""Test Cleanup module raises ValueError for invalid environment variables on import."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
reload(model_manager.activities.cleanup)
|
||||
|
||||
|
||||
# --- Temp Directory Cleanup Tests ---
|
||||
|
||||
|
||||
def test_cleanup_temp_directories_nonexistent_path(
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test temp directory cleanup with a non-existent path."""
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.warning = MagicMock()
|
||||
|
||||
cleanup.cleanup_temp_directories({'temp_path': '/nonexistent/path', 'metadata': {}})
|
||||
|
||||
cleanup.warning.assert_called_once()
|
||||
|
||||
|
||||
@patch.dict('model_manager.activities.cleanup.os.environ', {'CLEANUP_DRY_RUN': 'false'})
|
||||
def test_cleanup_temp_directories_success_with_deletions(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test successful deletion of old temporary directories."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
|
||||
old_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
|
||||
os.makedirs(old_dir)
|
||||
|
||||
recent_time = (datetime.now() - timedelta(hours=1)).strftime('%Y%m%d_%H%M%S_000000')
|
||||
recent_dir = os.path.join(temp_dir, f'recent_dir_{recent_time}')
|
||||
os.makedirs(recent_dir)
|
||||
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
assert not os.path.exists(old_dir)
|
||||
assert os.path.exists(recent_dir)
|
||||
|
||||
|
||||
@patch.dict(os.environ, {'CLEANUP_DRY_RUN': 'true'})
|
||||
def test_cleanup_temp_directories_dry_run(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test temp directory cleanup in dry_run mode does not delete."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
|
||||
old_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
|
||||
os.makedirs(old_dir)
|
||||
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
assert os.path.exists(old_dir)
|
||||
|
||||
|
||||
@patch.dict('model_manager.activities.cleanup.os.environ', {'CLEANUP_DRY_RUN': 'false'})
|
||||
def test_cleanup_temp_directories_delete_error(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test graceful handling of errors during directory deletion."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.error = MagicMock()
|
||||
|
||||
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
|
||||
old_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
|
||||
os.makedirs(old_dir)
|
||||
|
||||
with patch('shutil.rmtree', side_effect=OSError('Permission Denied')):
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
cleanup.error.assert_called_once()
|
||||
|
||||
|
||||
# --- Utility Tests ---
|
||||
|
||||
|
||||
def test_cleanup_temp_directories_with_files_and_unmatched_dirs(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test that files and directories with non-matching names are skipped."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.debug = MagicMock()
|
||||
|
||||
# Create a file and a directory with a non-matching name
|
||||
with open(os.path.join(temp_dir, 'a_file.txt'), 'w') as f:
|
||||
f.write('hello')
|
||||
os.makedirs(os.path.join(temp_dir, 'a_directory_with_no_timestamp'))
|
||||
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
# Ensure the debug message for skipping was called for the unmatched directory
|
||||
cleanup.debug.assert_called_with(
|
||||
'Skipping directory without timestamp pattern: a_directory_with_no_timestamp', {}
|
||||
)
|
||||
|
||||
|
||||
def test_cleanup_temp_directories_invalid_timestamp_format(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test that a directory with an invalid timestamp format is handled correctly."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.error = MagicMock()
|
||||
|
||||
# Create a directory with a malformed timestamp that matches the regex but fails parsing
|
||||
malformed_dir_name = 'dir_20239999_999999_999999'
|
||||
os.makedirs(os.path.join(temp_dir, malformed_dir_name))
|
||||
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
cleanup.error.assert_called_once()
|
||||
|
||||
|
||||
def test_cleanup_temp_directories_generic_exception(
|
||||
temp_dir,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test that a generic exception during directory cleanup is handled."""
|
||||
import model_manager.activities.cleanup
|
||||
|
||||
reload(model_manager.activities.cleanup)
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
|
||||
cleanup = Cleanup(
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.send_notification = MagicMock()
|
||||
|
||||
with patch('os.listdir', side_effect=Exception('Unexpected OS Error')):
|
||||
with pytest.raises(Exception, match='Unexpected OS Error'):
|
||||
cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
|
||||
|
||||
cleanup.send_notification.assert_called_once()
|
||||
643
tests/activities/test_experiment_tracking.py
Normal file
643
tests/activities/test_experiment_tracking.py
Normal file
@@ -0,0 +1,643 @@
|
||||
"""Unit tests for ExperimentTracking class with 100% coverage."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Create a mock logger."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_notification_handler():
|
||||
"""Create a mock notification handler."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_metrics_controller():
|
||||
"""Create a mock metrics controller."""
|
||||
controller = MagicMock()
|
||||
controller.shutdown = AsyncMock()
|
||||
controller.emit = AsyncMock()
|
||||
return controller
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db_config():
|
||||
"""Create a valid database configuration."""
|
||||
return {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
'user': 'testuser',
|
||||
'password': 'testpass',
|
||||
'dbname': 'testdb',
|
||||
'min_connections': 1,
|
||||
'max_connections': 10,
|
||||
}
|
||||
|
||||
|
||||
def test_experiment_tracking_init(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test ExperimentTracking initialization."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
|
||||
def test_experiment_tracking_del_without_engine(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test __del__ when engine attribute does not exist."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
if hasattr(et, 'engine'):
|
||||
delattr(et, 'engine')
|
||||
|
||||
et.__del__()
|
||||
|
||||
|
||||
def test_experiment_tracking_del_with_engine(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test __del__ when engine exists."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.engine = MagicMock()
|
||||
|
||||
class MockSuper:
|
||||
def __del__(self):
|
||||
pass
|
||||
|
||||
with patch('builtins.super', return_value=MockSuper()):
|
||||
et.__del__()
|
||||
|
||||
|
||||
def test_experiment_tracking_del_with_engine_exception(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test __del__ catches exceptions."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.engine = MagicMock()
|
||||
|
||||
class MockSuperWithError:
|
||||
_should_raise: bool
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._should_raise = False
|
||||
|
||||
def __del__(self):
|
||||
# Only raise error if not being cleaned up by garbage collector
|
||||
# This prevents the PytestUnraisableExceptionWarning
|
||||
if hasattr(self, '_should_raise') and self._should_raise:
|
||||
raise RuntimeError('Test error')
|
||||
|
||||
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||
|
||||
mock_super = MockSuperWithError()
|
||||
mock_super._should_raise = True
|
||||
try:
|
||||
with patch('builtins.super', return_value=mock_super):
|
||||
et.__del__()
|
||||
finally:
|
||||
# Prevent the exception from being raised during garbage collection
|
||||
mock_super._should_raise = False
|
||||
|
||||
|
||||
def test_execute_update_success(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test _execute_update executes query successfully."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_connection = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.rowcount = 1
|
||||
mock_connection.execute.return_value = mock_result
|
||||
mock_engine = MagicMock()
|
||||
mock_engine.begin.return_value.__enter__.return_value = mock_connection
|
||||
et.engine = mock_engine
|
||||
|
||||
result = et._execute_update('UPDATE test SET x = :x', {'x': 1})
|
||||
|
||||
assert result == {'rowcount': 1}
|
||||
mock_connection.execute.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_success(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with STATUS update type."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
et._execute_update = mock_execute_update # type: ignore[method-assign]
|
||||
et.info = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS,
|
||||
'status': 'running',
|
||||
}
|
||||
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
call_args = mock_execute.call_args
|
||||
assert 'status' in call_args[0][1]
|
||||
assert call_args[0][1]['status'] == 'running'
|
||||
assert call_args[0][1]['experiment_run_id'] == 1
|
||||
et.info.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_missing_status(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with STATUS but missing status parameter."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS,
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_success(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with STATUS_WITH_ERROR update type."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
et._execute_update = mock_execute_update # type: ignore[method-assign]
|
||||
et.info = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'failed',
|
||||
'error_message': 'Test error',
|
||||
}
|
||||
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
call_args = mock_execute.call_args
|
||||
assert 'status' in call_args[0][1]
|
||||
assert call_args[0][1]['status'] == 'failed'
|
||||
assert call_args[0][1]['error_message'] == 'Test error'
|
||||
et.info.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_truncate_message(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run truncates error message if too long."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
et._execute_update = mock_execute_update # type: ignore[method-assign]
|
||||
et.info = MagicMock()
|
||||
|
||||
long_error = 'x' * 2000
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'failed',
|
||||
'error_message': long_error,
|
||||
}
|
||||
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
call_args = mock_execute.call_args
|
||||
assert len(call_args[0][1]['error_message']) == 1024
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_missing_error_message(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with STATUS_WITH_ERROR but missing error_message."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'failed',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_model_saved_success(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with MODEL_SAVED update type."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
et._execute_update = mock_execute_update # type: ignore[method-assign]
|
||||
et.info = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.MODEL_SAVED,
|
||||
'status': 'completed',
|
||||
'run_name': 'run_001',
|
||||
}
|
||||
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
call_args = mock_execute.call_args
|
||||
assert 'run_name' in call_args[0][1]
|
||||
assert call_args[0][1]['run_name'] == 'run_001'
|
||||
assert call_args[0][1]['status'] == 'completed'
|
||||
et.info.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_model_saved_missing_run_name(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with MODEL_SAVED but missing run_name."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.MODEL_SAVED,
|
||||
'status': 'completed',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_invalid_update_type(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with invalid update_type."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': 'invalid_type',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_no_rows_updated(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run raises error when no rows are updated."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
return {'rowcount': 0}
|
||||
|
||||
et._execute_update = mock_execute_update # type: ignore[method-assign]
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 999,
|
||||
'update_type': UpdateType.STATUS,
|
||||
'status': 'running',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_missing_status(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with STATUS_WITH_ERROR but missing status - covers line 179."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'error_message': 'Some error',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_model_saved_missing_status(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test update_experiment_run with MODEL_SAVED but missing status - covers line 204."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 1,
|
||||
'update_type': UpdateType.MODEL_SAVED,
|
||||
'run_name': 'run_001',
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_experiment_tracking_del_with_engine_no_super_del(
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test __del__ when engine exists but super has no __del__ - covers line 103."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
|
||||
et = ExperimentTracking(
|
||||
host=db_config['host'],
|
||||
port=db_config['port'],
|
||||
user=db_config['user'],
|
||||
password=db_config['password'],
|
||||
dbname=db_config['dbname'],
|
||||
min_connections=db_config['min_connections'],
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.engine = MagicMock()
|
||||
|
||||
class MockSuperNoDel:
|
||||
pass
|
||||
|
||||
with patch('builtins.super', return_value=MockSuperNoDel()):
|
||||
et.__del__()
|
||||
630
tests/activities/test_training.py
Normal file
630
tests/activities/test_training.py
Normal file
@@ -0,0 +1,630 @@
|
||||
"""Unit tests for Training activities."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
from model_manager.utils.models.train_model_result import TrainModelResult
|
||||
|
||||
|
||||
def _minimal_params_dict():
|
||||
return {
|
||||
'variable_columns': ['a'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
'file_name': 'f.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def training():
|
||||
from model_manager.activities.training import Training
|
||||
|
||||
return Training(
|
||||
mlflow_repository=MagicMock(),
|
||||
plugin_store=MagicMock(),
|
||||
minio_repository=MagicMock(),
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
def test_load_model_metadata_success(training):
|
||||
training.plugin_store.get_model_index = MagicMock(
|
||||
return_value={'schemas': {'components': {'schemas': {}}}}
|
||||
)
|
||||
inp = {**_minimal_params_dict(), 'metadata': {'w': '1'}}
|
||||
out = training.load_model_metadata(inp)
|
||||
assert 'model_metadata' in out
|
||||
assert out['model_metadata']['schemas']
|
||||
|
||||
|
||||
def test_load_model_metadata_notifies_on_error(training):
|
||||
training.plugin_store.get_model_index = MagicMock(side_effect=RuntimeError('idx'))
|
||||
training.send_notification = MagicMock()
|
||||
inp = {**_minimal_params_dict(), 'metadata': {}}
|
||||
with pytest.raises(RuntimeError, match='idx'):
|
||||
training.load_model_metadata(inp)
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_validate_train_params_success(training):
|
||||
pdict = _minimal_params_dict()
|
||||
pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}}
|
||||
inp = {**pdict, 'metadata': {}}
|
||||
out = training.validate_train_params(inp)
|
||||
assert isinstance(out, dict)
|
||||
assert out['target_variable'] == 't'
|
||||
|
||||
|
||||
def test_validate_train_params_notifies(training):
|
||||
training.send_notification = MagicMock()
|
||||
inp = {'metadata': {}, 'experiment_run_id': 1}
|
||||
with pytest.raises((KeyError, ValueError, TypeError)):
|
||||
training.validate_train_params(inp)
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_train_model_download_fails_notifies(training):
|
||||
"""train_model notifies and re-raises when MinIO download fails."""
|
||||
tp = TrainModelParams.from_dict(
|
||||
{
|
||||
**_minimal_params_dict(),
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
training.minio_repository.download_file = MagicMock(side_effect=OSError('minio'))
|
||||
training.send_notification = MagicMock()
|
||||
with pytest.raises(OSError, match='minio'):
|
||||
training.train_model({'metadata': {'pod': 'x'}, 'train_params': tp.to_dict()})
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
def test_cleanup_resources(training):
|
||||
training.data_manager_repository.cleanup_run_directory = MagicMock()
|
||||
training.cleanup_resources({'metadata': {}, 'run_dir': '/tmp/x'})
|
||||
training.data_manager_repository.cleanup_run_directory.assert_called_once_with('/tmp/x', {})
|
||||
|
||||
|
||||
def test_cleanup_resources_notifies_on_error(training):
|
||||
training.data_manager_repository.cleanup_run_directory = MagicMock(
|
||||
side_effect=RuntimeError('rm')
|
||||
)
|
||||
training.send_notification = MagicMock()
|
||||
with pytest.raises(RuntimeError, match='rm'):
|
||||
training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'})
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_success_serializes_result(mock_mlflow, training):
|
||||
"""Exercise train_model happy path with mocks (MinIO, plugin wrapper, MLflow)."""
|
||||
tp = TrainModelParams.from_dict(
|
||||
{
|
||||
**_minimal_params_dict(),
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
|
||||
def _set_metrics(x, _w, **_kw):
|
||||
x.mse_val = 0.1
|
||||
x.mae_val = 0.2
|
||||
x.r2_val = 0.9
|
||||
return x
|
||||
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=_set_metrics
|
||||
)
|
||||
|
||||
def _fill_report(x, **_kw):
|
||||
x.report_path = '/tmp/report.html'
|
||||
x.train_data_path = '/tmp/train.csv'
|
||||
x.test_data_path = '/tmp/test.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
||||
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
pred_train = pd.DataFrame({'p': [1.0, 2.0]})
|
||||
pred_val = pd.DataFrame({'p': [1.0]})
|
||||
wrapper.predict = MagicMock(side_effect=[(pred_train, None), (pred_val, None)])
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'run-n'
|
||||
info.run_id = 'run-i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
|
||||
out = training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp.to_dict()})
|
||||
assert out['run_name'] is None
|
||||
assert out['run_id'] == 'run-i'
|
||||
assert out['run_dir'] == '/tmp/run'
|
||||
mock_mlflow.log_param.assert_any_call('mse_val', 0.1)
|
||||
mock_mlflow.log_param.assert_any_call('mae_val', 0.2)
|
||||
mock_mlflow.log_param.assert_any_call('r2_val', 0.9)
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_without_logger_does_not_set_wrapper_logger(_mock_mlflow, training):
|
||||
"""Covers branch where activity logger is None."""
|
||||
training.logger = None
|
||||
tp = TrainModelParams.from_dict(
|
||||
{
|
||||
**_minimal_params_dict(),
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill_report(x, **_kw):
|
||||
x.report_path = '/tmp/report.html'
|
||||
x.train_data_path = '/tmp/train.csv'
|
||||
x.test_data_path = '/tmp/test.csv'
|
||||
x.equation_path = '/tmp/eq.json'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
info.run_id = 'i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_train_params_as_dict(mock_mlflow, training):
|
||||
"""train_params may arrive as dict and is coerced via TrainModelParams.from_dict."""
|
||||
d = {
|
||||
**_minimal_params_dict(),
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tp = TrainModelParams.from_dict(d)
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill_report2(x, **_kw):
|
||||
x.report_path = '/tmp/report.html'
|
||||
x.train_data_path = '/tmp/train.csv'
|
||||
x.test_data_path = '/tmp/test.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report2)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
info.run_id = 'i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': d})
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_downloads_validation_file_when_set(mock_mlflow, training):
|
||||
"""Second MinIO download when val_file_name is set (covers val_bytes branch)."""
|
||||
d = {
|
||||
**_minimal_params_dict(),
|
||||
'val_file_name': 'val.csv',
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
tp = TrainModelParams.from_dict(d)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
def _dl(object_name, **_kwargs):
|
||||
if object_name == tp.file_name:
|
||||
return b'train'
|
||||
if object_name == 'val.csv':
|
||||
return b'val'
|
||||
raise AssertionError(object_name)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(side_effect=_dl)
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill(x, **_kw):
|
||||
x.report_path = '/tmp/report.html'
|
||||
x.train_data_path = '/tmp/train.csv'
|
||||
x.test_data_path = '/tmp/test.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
info.run_id = 'i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
assert training.minio_repository.download_file.call_count == 2
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
def test_prepare_data_observes_lag_on_success(training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
from model_manager import metrics as mm_metrics
|
||||
|
||||
training._prepare_data(b'csv', None, tp, {})
|
||||
|
||||
training.observe_lag_sync.assert_called_once()
|
||||
call_args = training.observe_lag_sync.call_args
|
||||
assert call_args.args[1] is mm_metrics.SIENTIA_TRAINING_DATA_PREPARATION_LAG
|
||||
training.emit_metric_sync.assert_not_called()
|
||||
|
||||
|
||||
def test_prepare_data_increments_error_counter_and_still_observes_lag_on_failure(training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(
|
||||
side_effect=RuntimeError('prep-fail')
|
||||
)
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
from model_manager import metrics as mm_metrics
|
||||
|
||||
with pytest.raises(RuntimeError, match='prep-fail'):
|
||||
training._prepare_data(b'csv', None, tp, {})
|
||||
|
||||
training.observe_lag_sync.assert_called_once()
|
||||
training.emit_metric_sync.assert_called_once()
|
||||
call_args = training.emit_metric_sync.call_args
|
||||
assert call_args.kwargs['metric_object'] is mm_metrics.SIENTIA_TRAINING_DATA_PREPARATION_ERROR_COUNT_TOTAL
|
||||
|
||||
|
||||
def test_fit_model_observes_lag_on_success(training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
from model_manager import metrics as mm_metrics
|
||||
|
||||
training._fit_model(wrapper, tmr, tp, {})
|
||||
|
||||
training.observe_lag_sync.assert_called_once()
|
||||
call_args = training.observe_lag_sync.call_args
|
||||
assert call_args.args[1] is mm_metrics.SIENTIA_TRAINING_MODEL_FIT_LAG
|
||||
training.emit_metric_sync.assert_not_called()
|
||||
|
||||
|
||||
def test_fit_model_increments_error_counter_on_failure(training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
wrapper = MagicMock()
|
||||
wrapper.train = MagicMock(side_effect=RuntimeError('fit-fail'))
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
from model_manager import metrics as mm_metrics
|
||||
|
||||
with pytest.raises(RuntimeError, match='fit-fail'):
|
||||
training._fit_model(wrapper, tmr, tp, {})
|
||||
|
||||
training.observe_lag_sync.assert_called_once()
|
||||
training.emit_metric_sync.assert_called_once()
|
||||
call_args = training.emit_metric_sync.call_args
|
||||
assert call_args.kwargs['metric_object'] is mm_metrics.SIENTIA_TRAINING_MODEL_FIT_ERROR_COUNT_TOTAL
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mm_metrics')
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_sets_quality_gauges_after_compute_metrics(mock_mlflow, mock_mm_metrics, training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
|
||||
def _set_metrics(x, _w, **_kw):
|
||||
x.mse_val = 0.5
|
||||
x.mae_val = 0.3
|
||||
x.r2_val = -0.1
|
||||
return x
|
||||
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=_set_metrics
|
||||
)
|
||||
|
||||
def _fill_report(x, **_kw):
|
||||
x.report_path = '/tmp/r.html'
|
||||
x.train_data_path = '/tmp/tr.csv'
|
||||
x.test_data_path = '/tmp/te.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_id = 'rid'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_MSE.labels.return_value.set.assert_called_once_with(0.5)
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_MAE.labels.return_value.set.assert_called_once_with(0.3)
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_R2.labels.return_value.set.assert_called_once_with(-0.1)
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mm_metrics')
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_skips_quality_gauges_when_none(_mock_mlflow, mock_mm_metrics, training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill_report(x, **_kw):
|
||||
x.report_path = '/tmp/r.html'
|
||||
x.train_data_path = '/tmp/tr.csv'
|
||||
x.test_data_path = '/tmp/te.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_id = 'rid'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_MSE.labels.return_value.set.assert_not_called()
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_MAE.labels.return_value.set.assert_not_called()
|
||||
mock_mm_metrics.SIENTIA_TRAINING_MODEL_QUALITY_R2.labels.return_value.set.assert_not_called()
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.mm_metrics')
|
||||
@patch('model_manager.activities.training.mlflow')
|
||||
def test_train_model_increments_trained_total_on_success(_mock_mlflow, mock_mm_metrics, training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill_report(x, **_kw):
|
||||
x.report_path = '/tmp/r.html'
|
||||
x.train_data_path = '/tmp/tr.csv'
|
||||
x.test_data_path = '/tmp/te.csv'
|
||||
x.run_dir = '/tmp/run'
|
||||
return x
|
||||
|
||||
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_id = 'rid'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
training.observe_lag_sync = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
training.emit_metric_sync.assert_called_once_with(
|
||||
metric_object=mock_mm_metrics.SIENTIA_TRAINING_MODEL_TRAINED_TOTAL,
|
||||
tags=training._get_training_labels(tp),
|
||||
)
|
||||
|
||||
|
||||
def test_train_model_does_not_increment_trained_total_on_failure(training):
|
||||
tp = TrainModelParams.from_dict(
|
||||
{**_minimal_params_dict(), 'model_metadata': {'schemas': {'components': {'schemas': {}}}}}
|
||||
)
|
||||
training.minio_repository.download_file = MagicMock(side_effect=RuntimeError('dl-fail'))
|
||||
training.send_notification = MagicMock()
|
||||
training.emit_metric_sync = MagicMock()
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
training.emit_metric_sync.assert_not_called()
|
||||
|
||||
|
||||
def test_train_model_value_error_when_paths_missing_after_report(training):
|
||||
"""Raises ValueError when report paths are not populated after generate_report."""
|
||||
tp = TrainModelParams.from_dict(
|
||||
{
|
||||
**_minimal_params_dict(),
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
||||
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
||||
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
||||
|
||||
training.minio_repository.download_file = MagicMock(return_value=b'x')
|
||||
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
||||
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
training.data_manager_repository.generate_report = MagicMock(return_value=tmr)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
||||
)
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
info.run_id = 'i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
training.send_notification = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError, match='Report path'):
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
103
tests/conftest.py
Normal file
103
tests/conftest.py
Normal file
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
Test bootstrap: stub optional `sientia_do` submodules not shipped in minimal installs.
|
||||
|
||||
Must run before importing `model_manager.sientia.models` (pulled in via TrainModelParams).
|
||||
Stubs Evidently submodules so `model_manager.sientia.reports` imports (via DataManagerRepository).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from types import ModuleType
|
||||
|
||||
|
||||
def _make_dummy(name: str) -> type:
|
||||
return type(name, (), {})
|
||||
|
||||
|
||||
def _stub_evidently() -> None:
|
||||
"""Minimal Evidently API surface required to import `model_manager.sientia.reports`."""
|
||||
ev = ModuleType('evidently')
|
||||
sys.modules['evidently'] = ev
|
||||
|
||||
mp = ModuleType('evidently.metric_preset')
|
||||
mp.DataDriftPreset = _make_dummy('DataDriftPreset') # type: ignore[attr-defined]
|
||||
sys.modules['evidently.metric_preset'] = mp
|
||||
|
||||
metrics = ModuleType('evidently.metrics')
|
||||
_metric_names = (
|
||||
'ColumnSummaryMetric',
|
||||
'ConflictTargetMetric',
|
||||
'DatasetCorrelationsMetric',
|
||||
'DatasetSummaryMetric',
|
||||
'RegressionAbsPercentageErrorPlot',
|
||||
'RegressionDummyMetric',
|
||||
'RegressionErrorDistribution',
|
||||
'RegressionErrorPlot',
|
||||
'RegressionPerformanceMetrics',
|
||||
'RegressionPredictedVsActualPlot',
|
||||
'RegressionPredictedVsActualScatter',
|
||||
)
|
||||
for n in _metric_names:
|
||||
setattr(metrics, n, _make_dummy(n))
|
||||
sys.modules['evidently.metrics'] = metrics
|
||||
|
||||
base = ModuleType('evidently.metrics.base_metric')
|
||||
|
||||
def generate_column_metrics(*_a, **_k):
|
||||
return []
|
||||
|
||||
base.generate_column_metrics = generate_column_metrics # type: ignore[attr-defined]
|
||||
sys.modules['evidently.metrics.base_metric'] = base
|
||||
|
||||
opt = ModuleType('evidently.options')
|
||||
opt.ColorOptions = _make_dummy('ColorOptions') # type: ignore[attr-defined]
|
||||
sys.modules['evidently.options'] = opt
|
||||
|
||||
pipeline = ModuleType('evidently.pipeline')
|
||||
sys.modules['evidently.pipeline'] = pipeline
|
||||
|
||||
colmap = ModuleType('evidently.pipeline.column_mapping')
|
||||
colmap.ColumnMapping = _make_dummy('ColumnMapping') # type: ignore[attr-defined]
|
||||
sys.modules['evidently.pipeline.column_mapping'] = colmap
|
||||
|
||||
rep = ModuleType('evidently.report')
|
||||
rep.Report = _make_dummy('Report') # type: ignore[attr-defined]
|
||||
sys.modules['evidently.report'] = rep
|
||||
|
||||
|
||||
def pytest_configure(config) -> None: # noqa: ARG001
|
||||
"""Register stub modules so imports used by production code resolve in CI/dev venvs."""
|
||||
_stub_evidently()
|
||||
|
||||
if 'sientia_do.operations.df_preprocessor' not in sys.modules:
|
||||
df_pre = ModuleType('sientia_do.operations.df_preprocessor')
|
||||
|
||||
def create_features(input_data, *_a, **_k):
|
||||
return input_data
|
||||
|
||||
def limit_dataset(input_data, low_lim, upp_lim, *_a, **_k):
|
||||
return input_data, low_lim, upp_lim
|
||||
|
||||
def treat_nan(input_data, *_a, **_k):
|
||||
return input_data
|
||||
|
||||
df_pre.create_features = create_features # type: ignore[attr-defined]
|
||||
df_pre.limit_dataset = limit_dataset # type: ignore[attr-defined]
|
||||
df_pre.treat_nan = treat_nan # type: ignore[attr-defined]
|
||||
sys.modules['sientia_do.operations.df_preprocessor'] = df_pre
|
||||
|
||||
sys.modules.setdefault('sientia_do.operations', ModuleType('sientia_do.operations'))
|
||||
|
||||
if 'sientia_do.timeseries.analyzer' not in sys.modules:
|
||||
ts_an = ModuleType('sientia_do.timeseries.analyzer')
|
||||
|
||||
class TimeSeriesDiscontinuityAnalyzer: # noqa: D401
|
||||
"""Stub for tests."""
|
||||
|
||||
pass
|
||||
|
||||
ts_an.TimeSeriesDiscontinuityAnalyzer = TimeSeriesDiscontinuityAnalyzer # type: ignore[attr-defined]
|
||||
sys.modules['sientia_do.timeseries.analyzer'] = ts_an
|
||||
|
||||
sys.modules.setdefault('sientia_do.timeseries', ModuleType('sientia_do.timeseries'))
|
||||
0
tests/schedules/__init__.py
Normal file
0
tests/schedules/__init__.py
Normal file
469
tests/schedules/test_cleanup_schedule.py
Normal file
469
tests/schedules/test_cleanup_schedule.py
Normal file
@@ -0,0 +1,469 @@
|
||||
"""Tests for cleanup schedule management."""
|
||||
|
||||
import os
|
||||
from datetime import timedelta
|
||||
from importlib import reload
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_temporal_client():
|
||||
"""Fixture for a mock Temporal client."""
|
||||
client = AsyncMock()
|
||||
client.list_schedules = AsyncMock()
|
||||
client.create_schedule = AsyncMock()
|
||||
handle = AsyncMock()
|
||||
handle.delete = AsyncMock()
|
||||
schedule = MagicMock()
|
||||
schedule.action.task_queue = 'cleanup_files-model-manager-worker-queue'
|
||||
schedule.action.execution_timeout = timedelta(hours=1)
|
||||
schedule.spec.cron_expressions = ['0 0 * * *']
|
||||
schedule.spec.time_zone_name = 'UTC'
|
||||
handle.describe = AsyncMock(return_value=MagicMock(schedule=schedule))
|
||||
client.get_schedule_handle = MagicMock(return_value=handle)
|
||||
return client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Fixture for a mock Sientia logger."""
|
||||
logger = MagicMock()
|
||||
logger.custom_info = MagicMock()
|
||||
logger.custom_error = MagicMock()
|
||||
return logger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def metadata():
|
||||
"""Fixture for metadata dict."""
|
||||
return {'pod_id': 'test-pod', 'project_name': 'test-project'}
|
||||
|
||||
|
||||
# --- schedule_exists Tests ---
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_exists_returns_true_when_schedule_found(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that schedule_exists returns True when schedule is found."""
|
||||
from model_manager.schedules.cleanup_schedule import schedule_exists
|
||||
|
||||
# Mock schedule list with matching schedule
|
||||
mock_schedule = MagicMock()
|
||||
mock_schedule.id = 'test-schedule-id'
|
||||
|
||||
async def mock_list_schedules():
|
||||
yield mock_schedule
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
result = await schedule_exists(mock_temporal_client, 'test-schedule-id', mock_logger, metadata)
|
||||
|
||||
assert result is True
|
||||
mock_temporal_client.list_schedules.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_exists_returns_false_when_schedule_not_found(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that schedule_exists returns False when schedule is not found."""
|
||||
from model_manager.schedules.cleanup_schedule import schedule_exists
|
||||
|
||||
# Mock empty schedule list
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield # Make it an async generator
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
result = await schedule_exists(
|
||||
mock_temporal_client, 'nonexistent-schedule', mock_logger, metadata
|
||||
)
|
||||
|
||||
assert result is False
|
||||
mock_temporal_client.list_schedules.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_exists_returns_false_when_different_schedule_found(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that schedule_exists returns False when only different schedules exist."""
|
||||
from model_manager.schedules.cleanup_schedule import schedule_exists
|
||||
|
||||
# Mock schedule list with non-matching schedule
|
||||
mock_schedule = MagicMock()
|
||||
mock_schedule.id = 'different-schedule-id'
|
||||
|
||||
async def mock_list_schedules():
|
||||
yield mock_schedule
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
result = await schedule_exists(mock_temporal_client, 'test-schedule-id', mock_logger, metadata)
|
||||
|
||||
assert result is False
|
||||
mock_temporal_client.list_schedules.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_exists_handles_exception(mock_temporal_client, mock_logger, metadata):
|
||||
"""Test that schedule_exists handles exceptions gracefully."""
|
||||
from model_manager.schedules.cleanup_schedule import schedule_exists
|
||||
|
||||
# Mock list_schedules to raise an exception
|
||||
mock_temporal_client.list_schedules.side_effect = Exception('Connection error')
|
||||
|
||||
result = await schedule_exists(mock_temporal_client, 'test-schedule-id', mock_logger, metadata)
|
||||
|
||||
assert result is False
|
||||
mock_logger.custom_error.assert_called_once()
|
||||
assert 'Error checking if schedule exists' in mock_logger.custom_error.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_schedule_reconcile_handles_describe_exception(mock_logger, metadata):
|
||||
"""Test _needs_schedule_reconcile returns True and logs when describe fails."""
|
||||
from model_manager.schedules.cleanup_schedule import _needs_schedule_reconcile
|
||||
|
||||
handle = AsyncMock()
|
||||
handle.describe = AsyncMock(side_effect=RuntimeError('describe failed'))
|
||||
|
||||
needs_reconcile = await _needs_schedule_reconcile(
|
||||
schedule_handle=handle,
|
||||
cleanup_task_queue='cleanup_files-model-manager-worker-queue',
|
||||
logger=mock_logger,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert needs_reconcile is True
|
||||
mock_logger.custom_error.assert_called_once()
|
||||
assert (
|
||||
'Error describing cleanup schedule for reconcile'
|
||||
in mock_logger.custom_error.call_args[0][0]
|
||||
)
|
||||
|
||||
|
||||
# --- create_cleanup_schedule Tests ---
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_reconciles_when_exists(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that create_cleanup_schedule recreates schedule when it already exists."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
# Mock schedule already exists
|
||||
mock_schedule = MagicMock()
|
||||
mock_schedule.id = 'cleanup-files-model-manager-worker-daily'
|
||||
|
||||
async def mock_list_schedules():
|
||||
yield mock_schedule
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
# Force reconcile by diverging task queue
|
||||
mock_temporal_client.get_schedule_handle.return_value.describe.return_value.schedule.action.task_queue = 'different-queue'
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
# Verify schedule was reconciled via delete + create
|
||||
mock_temporal_client.get_schedule_handle.assert_called_once_with(
|
||||
'cleanup-files-model-manager-worker-daily'
|
||||
)
|
||||
mock_temporal_client.get_schedule_handle.return_value.delete.assert_called_once()
|
||||
mock_temporal_client.create_schedule.assert_called_once()
|
||||
|
||||
mock_logger.custom_info.assert_called_once()
|
||||
assert 'reconciled successfully' in mock_logger.custom_info.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_noop_when_schedule_is_up_to_date(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test no-op reconcile when existing schedule already matches current config."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
mock_schedule = MagicMock()
|
||||
mock_schedule.id = 'cleanup-files-model-manager-worker-daily'
|
||||
|
||||
async def mock_list_schedules():
|
||||
yield mock_schedule
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
mock_temporal_client.get_schedule_handle.assert_called_once_with(
|
||||
'cleanup-files-model-manager-worker-daily'
|
||||
)
|
||||
mock_temporal_client.get_schedule_handle.return_value.delete.assert_not_called()
|
||||
mock_temporal_client.create_schedule.assert_not_called()
|
||||
mock_logger.custom_info.assert_called_once()
|
||||
assert 'no-op reconcile' in mock_logger.custom_info.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'CLEANUP_CRON': '0 2 * * *',
|
||||
'CLEANUP_TIMEZONE': 'America/Sao_Paulo',
|
||||
'CLEANUP_EXECUTION_TIMEOUT_HOURS': '2',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_creates_with_custom_config(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that create_cleanup_schedule creates schedule with custom configuration."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
# Mock schedule does not exist (empty list)
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield # Make it an async generator
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
# Verify schedule creation was called
|
||||
mock_temporal_client.create_schedule.assert_called_once()
|
||||
|
||||
# Verify schedule parameters
|
||||
call_args = mock_temporal_client.create_schedule.call_args
|
||||
schedule_id = call_args[0][0]
|
||||
schedule_obj = call_args[0][1]
|
||||
|
||||
assert schedule_id == 'cleanup-files-model-manager-worker-daily'
|
||||
assert schedule_obj.action.workflow == 'cleanup_files'
|
||||
assert schedule_obj.action.task_queue == 'cleanup_files-model-manager-worker-queue'
|
||||
assert schedule_obj.action.execution_timeout == timedelta(hours=2)
|
||||
assert schedule_obj.spec.cron_expressions == ['0 2 * * *']
|
||||
assert schedule_obj.spec.time_zone_name == 'America/Sao_Paulo'
|
||||
|
||||
# Verify success log was called
|
||||
assert mock_logger.custom_info.call_count == 1
|
||||
assert 'created successfully' in mock_logger.custom_info.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_uses_defaults(mock_temporal_client, mock_logger, metadata):
|
||||
"""Test that create_cleanup_schedule uses default values when env vars not set."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
# Remove optional env vars to test defaults
|
||||
for key in [
|
||||
'CLEANUP_CRON',
|
||||
'CLEANUP_TIMEZONE',
|
||||
'CLEANUP_EXECUTION_TIMEOUT_HOURS',
|
||||
]:
|
||||
os.environ.pop(key, None)
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
# Mock schedule does not exist (empty list)
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield # Make it an async generator
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
# Verify schedule creation was called
|
||||
mock_temporal_client.create_schedule.assert_called_once()
|
||||
|
||||
# Verify default parameters
|
||||
call_args = mock_temporal_client.create_schedule.call_args
|
||||
schedule_obj = call_args[0][1]
|
||||
|
||||
assert schedule_obj.spec.cron_expressions == ['0 0 * * *'] # Default midnight
|
||||
assert schedule_obj.spec.time_zone_name == 'UTC' # Default UTC
|
||||
assert schedule_obj.action.task_queue == 'cleanup_files-model-manager-worker-queue'
|
||||
assert schedule_obj.action.execution_timeout == timedelta(hours=1) # Default 1 hour
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_workflow_id_format(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that workflow ID is correctly formatted with schedule ID."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
# Mock schedule does not exist (empty list)
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield # Make it an async generator
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
# Verify workflow ID format
|
||||
call_args = mock_temporal_client.create_schedule.call_args
|
||||
schedule_obj = call_args[0][1]
|
||||
|
||||
expected_workflow_id = 'cleanup-files-scheduled-cleanup-files-model-manager-worker-daily'
|
||||
assert schedule_obj.action.id == expected_workflow_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
},
|
||||
)
|
||||
async def test_create_cleanup_schedule_empty_workflow_args(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test that workflow is created with empty args (uses env defaults)."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
# Mock schedule does not exist (empty list)
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield # Make it an async generator
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
# Verify workflow args are empty (it's a list with one empty dict)
|
||||
call_args = mock_temporal_client.create_schedule.call_args
|
||||
schedule_obj = call_args[0][1]
|
||||
|
||||
# The args are passed as positional args, so it's a list with one element
|
||||
assert schedule_obj.action.args == [{}]
|
||||
|
||||
|
||||
# --- Environment Variable Configuration Tests ---
|
||||
|
||||
|
||||
@patch.dict(
|
||||
'model_manager.schedules.cleanup_schedule.os.environ',
|
||||
{
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'CLEANUP_CRON': '30 3 * * 1',
|
||||
'CLEANUP_TIMEZONE': 'Europe/London',
|
||||
'CLEANUP_EXECUTION_TIMEOUT_HOURS': '3',
|
||||
},
|
||||
)
|
||||
def test_environment_variables_loaded_correctly():
|
||||
"""Test that environment variables are loaded correctly."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import (
|
||||
CLEANUP_CRON,
|
||||
CLEANUP_EXECUTION_TIMEOUT_HOURS,
|
||||
CLEANUP_TIMEZONE,
|
||||
build_cleanup_schedule_id,
|
||||
)
|
||||
|
||||
assert (
|
||||
build_cleanup_schedule_id('model-manager-worker')
|
||||
== 'cleanup-files-model-manager-worker-daily'
|
||||
)
|
||||
assert CLEANUP_CRON == '30 3 * * 1'
|
||||
assert CLEANUP_TIMEZONE == 'Europe/London'
|
||||
assert CLEANUP_EXECUTION_TIMEOUT_HOURS == 3
|
||||
|
||||
|
||||
def test_environment_variables_use_defaults_when_not_set():
|
||||
"""Test that default values are used when environment variables are not set."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
# Remove all env vars
|
||||
for key in [
|
||||
'RUNTIME',
|
||||
'CLEANUP_CRON',
|
||||
'CLEANUP_TIMEZONE',
|
||||
'CLEANUP_EXECUTION_TIMEOUT_HOURS',
|
||||
]:
|
||||
os.environ.pop(key, None)
|
||||
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import (
|
||||
CLEANUP_CRON,
|
||||
CLEANUP_EXECUTION_TIMEOUT_HOURS,
|
||||
CLEANUP_TIMEZONE,
|
||||
build_cleanup_schedule_id,
|
||||
)
|
||||
|
||||
assert build_cleanup_schedule_id(None) == 'cleanup-files-single-daily'
|
||||
assert CLEANUP_CRON == '0 0 * * *'
|
||||
assert CLEANUP_TIMEZONE == 'UTC'
|
||||
assert CLEANUP_EXECUTION_TIMEOUT_HOURS == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_cleanup_schedule_uses_single_runtime_when_runtime_missing(
|
||||
mock_temporal_client, mock_logger, metadata
|
||||
):
|
||||
"""Test create_cleanup_schedule uses single runtime fallback."""
|
||||
import model_manager.schedules.cleanup_schedule
|
||||
|
||||
os.environ.pop('RUNTIME', None)
|
||||
reload(model_manager.schedules.cleanup_schedule)
|
||||
from model_manager.schedules.cleanup_schedule import create_cleanup_schedule
|
||||
|
||||
async def mock_list_schedules():
|
||||
return
|
||||
yield
|
||||
|
||||
mock_temporal_client.list_schedules.return_value = mock_list_schedules()
|
||||
|
||||
await create_cleanup_schedule(mock_temporal_client, mock_logger, metadata)
|
||||
|
||||
call_args = mock_temporal_client.create_schedule.call_args
|
||||
schedule_obj = call_args[0][1]
|
||||
assert schedule_obj.action.task_queue == 'cleanup_files-single-queue'
|
||||
0
tests/sientia/__init__.py
Normal file
0
tests/sientia/__init__.py
Normal file
9
tests/sientia/test_exceptions.py
Normal file
9
tests/sientia/test_exceptions.py
Normal file
@@ -0,0 +1,9 @@
|
||||
"""Unit tests for custom exception aliases."""
|
||||
|
||||
from mlflow.exceptions import MlflowException
|
||||
|
||||
from model_manager.sientia.exceptions import SientiaMlException
|
||||
|
||||
|
||||
def test_sientia_ml_exception_is_mlflow_exception_alias():
|
||||
assert SientiaMlException is MlflowException
|
||||
481
tests/sientia/test_metrics.py
Normal file
481
tests/sientia/test_metrics.py
Normal file
@@ -0,0 +1,481 @@
|
||||
"""Unit tests for sientia metrics module."""
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from model_manager.sientia.metrics import (
|
||||
mae,
|
||||
mse,
|
||||
r2,
|
||||
rce_drift,
|
||||
rce_test,
|
||||
rce_train,
|
||||
silverman_radius,
|
||||
)
|
||||
|
||||
|
||||
def test_mse_perfect_predictions():
|
||||
"""Test MSE with perfect predictions returns 0.0."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
assert result == 0.0
|
||||
|
||||
|
||||
def test_mse_with_errors():
|
||||
"""Test MSE calculation with prediction errors."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.5, 2.5, 3.5, 4.5, 5.5])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
# MSE = mean((0.5^2, 0.5^2, 0.5^2, 0.5^2, 0.5^2)) = 0.25
|
||||
assert result == 0.25
|
||||
|
||||
|
||||
def test_mse_with_integer_input():
|
||||
"""Test MSE handles integer input and converts to float64."""
|
||||
real_data = pd.Series([1, 2, 3, 4, 5])
|
||||
predictions = pd.Series([2, 3, 4, 5, 6])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
# MSE = mean((1^2, 1^2, 1^2, 1^2, 1^2)) = 1.0
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_mse_with_large_errors():
|
||||
"""Test MSE with large prediction errors."""
|
||||
real_data = pd.Series([10.0, 20.0, 30.0])
|
||||
predictions = pd.Series([5.0, 15.0, 25.0])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
# MSE = mean((25, 25, 25)) = 25.0
|
||||
assert result == 25.0
|
||||
|
||||
|
||||
def test_mse_rounds_to_two_decimals():
|
||||
"""Test MSE rounds result to 2 decimal places."""
|
||||
real_data = pd.Series([1.111, 2.222, 3.333])
|
||||
predictions = pd.Series([1.222, 2.333, 3.444])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
# Result should be rounded to 2 decimals
|
||||
assert isinstance(result, float)
|
||||
assert len(str(result).split('.')[-1]) <= 2
|
||||
|
||||
|
||||
def test_mae_perfect_predictions():
|
||||
"""Test MAE with perfect predictions returns 0.0."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
assert result == 0.0
|
||||
|
||||
|
||||
def test_mae_with_errors():
|
||||
"""Test MAE calculation with prediction errors."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.5, 2.5, 3.5, 4.5, 5.5])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
# MAE = mean(|0.5|, |0.5|, |0.5|, |0.5|, |0.5|) = 0.5
|
||||
assert result == 0.5
|
||||
|
||||
|
||||
def test_mae_with_integer_input():
|
||||
"""Test MAE handles integer input and converts to float64."""
|
||||
real_data = pd.Series([1, 2, 3, 4, 5])
|
||||
predictions = pd.Series([2, 3, 4, 5, 6])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
# MAE = mean(|1|, |1|, |1|, |1|, |1|) = 1.0
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_mae_with_negative_errors():
|
||||
"""Test MAE with negative prediction errors (absolute value)."""
|
||||
real_data = pd.Series([10.0, 20.0, 30.0])
|
||||
predictions = pd.Series([15.0, 25.0, 35.0])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
# MAE = mean(|5|, |5|, |5|) = 5.0
|
||||
assert result == 5.0
|
||||
|
||||
|
||||
def test_mae_rounds_to_two_decimals():
|
||||
"""Test MAE rounds result to 2 decimal places."""
|
||||
real_data = pd.Series([1.111, 2.222, 3.333])
|
||||
predictions = pd.Series([1.222, 2.333, 3.444])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
# Result should be rounded to 2 decimals
|
||||
assert isinstance(result, float)
|
||||
assert len(str(result).split('.')[-1]) <= 2
|
||||
|
||||
|
||||
def test_r2_perfect_predictions():
|
||||
"""Test R2 with perfect predictions returns 1.0."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_r2_with_good_predictions():
|
||||
"""Test R2 calculation with good predictions."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([1.1, 2.1, 2.9, 4.1, 4.9])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
# R2 should be close to 1.0 for good predictions
|
||||
assert result > 0.9
|
||||
assert result <= 1.0
|
||||
|
||||
|
||||
def test_r2_with_integer_input():
|
||||
"""Test R2 handles integer input and converts to float64."""
|
||||
real_data = pd.Series([1, 2, 3, 4, 5])
|
||||
predictions = pd.Series([1, 2, 3, 4, 5])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_r2_with_poor_predictions():
|
||||
"""Test R2 with poor predictions returns low score."""
|
||||
real_data = pd.Series([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
predictions = pd.Series([5.0, 4.0, 3.0, 2.0, 1.0])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
# R2 should be negative for predictions worse than mean
|
||||
assert result < 0
|
||||
|
||||
|
||||
def test_r2_rounds_to_two_decimals():
|
||||
"""Test R2 rounds result to 2 decimal places."""
|
||||
real_data = pd.Series([1.111, 2.222, 3.333, 4.444, 5.555])
|
||||
predictions = pd.Series([1.222, 2.333, 3.444, 4.555, 5.666])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
# Result should be rounded to 2 decimals
|
||||
assert isinstance(result, float)
|
||||
assert len(str(result).split('.')[-1]) <= 2
|
||||
|
||||
|
||||
def test_mse_with_mixed_positive_negative():
|
||||
"""Test MSE with mixed positive and negative values."""
|
||||
real_data = pd.Series([-5.0, -2.0, 0.0, 3.0, 7.0])
|
||||
predictions = pd.Series([-4.0, -1.0, 1.0, 4.0, 8.0])
|
||||
|
||||
result = mse(real_data, predictions)
|
||||
|
||||
# MSE = mean((1^2, 1^2, 1^2, 1^2, 1^2)) = 1.0
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_mae_with_mixed_positive_negative():
|
||||
"""Test MAE with mixed positive and negative values."""
|
||||
real_data = pd.Series([-5.0, -2.0, 0.0, 3.0, 7.0])
|
||||
predictions = pd.Series([-4.0, -1.0, 1.0, 4.0, 8.0])
|
||||
|
||||
result = mae(real_data, predictions)
|
||||
|
||||
# MAE = mean(|1|, |1|, |1|, |1|, |1|) = 1.0
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
def test_r2_with_mixed_positive_negative():
|
||||
"""Test R2 with mixed positive and negative values."""
|
||||
real_data = pd.Series([-5.0, -2.0, 0.0, 3.0, 7.0])
|
||||
predictions = pd.Series([-5.0, -2.0, 0.0, 3.0, 7.0])
|
||||
|
||||
result = r2(real_data, predictions)
|
||||
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for silverman_radius
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_silverman_radius_basic():
|
||||
"""Test silverman_radius returns a positive float."""
|
||||
data = np.array([1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0])
|
||||
|
||||
result = silverman_radius(data)
|
||||
|
||||
assert isinstance(result, float)
|
||||
assert result > 0
|
||||
|
||||
|
||||
def test_silverman_radius_uniform_data():
|
||||
"""Test silverman_radius with uniformly distributed data."""
|
||||
data = np.linspace(0, 100, 50)
|
||||
|
||||
result = silverman_radius(data)
|
||||
|
||||
assert result > 0
|
||||
assert np.isfinite(result)
|
||||
|
||||
|
||||
def test_silverman_radius_normal_distribution():
|
||||
"""Test silverman_radius with normally distributed data."""
|
||||
np.random.seed(42)
|
||||
data = np.random.normal(loc=50, scale=10, size=100)
|
||||
|
||||
result = silverman_radius(data)
|
||||
|
||||
assert result > 0
|
||||
assert np.isfinite(result)
|
||||
|
||||
|
||||
def test_silverman_radius_small_dataset():
|
||||
"""Test silverman_radius with small dataset."""
|
||||
data = np.array([1.0, 2.0, 3.0])
|
||||
|
||||
result = silverman_radius(data)
|
||||
|
||||
assert result > 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for rce_train
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_rce_train_returns_dataframe():
|
||||
"""Test rce_train returns a DataFrame."""
|
||||
training_set = pd.DataFrame({'a': [1.0, 2.0, 3.0, 4.0, 5.0], 'b': [2.0, 3.0, 4.0, 5.0, 6.0]})
|
||||
|
||||
result = rce_train(training_set, 0.1)
|
||||
|
||||
assert isinstance(result, pd.DataFrame)
|
||||
|
||||
|
||||
def test_rce_train_includes_first_vector():
|
||||
"""Test rce_train always includes the first vector as a prototype."""
|
||||
training_set = pd.DataFrame({'a': [1.0, 2.0, 3.0], 'b': [1.0, 2.0, 3.0]})
|
||||
|
||||
result = rce_train(training_set, 0.1)
|
||||
|
||||
assert len(result) >= 1
|
||||
assert result.iloc[0].tolist() == [1.0, 1.0]
|
||||
|
||||
|
||||
def test_rce_train_with_identical_vectors():
|
||||
"""Test rce_train with identical vectors returns single prototype."""
|
||||
training_set = pd.DataFrame({'a': [1.0, 1.0, 1.0], 'b': [2.0, 2.0, 2.0]})
|
||||
|
||||
result = rce_train(training_set, 0.1)
|
||||
|
||||
# All vectors are identical, so only one prototype should be created
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
def test_rce_train_with_distant_vectors():
|
||||
"""Test rce_train with very distant vectors creates multiple prototypes."""
|
||||
training_set = pd.DataFrame({'a': [0.0, 100.0, 200.0], 'b': [0.0, 100.0, 200.0]})
|
||||
|
||||
result = rce_train(training_set, 0.1)
|
||||
|
||||
# Distant vectors should create multiple prototypes
|
||||
assert len(result) >= 1
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for rce_test
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_rce_test_returns_series():
|
||||
"""Test rce_test returns a pandas Series."""
|
||||
test_set = pd.DataFrame({'a': [1.5, 2.5], 'b': [1.5, 2.5]})
|
||||
prototypes = pd.DataFrame({'a': [1.0, 3.0], 'b': [1.0, 3.0]})
|
||||
|
||||
result = rce_test(test_set, prototypes)
|
||||
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == len(test_set)
|
||||
|
||||
|
||||
def test_rce_test_with_exact_match():
|
||||
"""Test rce_test with test vector matching a prototype."""
|
||||
test_set = pd.DataFrame({'a': [1.0], 'b': [2.0]})
|
||||
prototypes = pd.DataFrame({'a': [1.0], 'b': [2.0]})
|
||||
|
||||
result = rce_test(test_set, prototypes)
|
||||
|
||||
# Distance should be 0 for exact match
|
||||
assert result.iloc[0] == 0.0
|
||||
|
||||
|
||||
def test_rce_test_multiple_prototypes():
|
||||
"""Test rce_test finds closest prototype."""
|
||||
test_set = pd.DataFrame({'a': [1.1], 'b': [1.1]})
|
||||
prototypes = pd.DataFrame({'a': [1.0, 10.0], 'b': [1.0, 10.0]})
|
||||
|
||||
result = rce_test(test_set, prototypes)
|
||||
|
||||
# Should find the closest prototype (1.0, 1.0)
|
||||
assert len(result) == 1
|
||||
assert np.isfinite(result.iloc[0])
|
||||
|
||||
|
||||
def test_rce_test_signed_distances():
|
||||
"""Test rce_test returns signed distances."""
|
||||
test_set = pd.DataFrame({'a': [0.0, 5.0], 'b': [0.0, 5.0]})
|
||||
prototypes = pd.DataFrame({'a': [2.0], 'b': [2.0]})
|
||||
|
||||
result = rce_test(test_set, prototypes)
|
||||
|
||||
assert len(result) == 2
|
||||
# First test vector (0,0) is less than prototype (2,2) - should be negative
|
||||
# Second test vector (5,5) is greater than prototype (2,2) - should be positive
|
||||
assert result.iloc[0] < 0
|
||||
assert result.iloc[1] > 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for rce_drift
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_rce_drift_returns_series():
|
||||
"""Test rce_drift returns a pandas Series."""
|
||||
reference_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.0, 2.0, 3.0, 4.0, 5.0],
|
||||
'feature2': [2.0, 3.0, 4.0, 5.0, 6.0],
|
||||
'target': [10.0, 20.0, 30.0, 40.0, 50.0],
|
||||
'prediction': [11.0, 21.0, 31.0, 41.0, 51.0],
|
||||
}
|
||||
)
|
||||
real_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.5, 2.5],
|
||||
'feature2': [2.5, 3.5],
|
||||
'target': [15.0, 25.0],
|
||||
'prediction': [16.0, 26.0],
|
||||
}
|
||||
)
|
||||
|
||||
result = rce_drift(reference_data, real_data, 'target')
|
||||
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == len(real_data)
|
||||
|
||||
|
||||
def test_rce_drift_with_target_column():
|
||||
"""Test rce_drift using target column (drops prediction)."""
|
||||
reference_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.0, 2.0, 3.0],
|
||||
'target': [10.0, 20.0, 30.0],
|
||||
'prediction': [11.0, 21.0, 31.0],
|
||||
}
|
||||
)
|
||||
real_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.5],
|
||||
'target': [15.0],
|
||||
'prediction': [16.0],
|
||||
}
|
||||
)
|
||||
|
||||
result = rce_drift(reference_data, real_data, 'target')
|
||||
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
def test_rce_drift_with_prediction_column():
|
||||
"""Test rce_drift using prediction column (drops target)."""
|
||||
reference_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.0, 2.0, 3.0],
|
||||
'target': [10.0, 20.0, 30.0],
|
||||
'prediction': [11.0, 21.0, 31.0],
|
||||
}
|
||||
)
|
||||
real_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.5],
|
||||
'target': [15.0],
|
||||
'prediction': [16.0],
|
||||
}
|
||||
)
|
||||
|
||||
result = rce_drift(reference_data, real_data, 'prediction')
|
||||
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
def test_rce_drift_normalized_output():
|
||||
"""Test rce_drift returns normalized distances."""
|
||||
reference_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.0, 2.0, 3.0, 4.0, 5.0],
|
||||
'target': [10.0, 20.0, 30.0, 40.0, 50.0],
|
||||
'prediction': [10.0, 20.0, 30.0, 40.0, 50.0],
|
||||
}
|
||||
)
|
||||
real_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [2.5, 3.5],
|
||||
'target': [25.0, 35.0],
|
||||
'prediction': [25.0, 35.0],
|
||||
}
|
||||
)
|
||||
|
||||
result = rce_drift(reference_data, real_data, 'target')
|
||||
|
||||
# Result should be a Series with same length as real_data
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == len(real_data)
|
||||
|
||||
|
||||
def test_rce_drift_handles_common_columns():
|
||||
"""Test rce_drift correctly handles common columns between datasets."""
|
||||
reference_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.0, 2.0, 3.0],
|
||||
'feature2': [2.0, 3.0, 4.0],
|
||||
'extra_ref': [100.0, 200.0, 300.0],
|
||||
'target': [10.0, 20.0, 30.0],
|
||||
'prediction': [11.0, 21.0, 31.0],
|
||||
}
|
||||
)
|
||||
real_data = pd.DataFrame(
|
||||
{
|
||||
'feature1': [1.5],
|
||||
'feature2': [2.5],
|
||||
'extra_real': [150.0],
|
||||
'target': [15.0],
|
||||
'prediction': [16.0],
|
||||
}
|
||||
)
|
||||
|
||||
result = rce_drift(reference_data, real_data, 'target')
|
||||
|
||||
# Should work with only common columns
|
||||
assert isinstance(result, pd.Series)
|
||||
assert len(result) == 1
|
||||
437
tests/sientia/test_reports.py
Normal file
437
tests/sientia/test_reports.py
Normal file
@@ -0,0 +1,437 @@
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
try:
|
||||
from model_manager.sientia import reports
|
||||
except ImportError as exc:
|
||||
pytest.skip(
|
||||
f'reports requires Evidently API matching production pin: {exc}',
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_color_options(monkeypatch):
|
||||
def fake_color_options(**kwargs):
|
||||
return dict(kwargs)
|
||||
|
||||
monkeypatch.setattr(reports, 'ColorOptions', fake_color_options)
|
||||
|
||||
|
||||
def test_load_html_from_file_success(tmp_path):
|
||||
sample_file = tmp_path / 'sample.html'
|
||||
sample_file.write_text('<p>Hello</p>', encoding='utf-8')
|
||||
|
||||
content = reports.load_html_from_file(str(sample_file))
|
||||
|
||||
assert content == '<p>Hello</p>'
|
||||
|
||||
|
||||
def test_load_html_from_file_missing_file():
|
||||
with pytest.raises(FileNotFoundError):
|
||||
reports.load_html_from_file('non-existent.html')
|
||||
|
||||
|
||||
def test_load_html_from_file_os_error(monkeypatch):
|
||||
def fake_open(*_args, **_kwargs):
|
||||
raise OSError('boom')
|
||||
|
||||
monkeypatch.setattr('builtins.open', fake_open)
|
||||
|
||||
with pytest.raises(OSError, match='boom'):
|
||||
reports.load_html_from_file('path.html')
|
||||
|
||||
|
||||
def test_inject_content_replaces_section():
|
||||
main_html = "<html><body><div id='target'>old</div></body></html>"
|
||||
content = '<span>new</span>'
|
||||
|
||||
result = reports.inject_content(main_html, 'target', content)
|
||||
|
||||
soup = BeautifulSoup(result, 'html.parser')
|
||||
section = soup.find(id='target')
|
||||
assert section is not None
|
||||
assert section.find('span').text == 'new'
|
||||
|
||||
|
||||
def test_inject_content_missing_section():
|
||||
main_html = "<html><body><div id='other'>keep</div></body></html>"
|
||||
|
||||
result = reports.inject_content(main_html, 'missing', '<p>ignored</p>')
|
||||
|
||||
# Content should be unchanged when section is missing
|
||||
soup = BeautifulSoup(result, 'html.parser')
|
||||
assert soup.find(id='other') is not None
|
||||
assert soup.find(id='other').text == 'keep'
|
||||
|
||||
|
||||
def test_reports_init_sets_defaults(stub_color_options):
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
|
||||
assert report.metrics == []
|
||||
assert isinstance(report.options, list) and len(report.options) == 1
|
||||
assert report.sections == {}
|
||||
assert report.base_path is None
|
||||
|
||||
|
||||
def test_add_data_quality_section_without_run(monkeypatch, stub_color_options):
|
||||
monkeypatch.setattr(reports, 'DatasetSummaryMetric', lambda: 'summary')
|
||||
monkeypatch.setattr(
|
||||
reports,
|
||||
'generate_column_metrics',
|
||||
lambda *args, **kwargs: ('columns', kwargs),
|
||||
)
|
||||
monkeypatch.setattr(reports, 'ConflictTargetMetric', lambda: 'conflict')
|
||||
monkeypatch.setattr(reports, 'DatasetCorrelationsMetric', lambda: 'correlations')
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
report.add_data_quality_section(columns=['col'], run=False)
|
||||
|
||||
assert report.metrics[-4:] == [
|
||||
'summary',
|
||||
('columns', {'columns': ['col'], 'skip_id_column': True}),
|
||||
'conflict',
|
||||
'correlations',
|
||||
]
|
||||
assert 'data_quality' not in report.sections
|
||||
|
||||
|
||||
def test_add_data_quality_section_with_run(monkeypatch, tmp_path, stub_color_options):
|
||||
summary = object()
|
||||
column_metrics = object()
|
||||
conflict = object()
|
||||
correlations = object()
|
||||
monkeypatch.setattr(reports, 'DatasetSummaryMetric', lambda: summary)
|
||||
|
||||
def fake_generate_column_metrics(*_args, **kwargs):
|
||||
return column_metrics
|
||||
|
||||
monkeypatch.setattr(reports, 'generate_column_metrics', fake_generate_column_metrics)
|
||||
monkeypatch.setattr(reports, 'ConflictTargetMetric', lambda: conflict)
|
||||
monkeypatch.setattr(reports, 'DatasetCorrelationsMetric', lambda: correlations)
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'data_quality'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(
|
||||
reference_data='ref', current_data='cur', target_name='target', base_path=str(tmp_path)
|
||||
)
|
||||
report.add_data_quality_section(columns=['c1'], run=True)
|
||||
|
||||
assert report.metrics[-4:] == [summary, column_metrics, conflict, correlations]
|
||||
assert report.sections['data_quality'] == {'result': 'data_quality'}
|
||||
ReportMock.assert_called_once_with(
|
||||
metrics=[summary, column_metrics, conflict, correlations], options=report.options
|
||||
)
|
||||
run_kwargs = report_instance.run.call_args.kwargs
|
||||
assert run_kwargs['reference_data'] == 'ref'
|
||||
assert run_kwargs['current_data'] == 'cur'
|
||||
assert run_kwargs['column_mapping'].target == 'target'
|
||||
report_instance.save_html.assert_called_once_with(
|
||||
os.path.join(str(tmp_path), 'data_quality.html')
|
||||
)
|
||||
|
||||
|
||||
def test_add_data_quality_section_run_without_base_path(monkeypatch, stub_color_options):
|
||||
summary = object()
|
||||
column_metrics = object()
|
||||
conflict = object()
|
||||
correlations = object()
|
||||
monkeypatch.setattr(reports, 'DatasetSummaryMetric', lambda: summary)
|
||||
monkeypatch.setattr(
|
||||
reports,
|
||||
'generate_column_metrics',
|
||||
lambda *args, **kwargs: column_metrics,
|
||||
)
|
||||
monkeypatch.setattr(reports, 'ConflictTargetMetric', lambda: conflict)
|
||||
monkeypatch.setattr(reports, 'DatasetCorrelationsMetric', lambda: correlations)
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'quality'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
report.add_data_quality_section(run=True)
|
||||
|
||||
assert report.sections['data_quality'] == {'result': 'quality'}
|
||||
report_instance.save_html.assert_not_called()
|
||||
|
||||
|
||||
def test_add_data_quality_section_non_default_target_keeps_conflict_metric(
|
||||
monkeypatch, stub_color_options
|
||||
):
|
||||
summary = object()
|
||||
column_metrics = object()
|
||||
conflict = object()
|
||||
correlations = object()
|
||||
|
||||
monkeypatch.setattr(reports, 'DatasetSummaryMetric', lambda: summary)
|
||||
monkeypatch.setattr(
|
||||
reports,
|
||||
'generate_column_metrics',
|
||||
lambda *args, **kwargs: column_metrics,
|
||||
)
|
||||
monkeypatch.setattr(reports, 'ConflictTargetMetric', lambda: conflict)
|
||||
monkeypatch.setattr(reports, 'DatasetCorrelationsMetric', lambda: correlations)
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='sales')
|
||||
report.add_data_quality_section(columns=['c1'], run=False)
|
||||
|
||||
assert report.metrics[-4:] == [
|
||||
summary,
|
||||
column_metrics,
|
||||
conflict,
|
||||
correlations,
|
||||
]
|
||||
|
||||
|
||||
def test_add_data_drift_section_paths(monkeypatch, tmp_path, stub_color_options):
|
||||
drift_instances = [object(), object(), object()]
|
||||
DataDriftPresetMock = MagicMock(side_effect=drift_instances)
|
||||
monkeypatch.setattr(reports, 'DataDriftPreset', DataDriftPresetMock)
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'data_drift'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(
|
||||
reference_data='ref', current_data='cur', target_name='target', base_path=str(tmp_path)
|
||||
)
|
||||
report.add_data_drift_section(columns=['c1'], run=False)
|
||||
assert report.metrics[-1] == drift_instances[0]
|
||||
assert 'data_drift' not in report.sections
|
||||
|
||||
report.add_data_drift_section(columns=['c1'], run=True)
|
||||
assert report.sections['data_drift'] == {'result': 'data_drift'}
|
||||
ReportMock.assert_called_with(metrics=[drift_instances[2]], options=report.options)
|
||||
run_kwargs = report_instance.run.call_args.kwargs
|
||||
assert run_kwargs['reference_data'] == 'ref'
|
||||
assert run_kwargs['current_data'] == 'cur'
|
||||
assert run_kwargs['column_mapping'].target == 'target'
|
||||
report_instance.save_html.assert_called_with(os.path.join(str(tmp_path), 'data_drift.html'))
|
||||
|
||||
|
||||
def test_add_data_drift_section_run_without_base_path(monkeypatch, stub_color_options):
|
||||
drift_instances = [object(), object(), object()]
|
||||
DataDriftPresetMock = MagicMock(side_effect=drift_instances)
|
||||
monkeypatch.setattr(reports, 'DataDriftPreset', DataDriftPresetMock)
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'drift'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
report.add_data_drift_section(run=True)
|
||||
|
||||
assert report.sections['data_drift'] == {'result': 'drift'}
|
||||
report_instance.save_html.assert_not_called()
|
||||
|
||||
|
||||
def test_add_regression_section(monkeypatch, tmp_path, stub_color_options):
|
||||
regression_metrics = [object() for _ in range(7)]
|
||||
monkeypatch.setattr(reports, 'RegressionPerformanceMetrics', lambda: regression_metrics[0])
|
||||
monkeypatch.setattr(reports, 'RegressionDummyMetric', lambda: regression_metrics[1])
|
||||
monkeypatch.setattr(
|
||||
reports, 'RegressionPredictedVsActualScatter', lambda: regression_metrics[2]
|
||||
)
|
||||
monkeypatch.setattr(reports, 'RegressionPredictedVsActualPlot', lambda: regression_metrics[3])
|
||||
monkeypatch.setattr(reports, 'RegressionErrorPlot', lambda: regression_metrics[4])
|
||||
monkeypatch.setattr(reports, 'RegressionAbsPercentageErrorPlot', lambda: regression_metrics[5])
|
||||
monkeypatch.setattr(reports, 'RegressionErrorDistribution', lambda: regression_metrics[6])
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'regression'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(
|
||||
reference_data='ref', current_data='cur', target_name='target', base_path=str(tmp_path)
|
||||
)
|
||||
|
||||
report.add_regression_section(run=False)
|
||||
assert report.metrics[-7:] == regression_metrics
|
||||
assert 'regression' not in report.sections
|
||||
|
||||
report.add_regression_section(run=True)
|
||||
assert report.sections['regression'] == {'result': 'regression'}
|
||||
ReportMock.assert_called_with(metrics=regression_metrics, options=report.options)
|
||||
report_instance.run.assert_called_with(
|
||||
reference_data='ref',
|
||||
current_data='cur',
|
||||
column_mapping=report_instance.run.call_args.kwargs['column_mapping'],
|
||||
)
|
||||
report_instance.save_html.assert_called_with(os.path.join(str(tmp_path), 'regression.html'))
|
||||
|
||||
|
||||
def test_add_regression_section_run_without_base_path(monkeypatch, stub_color_options):
|
||||
regression_metrics = [object() for _ in range(7)]
|
||||
monkeypatch.setattr(reports, 'RegressionPerformanceMetrics', lambda: regression_metrics[0])
|
||||
monkeypatch.setattr(reports, 'RegressionDummyMetric', lambda: regression_metrics[1])
|
||||
monkeypatch.setattr(
|
||||
reports, 'RegressionPredictedVsActualScatter', lambda: regression_metrics[2]
|
||||
)
|
||||
monkeypatch.setattr(reports, 'RegressionPredictedVsActualPlot', lambda: regression_metrics[3])
|
||||
monkeypatch.setattr(reports, 'RegressionErrorPlot', lambda: regression_metrics[4])
|
||||
monkeypatch.setattr(reports, 'RegressionAbsPercentageErrorPlot', lambda: regression_metrics[5])
|
||||
monkeypatch.setattr(reports, 'RegressionErrorDistribution', lambda: regression_metrics[6])
|
||||
|
||||
report_instance = MagicMock()
|
||||
report_instance.as_dict.return_value = {'result': 'reg'}
|
||||
ReportMock = MagicMock(return_value=report_instance)
|
||||
monkeypatch.setattr(reports, 'Report', ReportMock)
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
report.add_regression_section(run=True)
|
||||
|
||||
assert report.sections['regression'] == {'result': 'reg'}
|
||||
report_instance.save_html.assert_not_called()
|
||||
|
||||
|
||||
def test_set_color_options_appends(monkeypatch):
|
||||
calls = []
|
||||
|
||||
def color_options_mock(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return kwargs
|
||||
|
||||
monkeypatch.setattr(reports, 'ColorOptions', color_options_mock)
|
||||
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
report.set_color_options(primary_color='#111', secondary_color='#222')
|
||||
|
||||
options = report.options
|
||||
assert options is not None
|
||||
assert len(options) == 2
|
||||
assert calls[0]['primary_color'] == '#0F4C81'
|
||||
assert calls[1]['primary_color'] == '#111'
|
||||
assert options[1]['secondary_color'] == '#222'
|
||||
|
||||
|
||||
def test_save_all_sections_html_requires_base_path(stub_color_options):
|
||||
report = reports.Reports(reference_data='ref', current_data='cur', target_name='target')
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
report.save_all_sections_html('output/report.html')
|
||||
|
||||
|
||||
def test_save_all_sections_html_requires_template_path(stub_color_options, tmp_path):
|
||||
report = reports.Reports(
|
||||
reference_data='ref',
|
||||
current_data='cur',
|
||||
target_name='target',
|
||||
base_path=str(tmp_path),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match='template_path is required'):
|
||||
report.save_all_sections_html('output/report.html')
|
||||
|
||||
|
||||
def test_save_all_sections_html_writes_output(tmp_path, stub_color_options):
|
||||
base_dir = tmp_path / 'templates'
|
||||
base_dir.mkdir()
|
||||
(base_dir / 'header.html').write_text(
|
||||
"<html><body><div id='data_drift'></div><div id='data_quality'></div><div id='regression'></div></body></html>",
|
||||
encoding='utf-8',
|
||||
)
|
||||
(base_dir / 'data_drift.html').write_text('<p>Drift</p>', encoding='utf-8')
|
||||
(base_dir / 'data_quality.html').write_text('<p>Quality</p>', encoding='utf-8')
|
||||
(base_dir / 'regression.html').write_text('<p>Regression</p>', encoding='utf-8')
|
||||
|
||||
report = reports.Reports(
|
||||
reference_data='ref',
|
||||
current_data='cur',
|
||||
target_name='target',
|
||||
base_path=str(base_dir),
|
||||
template_path=str(base_dir),
|
||||
)
|
||||
output_path = tmp_path / 'reports' / 'combined.html'
|
||||
|
||||
report.save_all_sections_html(str(output_path))
|
||||
|
||||
assert output_path.exists()
|
||||
content = output_path.read_text(encoding='utf-8')
|
||||
assert '<p>Drift</p>' in content
|
||||
assert '<p>Quality</p>' in content
|
||||
assert '<p>Regression</p>' in content
|
||||
|
||||
|
||||
def test_save_all_sections_html_creates_directory(monkeypatch, tmp_path, stub_color_options):
|
||||
base_dir = tmp_path / 'templates'
|
||||
base_dir.mkdir()
|
||||
(base_dir / 'header.html').write_text(
|
||||
"<html><body><div id='data_drift'></div><div id='data_quality'></div><div id='regression'></div></body></html>",
|
||||
encoding='utf-8',
|
||||
)
|
||||
(base_dir / 'data_drift.html').write_text('<p>Drift</p>', encoding='utf-8')
|
||||
(base_dir / 'data_quality.html').write_text('<p>Quality</p>', encoding='utf-8')
|
||||
(base_dir / 'regression.html').write_text('<p>Regression</p>', encoding='utf-8')
|
||||
|
||||
make_dirs_called = []
|
||||
report = reports.Reports(
|
||||
reference_data='ref',
|
||||
current_data='cur',
|
||||
target_name='target',
|
||||
base_path=str(base_dir),
|
||||
template_path=str(base_dir),
|
||||
)
|
||||
output_path = tmp_path / 'nested' / 'report.html'
|
||||
output_dir = str(output_path.parent)
|
||||
|
||||
original_exists = os.path.exists
|
||||
original_makedirs = os.makedirs
|
||||
|
||||
def fake_exists(path):
|
||||
if path == output_dir:
|
||||
return False
|
||||
return original_exists(path)
|
||||
|
||||
def fake_makedirs(path, exist_ok=False):
|
||||
make_dirs_called.append((path, exist_ok))
|
||||
return original_makedirs(path, exist_ok=exist_ok)
|
||||
|
||||
monkeypatch.setattr(os.path, 'exists', fake_exists)
|
||||
monkeypatch.setattr(os, 'makedirs', fake_makedirs)
|
||||
|
||||
report.save_all_sections_html(str(output_path))
|
||||
|
||||
assert make_dirs_called == [(str(output_path.parent), True)]
|
||||
|
||||
|
||||
def test_save_all_sections_html_no_directory_needed(monkeypatch, tmp_path, stub_color_options):
|
||||
base_dir = tmp_path / 'templates'
|
||||
base_dir.mkdir()
|
||||
(base_dir / 'header.html').write_text(
|
||||
"<html><body><div id='data_drift'></div><div id='data_quality'></div><div id='regression'></div></body></html>",
|
||||
encoding='utf-8',
|
||||
)
|
||||
(base_dir / 'data_drift.html').write_text('<p>Drift</p>', encoding='utf-8')
|
||||
(base_dir / 'data_quality.html').write_text('<p>Quality</p>', encoding='utf-8')
|
||||
(base_dir / 'regression.html').write_text('<p>Regression</p>', encoding='utf-8')
|
||||
|
||||
mk_calls = []
|
||||
|
||||
def fake_makedirs(path, exist_ok=False):
|
||||
mk_calls.append((path, exist_ok))
|
||||
|
||||
monkeypatch.setattr(os, 'makedirs', fake_makedirs)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
report = reports.Reports(
|
||||
reference_data='ref',
|
||||
current_data='cur',
|
||||
target_name='target',
|
||||
base_path=str(base_dir),
|
||||
template_path=str(base_dir),
|
||||
)
|
||||
report.save_all_sections_html('report.html')
|
||||
|
||||
assert mk_calls == []
|
||||
assert (tmp_path / 'report.html').exists()
|
||||
379
tests/test_metrics.py
Normal file
379
tests/test_metrics.py
Normal file
@@ -0,0 +1,379 @@
|
||||
"""Unit tests for model_manager.metrics module.
|
||||
|
||||
This module tests the Prometheus metrics configuration used for
|
||||
monitoring and observability in the Sientia DataOps Model Manager.
|
||||
"""
|
||||
|
||||
|
||||
def test_app_up_metric_exists():
|
||||
"""Test that APP_UP metric is properly defined."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
assert APP_UP is not None
|
||||
assert APP_UP._name == 'app_up'
|
||||
assert (
|
||||
APP_UP._documentation == 'Indicates if the application is running (1) or shutting down (0)'
|
||||
)
|
||||
|
||||
|
||||
def test_app_up_metric_has_pod_id_label():
|
||||
"""Test that APP_UP metric has pod_id label."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
assert 'pod_id' in APP_UP._labelnames
|
||||
|
||||
|
||||
def test_app_up_metric_is_gauge():
|
||||
"""Test that APP_UP is a Gauge metric."""
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
assert isinstance(APP_UP, Gauge)
|
||||
|
||||
|
||||
def test_app_up_metric_can_be_set_to_one():
|
||||
"""Test that APP_UP metric can be set to 1 (running)."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
# Set metric to 1 for a specific pod
|
||||
APP_UP.labels(pod_id='test-pod-1').set(1)
|
||||
|
||||
# Verify the metric value
|
||||
metric_value = APP_UP.labels(pod_id='test-pod-1')._value._value
|
||||
assert metric_value == 1
|
||||
|
||||
|
||||
def test_app_up_metric_can_be_set_to_zero():
|
||||
"""Test that APP_UP metric can be set to 0 (shutting down)."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
# Set metric to 0 for a specific pod
|
||||
APP_UP.labels(pod_id='test-pod-2').set(0)
|
||||
|
||||
# Verify the metric value
|
||||
metric_value = APP_UP.labels(pod_id='test-pod-2')._value._value
|
||||
assert metric_value == 0
|
||||
|
||||
|
||||
def test_app_up_metric_multiple_pods():
|
||||
"""Test that APP_UP metric can track multiple pods independently."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
# Set different values for different pods
|
||||
APP_UP.labels(pod_id='pod-1').set(1)
|
||||
APP_UP.labels(pod_id='pod-2').set(0)
|
||||
APP_UP.labels(pod_id='pod-3').set(1)
|
||||
|
||||
# Verify each pod has correct value
|
||||
assert APP_UP.labels(pod_id='pod-1')._value._value == 1
|
||||
assert APP_UP.labels(pod_id='pod-2')._value._value == 0
|
||||
assert APP_UP.labels(pod_id='pod-3')._value._value == 1
|
||||
|
||||
|
||||
def test_app_up_metric_default_value():
|
||||
"""Test that APP_UP metric starts with no value set."""
|
||||
# Create a new label that hasn't been used yet
|
||||
import uuid
|
||||
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
unique_pod = f'test-pod-{uuid.uuid4()}'
|
||||
|
||||
# The metric should exist but not have a value until set
|
||||
metric = APP_UP.labels(pod_id=unique_pod)
|
||||
assert metric is not None
|
||||
|
||||
|
||||
def test_metrics_module_imports():
|
||||
"""Test that metrics module can be imported successfully."""
|
||||
import model_manager.metrics
|
||||
|
||||
assert hasattr(model_manager.metrics, 'APP_UP')
|
||||
assert hasattr(model_manager.metrics, 'Gauge')
|
||||
|
||||
|
||||
def test_metrics_module_docstring():
|
||||
"""Test that metrics module has proper documentation."""
|
||||
import model_manager.metrics
|
||||
|
||||
assert model_manager.metrics.__doc__ is not None
|
||||
assert 'Prometheus' in model_manager.metrics.__doc__
|
||||
assert 'metrics' in model_manager.metrics.__doc__
|
||||
|
||||
|
||||
def test_app_up_metric_can_increment():
|
||||
"""Test that APP_UP metric value can be incremented."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
pod_id = 'test-pod-increment'
|
||||
APP_UP.labels(pod_id=pod_id).set(0)
|
||||
|
||||
# Increment the metric
|
||||
APP_UP.labels(pod_id=pod_id).inc()
|
||||
|
||||
metric_value = APP_UP.labels(pod_id=pod_id)._value._value
|
||||
assert metric_value == 1
|
||||
|
||||
|
||||
def test_app_up_metric_can_decrement():
|
||||
"""Test that APP_UP metric value can be decremented."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
pod_id = 'test-pod-decrement'
|
||||
APP_UP.labels(pod_id=pod_id).set(1)
|
||||
|
||||
# Decrement the metric
|
||||
APP_UP.labels(pod_id=pod_id).dec()
|
||||
|
||||
metric_value = APP_UP.labels(pod_id=pod_id)._value._value
|
||||
assert metric_value == 0
|
||||
|
||||
|
||||
def test_app_up_metric_set_to_timestamp():
|
||||
"""Test that APP_UP metric can be set to current timestamp."""
|
||||
import time
|
||||
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
pod_id = 'test-pod-timestamp'
|
||||
current_time = time.time()
|
||||
|
||||
# Set to timestamp
|
||||
APP_UP.labels(pod_id=pod_id).set_to_current_time()
|
||||
|
||||
metric_value = APP_UP.labels(pod_id=pod_id)._value._value
|
||||
|
||||
# Should be close to current time
|
||||
assert abs(metric_value - current_time) < 2 # Within 2 seconds
|
||||
|
||||
|
||||
def test_app_up_metric_label_validation():
|
||||
"""Test that APP_UP metric validates label names."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
# Should work with valid label
|
||||
APP_UP.labels(pod_id='valid-pod-name').set(1)
|
||||
|
||||
# Should work with empty string (though not recommended)
|
||||
APP_UP.labels(pod_id='').set(1)
|
||||
|
||||
# Should work with special characters
|
||||
APP_UP.labels(pod_id='pod-123_test.example').set(1)
|
||||
|
||||
|
||||
def test_module_exports():
|
||||
"""Test that metrics module exports expected symbols."""
|
||||
import model_manager.metrics as metrics_module
|
||||
|
||||
# Check that module has the expected exports
|
||||
module_contents = dir(metrics_module)
|
||||
|
||||
assert 'APP_UP' in module_contents
|
||||
assert 'Gauge' in module_contents
|
||||
|
||||
|
||||
def test_app_up_metric_thread_safety():
|
||||
"""Test that APP_UP metric is thread-safe."""
|
||||
import threading
|
||||
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
pod_id = 'test-pod-threading'
|
||||
APP_UP.labels(pod_id=pod_id).set(0)
|
||||
|
||||
def increment_metric():
|
||||
for _ in range(100):
|
||||
APP_UP.labels(pod_id=pod_id).inc()
|
||||
|
||||
# Create multiple threads that increment the metric
|
||||
threads = [threading.Thread(target=increment_metric) for _ in range(5)]
|
||||
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
# Should have incremented 500 times total
|
||||
metric_value = APP_UP.labels(pod_id=pod_id)._value._value
|
||||
assert metric_value == 500
|
||||
|
||||
|
||||
def test_prometheus_client_gauge_import():
|
||||
"""Test that Gauge is properly imported from prometheus_client."""
|
||||
from prometheus_client import Gauge as PrometheusGauge
|
||||
|
||||
from model_manager.metrics import Gauge
|
||||
|
||||
assert Gauge is PrometheusGauge
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Training metrics — existence, type, and labels
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_TRAINING_LABEL_NAMES = ('pod_id', 'model_name', 'model_type')
|
||||
|
||||
|
||||
def _assert_training_labels(metric):
|
||||
for label in _TRAINING_LABEL_NAMES:
|
||||
assert label in metric._labelnames
|
||||
|
||||
|
||||
def test_sientia_training_data_preparation_lag_is_histogram():
|
||||
from prometheus_client import Histogram
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_DATA_PREPARATION_LAG
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_DATA_PREPARATION_LAG, Histogram)
|
||||
assert SIENTIA_TRAINING_DATA_PREPARATION_LAG._name == 'sientia_training_data_preparation_lag'
|
||||
_assert_training_labels(SIENTIA_TRAINING_DATA_PREPARATION_LAG)
|
||||
|
||||
|
||||
def test_sientia_training_data_preparation_error_count_total_is_counter():
|
||||
from prometheus_client import Counter
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_DATA_PREPARATION_ERROR_COUNT_TOTAL
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_DATA_PREPARATION_ERROR_COUNT_TOTAL, Counter)
|
||||
assert 'sientia_training_data_preparation_error_count' in SIENTIA_TRAINING_DATA_PREPARATION_ERROR_COUNT_TOTAL._name
|
||||
_assert_training_labels(SIENTIA_TRAINING_DATA_PREPARATION_ERROR_COUNT_TOTAL)
|
||||
|
||||
|
||||
def test_sientia_training_model_fit_lag_is_histogram():
|
||||
from prometheus_client import Histogram
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_FIT_LAG
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_FIT_LAG, Histogram)
|
||||
assert SIENTIA_TRAINING_MODEL_FIT_LAG._name == 'sientia_training_model_fit_lag'
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_FIT_LAG)
|
||||
|
||||
|
||||
def test_sientia_training_model_fit_error_count_total_is_counter():
|
||||
from prometheus_client import Counter
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_FIT_ERROR_COUNT_TOTAL
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_FIT_ERROR_COUNT_TOTAL, Counter)
|
||||
assert 'sientia_training_model_fit_error_count' in SIENTIA_TRAINING_MODEL_FIT_ERROR_COUNT_TOTAL._name
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_FIT_ERROR_COUNT_TOTAL)
|
||||
|
||||
|
||||
def test_sientia_training_model_quality_mse_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_QUALITY_MSE
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_QUALITY_MSE, Gauge)
|
||||
assert SIENTIA_TRAINING_MODEL_QUALITY_MSE._name == 'sientia_training_model_quality_mse'
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_QUALITY_MSE)
|
||||
|
||||
|
||||
def test_sientia_training_model_quality_mae_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_QUALITY_MAE
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_QUALITY_MAE, Gauge)
|
||||
assert SIENTIA_TRAINING_MODEL_QUALITY_MAE._name == 'sientia_training_model_quality_mae'
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_QUALITY_MAE)
|
||||
|
||||
|
||||
def test_sientia_training_model_quality_r2_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_QUALITY_R2
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_QUALITY_R2, Gauge)
|
||||
assert SIENTIA_TRAINING_MODEL_QUALITY_R2._name == 'sientia_training_model_quality_r2'
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_QUALITY_R2)
|
||||
|
||||
|
||||
def test_sientia_training_dataset_train_rows_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_DATASET_TRAIN_ROWS
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_DATASET_TRAIN_ROWS, Gauge)
|
||||
assert SIENTIA_TRAINING_DATASET_TRAIN_ROWS._name == 'sientia_training_dataset_train_rows'
|
||||
_assert_training_labels(SIENTIA_TRAINING_DATASET_TRAIN_ROWS)
|
||||
|
||||
|
||||
def test_sientia_training_dataset_val_rows_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_DATASET_VAL_ROWS
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_DATASET_VAL_ROWS, Gauge)
|
||||
assert SIENTIA_TRAINING_DATASET_VAL_ROWS._name == 'sientia_training_dataset_val_rows'
|
||||
_assert_training_labels(SIENTIA_TRAINING_DATASET_VAL_ROWS)
|
||||
|
||||
|
||||
def test_sientia_training_model_trained_total_is_counter():
|
||||
from prometheus_client import Counter
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_MODEL_TRAINED_TOTAL
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_MODEL_TRAINED_TOTAL, Counter)
|
||||
assert 'sientia_training_model_trained' in SIENTIA_TRAINING_MODEL_TRAINED_TOTAL._name
|
||||
_assert_training_labels(SIENTIA_TRAINING_MODEL_TRAINED_TOTAL)
|
||||
|
||||
|
||||
def test_sientia_training_feature_count_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_FEATURE_COUNT
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_FEATURE_COUNT, Gauge)
|
||||
assert SIENTIA_TRAINING_FEATURE_COUNT._name == 'sientia_training_feature_count'
|
||||
_assert_training_labels(SIENTIA_TRAINING_FEATURE_COUNT)
|
||||
|
||||
|
||||
def test_sientia_training_info_is_gauge():
|
||||
from prometheus_client import Gauge
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_INFO
|
||||
|
||||
assert isinstance(SIENTIA_TRAINING_INFO, Gauge)
|
||||
assert SIENTIA_TRAINING_INFO._name == 'sientia_training_info'
|
||||
|
||||
expected_labels = {
|
||||
'pod_id', 'model_name', 'model_type',
|
||||
'dataset_train_rows', 'dataset_val_rows', 'feature_count',
|
||||
'mse', 'mae', 'r2',
|
||||
}
|
||||
assert expected_labels == set(SIENTIA_TRAINING_INFO._labelnames)
|
||||
|
||||
|
||||
def test_sientia_training_info_set_value():
|
||||
import time
|
||||
|
||||
from model_manager.metrics import SIENTIA_TRAINING_INFO
|
||||
|
||||
ts = time.time() * 1000
|
||||
SIENTIA_TRAINING_INFO.labels(
|
||||
pod_id='test-pod',
|
||||
model_name='my_model',
|
||||
model_type='linear',
|
||||
dataset_train_rows='1000',
|
||||
dataset_val_rows='200',
|
||||
feature_count='5',
|
||||
mse='0.01',
|
||||
mae='0.08',
|
||||
r2='0.95',
|
||||
).set(ts)
|
||||
|
||||
value = SIENTIA_TRAINING_INFO.labels(
|
||||
pod_id='test-pod',
|
||||
model_name='my_model',
|
||||
model_type='linear',
|
||||
dataset_train_rows='1000',
|
||||
dataset_val_rows='200',
|
||||
feature_count='5',
|
||||
mse='0.01',
|
||||
mae='0.08',
|
||||
r2='0.95',
|
||||
)._value._value
|
||||
assert abs(value - ts) < 2000
|
||||
18
tests/test_runtime_paths.py
Normal file
18
tests/test_runtime_paths.py
Normal file
@@ -0,0 +1,18 @@
|
||||
"""Tests for runtime filesystem layout constants."""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def test_ensure_runtime_directories_creates_expected_paths():
|
||||
from model_manager.runtime_paths import (
|
||||
LOGS_DIR,
|
||||
REPORTS_ROOT,
|
||||
REPORTS_TEMP_DIR,
|
||||
ensure_runtime_directories,
|
||||
)
|
||||
|
||||
with patch('model_manager.runtime_paths.makedirs') as makedirs_mock:
|
||||
ensure_runtime_directories()
|
||||
|
||||
created = {call.args[0] for call in makedirs_mock.call_args_list}
|
||||
assert created == {REPORTS_ROOT, REPORTS_TEMP_DIR, LOGS_DIR}
|
||||
0
tests/utils/__init__.py
Normal file
0
tests/utils/__init__.py
Normal file
0
tests/utils/models/__init__.py
Normal file
0
tests/utils/models/__init__.py
Normal file
75
tests/utils/models/test_experiment_status.py
Normal file
75
tests/utils/models/test_experiment_status.py
Normal file
@@ -0,0 +1,75 @@
|
||||
"""Unit tests for ExperimentStatus enum."""
|
||||
|
||||
from model_manager.utils.models.experiment_status import ExperimentStatus
|
||||
|
||||
|
||||
def test_experiment_status_values():
|
||||
"""Test that all expected status values exist."""
|
||||
assert ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR == 'ORCHESTRATOR_VALIDATION_ERROR'
|
||||
assert ExperimentStatus.ORCHESTRATOR_WAITING_PROC == 'ORCHESTRATOR_WAITING_PROC'
|
||||
assert ExperimentStatus.TRAINING_SUCCESS == 'TRAINING_SUCCESS'
|
||||
assert ExperimentStatus.TRAINING_ERROR == 'TRAINING_ERROR'
|
||||
|
||||
|
||||
def test_experiment_status_count():
|
||||
"""Test that enum has exactly 4 status values."""
|
||||
assert len(ExperimentStatus) == 4
|
||||
|
||||
|
||||
def test_experiment_status_is_string():
|
||||
"""Test that enum values are strings."""
|
||||
for status in ExperimentStatus:
|
||||
assert isinstance(status.value, str)
|
||||
assert isinstance(status, str)
|
||||
|
||||
|
||||
def test_experiment_status_membership():
|
||||
"""Test membership checks for status values."""
|
||||
assert 'ORCHESTRATOR_VALIDATION_ERROR' in [s.value for s in ExperimentStatus]
|
||||
assert 'ORCHESTRATOR_WAITING_PROC' in [s.value for s in ExperimentStatus]
|
||||
assert 'TRAINING_SUCCESS' in [s.value for s in ExperimentStatus]
|
||||
assert 'TRAINING_ERROR' in [s.value for s in ExperimentStatus]
|
||||
|
||||
|
||||
def test_experiment_status_iteration():
|
||||
"""Test that enum can be iterated."""
|
||||
statuses = list(ExperimentStatus)
|
||||
assert len(statuses) == 4
|
||||
assert ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR in statuses
|
||||
assert ExperimentStatus.ORCHESTRATOR_WAITING_PROC in statuses
|
||||
assert ExperimentStatus.TRAINING_SUCCESS in statuses
|
||||
assert ExperimentStatus.TRAINING_ERROR in statuses
|
||||
|
||||
|
||||
def test_experiment_status_comparison():
|
||||
"""Test that enum values can be compared with strings."""
|
||||
assert ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR == 'ORCHESTRATOR_VALIDATION_ERROR'
|
||||
assert ExperimentStatus.ORCHESTRATOR_WAITING_PROC == 'ORCHESTRATOR_WAITING_PROC'
|
||||
assert ExperimentStatus.TRAINING_SUCCESS == 'TRAINING_SUCCESS'
|
||||
assert str(ExperimentStatus.TRAINING_ERROR) != 'TRAINING_SUCCESS'
|
||||
|
||||
|
||||
def test_experiment_status_access_by_name():
|
||||
"""Test accessing enum members by name."""
|
||||
assert (
|
||||
ExperimentStatus['ORCHESTRATOR_VALIDATION_ERROR']
|
||||
== ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR
|
||||
)
|
||||
assert (
|
||||
ExperimentStatus['ORCHESTRATOR_WAITING_PROC'] == ExperimentStatus.ORCHESTRATOR_WAITING_PROC
|
||||
)
|
||||
assert ExperimentStatus['TRAINING_SUCCESS'] == ExperimentStatus.TRAINING_SUCCESS
|
||||
assert ExperimentStatus['TRAINING_ERROR'] == ExperimentStatus.TRAINING_ERROR
|
||||
|
||||
|
||||
def test_experiment_status_access_by_value():
|
||||
"""Test accessing enum members by value."""
|
||||
assert (
|
||||
ExperimentStatus('ORCHESTRATOR_VALIDATION_ERROR')
|
||||
== ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR
|
||||
)
|
||||
assert (
|
||||
ExperimentStatus('ORCHESTRATOR_WAITING_PROC') == ExperimentStatus.ORCHESTRATOR_WAITING_PROC
|
||||
)
|
||||
assert ExperimentStatus('TRAINING_SUCCESS') == ExperimentStatus.TRAINING_SUCCESS
|
||||
assert ExperimentStatus('TRAINING_ERROR') == ExperimentStatus.TRAINING_ERROR
|
||||
51
tests/utils/models/test_init.py
Normal file
51
tests/utils/models/test_init.py
Normal file
@@ -0,0 +1,51 @@
|
||||
"""Unit tests for models __init__.py module."""
|
||||
|
||||
from model_manager.utils.models import (
|
||||
ExperimentStatus,
|
||||
TrainModelParams,
|
||||
TrainModelResult,
|
||||
)
|
||||
|
||||
|
||||
def test_experiment_status_import():
|
||||
"""Test that ExperimentStatus can be imported from models package."""
|
||||
assert ExperimentStatus is not None
|
||||
assert hasattr(ExperimentStatus, 'ORCHESTRATOR_WAITING_PROC')
|
||||
assert hasattr(ExperimentStatus, 'TRAINING_SUCCESS')
|
||||
|
||||
|
||||
def test_train_model_params_import():
|
||||
"""Test that TrainModelParams can be imported from models package."""
|
||||
assert TrainModelParams is not None
|
||||
assert callable(TrainModelParams)
|
||||
|
||||
|
||||
def test_train_model_result_import():
|
||||
"""Test that TrainModelResult can be imported from models package."""
|
||||
assert TrainModelResult is not None
|
||||
# Dataclasses have __dataclass_fields__
|
||||
assert hasattr(TrainModelResult, '__dataclass_fields__')
|
||||
|
||||
|
||||
def test_all_exports():
|
||||
"""Test that __all__ contains all expected exports."""
|
||||
from model_manager.utils.models import __all__
|
||||
|
||||
assert 'ExperimentStatus' in __all__
|
||||
assert 'TrainModelParams' in __all__
|
||||
assert 'TrainModelResult' in __all__
|
||||
assert len(__all__) == 3
|
||||
|
||||
|
||||
def test_no_extra_exports():
|
||||
"""Test that only expected items are exported."""
|
||||
import model_manager.utils.models as models_module
|
||||
|
||||
# Get all public attributes (not starting with _)
|
||||
public_attrs = [attr for attr in dir(models_module) if not attr.startswith('_')]
|
||||
|
||||
# Should only have the 3 main classes
|
||||
expected_public = {'ExperimentStatus', 'TrainModelParams', 'TrainModelResult'}
|
||||
|
||||
# Check that our expected classes are present
|
||||
assert expected_public.issubset(set(public_attrs))
|
||||
321
tests/utils/models/test_train_model_params.py
Normal file
321
tests/utils/models/test_train_model_params.py
Normal file
@@ -0,0 +1,321 @@
|
||||
"""Unit tests for TrainModelParams (current schema)."""
|
||||
|
||||
import copy
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import (
|
||||
DEFAULT_TRAIN_DATE_FORMAT,
|
||||
TrainModelParams,
|
||||
validate_frontend_date_format,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def minimal_model_metadata() -> dict:
|
||||
"""Minimal truthy metadata so validate_business_rules passes schema lookup."""
|
||||
return {'schemas': {'components': {'schemas': {}}}}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_train_params_dict(minimal_model_metadata) -> dict:
|
||||
"""Valid dictionary for TrainModelParams.from_dict."""
|
||||
return {
|
||||
'variable_columns': ['var1', 'var2'],
|
||||
'target_variable': 'target',
|
||||
'bucket_name': 'test-bucket',
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': minimal_model_metadata,
|
||||
}
|
||||
|
||||
|
||||
def test_from_dict_success(valid_train_params_dict):
|
||||
"""from_dict builds params and experiment_name from model_name."""
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
assert params.variable_columns == ['var1', 'var2']
|
||||
assert params.target_variable == 'target'
|
||||
assert params.bucket_name == 'test-bucket'
|
||||
assert params.experiment_run_id == 1
|
||||
assert params.experiment_name == 'Linear Regression'
|
||||
assert params.model_metadata is valid_train_params_dict['model_metadata']
|
||||
|
||||
|
||||
def test_from_dict_date_format_omitted_uses_default(valid_train_params_dict):
|
||||
"""Missing date_format defaults to DEFAULT_TRAIN_DATE_FORMAT."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['date_format']
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
|
||||
|
||||
|
||||
def test_from_dict_date_format_blank_uses_default(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
|
||||
|
||||
|
||||
def test_from_dict_superfluous_date_column_camel_key_is_ignored(valid_train_params_dict):
|
||||
"""Only snake_case keys are read; dateColumn does not populate date_column."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['dateColumn'] = 'wrong_name'
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_column == 'timestamp'
|
||||
|
||||
|
||||
def test_from_dict_missing_date_column_raises(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['date_column']
|
||||
with pytest.raises(ValueError, match='date_column is required'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
|
||||
def test_from_dict_date_format_non_string_raises(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = 12345
|
||||
with pytest.raises(TypeError, match='date_format must be a string'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
|
||||
def test_from_dict_coerces_experiment_run_id_string(valid_train_params_dict):
|
||||
"""Numeric string experiment_run_id is coerced to int."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['experiment_run_id'] = '42'
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.experiment_run_id == 42
|
||||
|
||||
|
||||
def test_from_dict_model_metadata_none(valid_train_params_dict):
|
||||
"""model_metadata may be None before load_model_metadata activity."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = None
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.model_metadata is None
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_rejects_bool():
|
||||
"""Boolean must not be accepted as experiment_run_id."""
|
||||
with pytest.raises(TypeError, match='experiment_run_id must be an integer'):
|
||||
TrainModelParams._coerce_experiment_run_id(True)
|
||||
|
||||
|
||||
def test_parse_optional_model_metadata_rejects_list():
|
||||
"""model_metadata must be dict or None."""
|
||||
with pytest.raises(TypeError, match='model_metadata must be a dict or None'):
|
||||
TrainModelParams._parse_optional_model_metadata([])
|
||||
|
||||
|
||||
def test_check_none_raises_value_error():
|
||||
with pytest.raises(ValueError, match='test_field is required'):
|
||||
TrainModelParams._check_none(None, str, 'test_field')
|
||||
|
||||
|
||||
def test_check_none_raises_type_error():
|
||||
with pytest.raises(TypeError, match='test_field must be of type str'):
|
||||
TrainModelParams._check_none(123, str, 'test_field')
|
||||
|
||||
|
||||
def test_validate_business_rules_success(valid_train_params_dict):
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_missing_model_metadata(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = None
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='model_metadata is required'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_out_of_range(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['train_size'] = 5
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='train_size must be between'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_variable_columns(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['variable_columns'] = []
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='variable_columns cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_target(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['target_variable'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='target_variable cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_whitespace_date_column(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_column'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='date_column cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_from_dict_missing_required_key(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['bucket_name']
|
||||
with pytest.raises(ValueError, match='bucket_name is required'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
|
||||
def test_to_dict_roundtrip_keys(valid_train_params_dict):
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
d = params.to_dict()
|
||||
assert 'variable_columns' in d
|
||||
assert d['experiment_run_id'] == 1
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_float():
|
||||
assert TrainModelParams._coerce_experiment_run_id(2.0) == 2
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_none_raises():
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required'):
|
||||
TrainModelParams._coerce_experiment_run_id(None)
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_invalid_type():
|
||||
with pytest.raises(TypeError, match='integer or numeric string'):
|
||||
TrainModelParams._coerce_experiment_run_id([1])
|
||||
|
||||
|
||||
def test_validate_model_param_schema_validation_error(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {
|
||||
'type': 'object',
|
||||
'properties': {'x': {'type': 'integer'}},
|
||||
'required': ['x'],
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
with pytest.raises(ValueError, match='Model parameters validation failed'):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_unexpected_validator_error(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {'type': 'object'},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m:
|
||||
m.return_value.validate.side_effect = RuntimeError('boom')
|
||||
with pytest.raises(RuntimeError, match='boom'):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_date_format_invalid(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = 'not-an-allowed-format'
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='Invalid date_format'):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_required_strings_whitespace_bucket_file_model(valid_train_params_dict):
|
||||
for field, msg in [
|
||||
('bucket_name', 'bucket_name cannot be empty'),
|
||||
('file_name', 'file_name cannot be empty'),
|
||||
('model_name', 'model_name cannot be empty'),
|
||||
]:
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d[field] = ' '
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match=msg):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_only_data_model_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'data_model': {'type': 'object'}}}}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_only_model_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.model_kwargs = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_only_opt_params_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'opt_params': {'type': 'object'}}}}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.opt_params = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_all_schema_branches(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {'type': 'object'},
|
||||
'model': {'type': 'object'},
|
||||
'opt_params': {'type': 'object'},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
p.model_kwargs = {}
|
||||
p.opt_params = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_frontend_date_format_whitespace_returns():
|
||||
validate_frontend_date_format(' ')
|
||||
|
||||
|
||||
def test_validate_frontend_date_format_valid_returns():
|
||||
validate_frontend_date_format('dd/MM/yyyy HH:mm:ss')
|
||||
71
tests/utils/models/test_train_model_result.py
Normal file
71
tests/utils/models/test_train_model_result.py
Normal file
@@ -0,0 +1,71 @@
|
||||
"""Unit tests for TrainModelResult dataclass."""
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
from model_manager.utils.models.train_model_result import TrainModelResult
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_params() -> TrainModelParams:
|
||||
"""Minimal TrainModelParams for TrainModelResult tests."""
|
||||
return TrainModelParams.from_dict(
|
||||
{
|
||||
'variable_columns': ['a'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
'file_name': 'f.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_frames():
|
||||
train = pd.DataFrame({'a': [1, 2], 't': [1.0, 2.0]})
|
||||
val = pd.DataFrame({'a': [3], 't': [3.0]})
|
||||
return train, val
|
||||
|
||||
|
||||
def test_train_model_result_creation(sample_params, sample_frames):
|
||||
train, val = sample_frames
|
||||
result = TrainModelResult(params=sample_params, train_data=train, val_data=val)
|
||||
assert result.params is sample_params
|
||||
assert result.train_data.equals(train)
|
||||
assert result.val_data.equals(val)
|
||||
assert result.run_name is None
|
||||
|
||||
|
||||
def test_train_model_result_optional_paths(sample_params, sample_frames):
|
||||
train, val = sample_frames
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
train_data=train,
|
||||
val_data=val,
|
||||
run_name='run-1',
|
||||
run_id='rid',
|
||||
run_dir='/tmp/x',
|
||||
mse_val=0.1,
|
||||
mae_val=0.2,
|
||||
r2_val=0.99,
|
||||
)
|
||||
assert result.run_name == 'run-1'
|
||||
assert result.run_id == 'rid'
|
||||
assert result.run_dir == '/tmp/x'
|
||||
assert result.mse_val == 0.1
|
||||
630
tests/utils/repository/test_data_manager_repository.py
Normal file
630
tests/utils/repository/test_data_manager_repository.py
Normal file
@@ -0,0 +1,630 @@
|
||||
"""Unit tests for DataManagerRepository and module helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from model_manager.runtime_paths import REPORTS_ROOT
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
from model_manager.utils.models.train_model_result import TrainModelResult
|
||||
from model_manager.utils.repository import data_manager_repository as dmr
|
||||
|
||||
|
||||
def test_train_test_split_dataframe_shuffle():
|
||||
df = pd.DataFrame({'a': range(10)})
|
||||
tr, te = dmr.train_test_split(df, train_size=0.7, random_state=0, shuffle=True)
|
||||
assert len(tr) == 7 and len(te) == 3
|
||||
|
||||
|
||||
def test_train_test_split_dataframe_no_shuffle():
|
||||
df = pd.DataFrame({'a': range(10)})
|
||||
tr, te = dmr.train_test_split(df, train_size=0.5, shuffle=False)
|
||||
assert list(tr['a']) == [0, 1, 2, 3, 4]
|
||||
|
||||
|
||||
def test_train_test_split_dataframe_returns_dataframes():
|
||||
df = pd.DataFrame(np.arange(20).reshape(10, 2), columns=['a', 'b'])
|
||||
tr, te = dmr.train_test_split(df, train_size=0.5, shuffle=False, random_state=None)
|
||||
assert isinstance(tr, pd.DataFrame)
|
||||
assert isinstance(te, pd.DataFrame)
|
||||
assert tr.shape[0] == 5 and te.shape[0] == 5
|
||||
|
||||
|
||||
def _params(**kwargs) -> TrainModelParams:
|
||||
base: dict[str, Any] = {
|
||||
'variable_columns': ['v1'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
'file_name': 'f.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
base.update(kwargs)
|
||||
return TrainModelParams.from_dict(base)
|
||||
|
||||
|
||||
def test_ensure_date_column_parsed_missing_column_raises():
|
||||
df = pd.DataFrame({'a': [1]})
|
||||
p = _params(date_column='missing')
|
||||
with pytest.raises(ValueError, match='not found in dataset columns'):
|
||||
dmr._ensure_date_column_parsed(df, p)
|
||||
|
||||
|
||||
def test_ensure_date_column_parsed_success():
|
||||
df = pd.DataFrame({'a': range(3), 'ts': ['2024-01-01 10:00:00'] * 3})
|
||||
p = _params(date_column='ts')
|
||||
out = dmr._ensure_date_column_parsed(df, p)
|
||||
assert pd.api.types.is_datetime64_any_dtype(out['ts'])
|
||||
|
||||
|
||||
def test_ensure_date_column_parsed_naive_with_frontend_format():
|
||||
"""CSV timestamps without timezone use params.date_format strftime mapping."""
|
||||
df = pd.DataFrame({'ts': ['2025-06-02 00:00:00', '2025-06-02 01:00:00']})
|
||||
p = _params(date_column='ts', date_format='yyyy-MM-dd HH:mm:ss')
|
||||
out = dmr._ensure_date_column_parsed(df, p)
|
||||
assert pd.api.types.is_datetime64_any_dtype(out['ts'])
|
||||
|
||||
|
||||
def test_ensure_date_column_parsed_invalid_raises():
|
||||
df = pd.DataFrame({'a': range(3), 'ts': ['not-a-date'] * 3})
|
||||
p = _params(date_column='ts')
|
||||
with pytest.raises(ValueError, match='Failed to parse date column'):
|
||||
dmr._ensure_date_column_parsed(df, p)
|
||||
|
||||
|
||||
def test_prepare_training_data_csv_load_failure():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = _params()
|
||||
with patch(
|
||||
'model_manager.utils.repository.data_manager_repository.pd.read_csv',
|
||||
side_effect=pd.errors.ParserError('bad'),
|
||||
):
|
||||
with pytest.raises(ValueError, match='Failed to load training CSV'):
|
||||
repo.prepare_training_data(b'x', None, p, {})
|
||||
|
||||
|
||||
def test_prepare_training_data_empty_after_load():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
# Headers only; timestamp column present but no data rows
|
||||
csv_bytes = b'timestamp,v1,t\n'
|
||||
with pytest.raises(ValueError, match='Training data view is empty after transformation'):
|
||||
repo.prepare_training_data(csv_bytes, None, p, {})
|
||||
|
||||
|
||||
def test_prepare_training_data_empty_after_transformation(monkeypatch):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
monkeypatch.setattr(
|
||||
repo,
|
||||
'_configure_datetime_index',
|
||||
lambda *_args, **_kwargs: pd.DataFrame(columns=['v1', 't']),
|
||||
)
|
||||
monkeypatch.setattr(repo, '_set_timezone_on_index', lambda data, *_args, **_kwargs: data)
|
||||
with pytest.raises(ValueError, match='Training data view is empty after transformation'):
|
||||
repo.prepare_training_data(b'timestamp,v1,t\n', None, p, {})
|
||||
|
||||
|
||||
def _minimal_dict_for_prepare():
|
||||
return {
|
||||
'variable_columns': ['v1'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
'file_name': 'f.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
|
||||
|
||||
def _csv_bytes_with_ts(n_rows: int = 20) -> bytes:
|
||||
"""CSV with leading timestamp column (naive, matches default date_format)."""
|
||||
lines = ['timestamp,v1,t']
|
||||
for i in range(n_rows):
|
||||
lines.append(f'2024-01-{i + 1:02d} 00:00:00,{i},{i + 1}')
|
||||
return '\n'.join(lines).encode()
|
||||
|
||||
|
||||
def test_prepare_training_data_validation_csv_invalid():
|
||||
from io import BytesIO
|
||||
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
train_csv = _csv_bytes_with_ts(5)
|
||||
train_df = pd.read_csv(BytesIO(train_csv), sep=',', decimal='.')
|
||||
with patch.object(
|
||||
dmr.pd,
|
||||
'read_csv',
|
||||
side_effect=[train_df, pd.errors.ParserError('bad val')],
|
||||
):
|
||||
with pytest.raises(ValueError, match='Failed to load validation CSV'):
|
||||
repo.prepare_training_data(train_csv, b'broken', p, {})
|
||||
|
||||
|
||||
def test_prepare_training_data_validation_empty_val():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
train_csv = _csv_bytes_with_ts(5)
|
||||
val_csv = b'timestamp,v1,t\n'
|
||||
with pytest.raises(ValueError, match='Validation data view is empty'):
|
||||
repo.prepare_training_data(train_csv, val_csv, p, {})
|
||||
|
||||
|
||||
def test_prepare_training_data_drops_row_with_blank_timestamp():
|
||||
"""Rows with empty date_column values are removed before datetime parsing."""
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
d = _minimal_dict_for_prepare()
|
||||
d['date_column'] = 'timestamp'
|
||||
d['date_format'] = 'yyyy-MM-dd HH:mm:ss'
|
||||
p = TrainModelParams.from_dict(d)
|
||||
lines = ['timestamp,v1,t']
|
||||
for i in range(10):
|
||||
if i == 3:
|
||||
lines.append(',1.0,2.0')
|
||||
else:
|
||||
lines.append(f'2025-06-01 {i:02d}:00:00,1.0,2.0')
|
||||
csv = '\n'.join(lines).encode()
|
||||
res = repo.prepare_training_data(csv, None, p, {})
|
||||
assert len(res.train_data) + len(res.val_data) == 9
|
||||
|
||||
|
||||
def test_prepare_training_data_split_path():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
train_csv = _csv_bytes_with_ts(20)
|
||||
res = repo.prepare_training_data(train_csv, None, p, {})
|
||||
assert res.train_data is not None and res.val_data is not None
|
||||
|
||||
|
||||
def test_prepare_training_data_explicit_validation_success():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
train_csv = _csv_bytes_with_ts(10)
|
||||
val_csv = _csv_bytes_with_ts(5)
|
||||
res = repo.prepare_training_data(train_csv, val_csv, p, {})
|
||||
assert len(res.val_data) == 5
|
||||
|
||||
|
||||
def test_coerce_non_timestamp_columns_to_numeric_success():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
idx = pd.date_range('2024-01-01', periods=2, freq='h', tz='UTC')
|
||||
df = pd.DataFrame(
|
||||
{'v1': ['1.25', '2.75'], 't': ['10', '11']},
|
||||
index=idx,
|
||||
)
|
||||
out = repo._coerce_non_timestamp_columns_to_numeric(df, p, {})
|
||||
assert pd.api.types.is_numeric_dtype(out['v1'])
|
||||
assert pd.api.types.is_numeric_dtype(out['t'])
|
||||
assert float(out['v1'].iloc[0]) == 1.25
|
||||
assert float(out['t'].iloc[1]) == 11.0
|
||||
|
||||
|
||||
def test_coerce_non_timestamp_columns_to_numeric_invalid_values_to_nan():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
idx = pd.date_range('2024-01-01', periods=2, freq='h', tz='UTC')
|
||||
df = pd.DataFrame(
|
||||
{'v1': ['1.25', 'oops'], 't': ['10', 'bad']},
|
||||
index=idx,
|
||||
)
|
||||
out = repo._coerce_non_timestamp_columns_to_numeric(df, p, {})
|
||||
assert np.isnan(out['v1'].iloc[1])
|
||||
assert np.isnan(out['t'].iloc[1])
|
||||
|
||||
|
||||
def test_prepare_training_data_coerces_non_timestamp_columns_to_numeric():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
lines = ['timestamp,v1,t']
|
||||
for i in range(10):
|
||||
v1 = 'bad' if i == 4 else f'{i + 0.5}'
|
||||
t = 'bad' if i == 7 else f'{i + 1.0}'
|
||||
lines.append(f'2025-06-01 {i:02d}:00:00,{v1},{t}')
|
||||
csv = '\n'.join(lines).encode()
|
||||
res = repo.prepare_training_data(csv, None, p, {})
|
||||
joined = pd.concat([res.train_data, res.val_data], axis=0).sort_index()
|
||||
assert pd.api.types.is_numeric_dtype(joined['v1'])
|
||||
assert pd.api.types.is_numeric_dtype(joined['t'])
|
||||
assert joined['v1'].isna().sum() == 1
|
||||
assert joined['t'].isna().sum() == 1
|
||||
|
||||
|
||||
def test_as_series_series():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
s = pd.Series([1.0, 2.0])
|
||||
assert repo._as_series(s).equals(s)
|
||||
|
||||
|
||||
def test_as_series_one_column_df():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
df = pd.DataFrame({'x': [1.0, 2.0]})
|
||||
out = repo._as_series(df)
|
||||
assert isinstance(out, pd.Series)
|
||||
|
||||
|
||||
def test_as_series_multi_column_raises():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
df = pd.DataFrame({'a': [1.0], 'b': [2.0]})
|
||||
with pytest.raises(ValueError, match='single-column'):
|
||||
repo._as_series(df)
|
||||
|
||||
|
||||
def test_compute_regression_metrics_requires_y_pred():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0]}),
|
||||
val_data=pd.DataFrame({'t': [1.0]}),
|
||||
y_pred=None,
|
||||
)
|
||||
with pytest.raises(ValueError, match='y_pred must be set'):
|
||||
repo.compute_regression_metrics(tmr, MagicMock())
|
||||
|
||||
|
||||
def test_compute_regression_metrics_no_overlap():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0]}),
|
||||
val_data=pd.DataFrame({'t': [1.0]}, index=[10]),
|
||||
y_pred=pd.DataFrame({'p': [1.0]}, index=[20]),
|
||||
)
|
||||
with pytest.raises(ValueError, match='No overlapping indices'):
|
||||
repo.compute_regression_metrics(tmr, MagicMock())
|
||||
|
||||
|
||||
def test_compute_regression_metrics_linear_equation():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_type = 'linear_regression'
|
||||
idx = pd.Index([0, 1])
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
val_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
y_pred=pd.DataFrame({'p': [1.0, 2.0]}, index=idx),
|
||||
)
|
||||
regr = MagicMock()
|
||||
regr.coef_ = np.array([0.5])
|
||||
regr.intercept_ = 1.0
|
||||
wrapper = MagicMock()
|
||||
wrapper.model = MagicMock()
|
||||
wrapper.model.regr = regr
|
||||
out = repo.compute_regression_metrics(tmr, wrapper)
|
||||
assert out.mse_val is not None and out.equation is not None
|
||||
|
||||
|
||||
def test_compute_regression_metrics_linear_skips_equation_without_sklearn_regr():
|
||||
"""E2E dummy wrappers expose model without sklearn .regr; metrics still compute."""
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_type = 'linear_regression'
|
||||
idx = pd.Index([0, 1])
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
val_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
y_pred=pd.DataFrame({'p': [1.0, 2.0]}, index=idx),
|
||||
)
|
||||
wrapper = MagicMock()
|
||||
wrapper.model = object()
|
||||
out = repo.compute_regression_metrics(tmr, wrapper)
|
||||
assert out.mse_val is not None and out.equation is None
|
||||
|
||||
|
||||
def test_compute_regression_metrics_non_linear_skips_equation():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_type = 'xgboost'
|
||||
idx = pd.Index([0, 1])
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
val_data=pd.DataFrame({'t': [1.0, 2.0]}, index=idx),
|
||||
y_pred=pd.DataFrame({'p': [1.0, 2.0]}, index=idx),
|
||||
)
|
||||
out = repo.compute_regression_metrics(tmr, MagicMock())
|
||||
assert out.mse_val is not None and out.equation is None
|
||||
|
||||
|
||||
def test_configure_datetime_index_none_raises():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
with pytest.raises(ValueError, match='Data is None'):
|
||||
repo._configure_datetime_index(None, p, {})
|
||||
|
||||
|
||||
def test_configure_datetime_index_already_datetime_index():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
existing_idx = pd.date_range('2024-01-01', periods=3, freq='h')
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
'timestamp': pd.to_datetime(
|
||||
['2024-01-03 00:00:00', '2024-01-01 00:00:00', '2024-01-02 00:00:00']
|
||||
),
|
||||
'v1': [1, 2, 3],
|
||||
't': [1, 2, 3],
|
||||
},
|
||||
index=existing_idx,
|
||||
)
|
||||
out = repo._configure_datetime_index(df, p, {})
|
||||
assert isinstance(out.index, pd.DatetimeIndex)
|
||||
assert out.index.equals(pd.DatetimeIndex(pd.to_datetime(sorted(df['timestamp'].tolist()))))
|
||||
|
||||
|
||||
def test_configure_datetime_index_from_date_column():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict({**_minimal_dict_for_prepare(), 'date_column': 'mydate'})
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
'mydate': pd.date_range('2024-01-01', periods=3, freq='D'),
|
||||
'v1': [1, 2, 3],
|
||||
't': [1, 2, 3],
|
||||
}
|
||||
)
|
||||
out = repo._configure_datetime_index(df, p, {})
|
||||
assert isinstance(out.index, pd.DatetimeIndex)
|
||||
|
||||
|
||||
def test_configure_datetime_index_missing_date_column_raises():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict({**_minimal_dict_for_prepare(), 'date_column': 'mydate'})
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
'timestamp': pd.date_range('2024-01-01', periods=3, freq='D'),
|
||||
'v1': [1, 2, 3],
|
||||
't': [1, 2, 3],
|
||||
}
|
||||
)
|
||||
with pytest.raises(ValueError, match='date_column "mydate" not found'):
|
||||
repo._configure_datetime_index(df, p, {})
|
||||
|
||||
|
||||
def test_configure_datetime_index_non_datetime_date_column_raises():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict({**_minimal_dict_for_prepare(), 'date_column': 'mydate'})
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
'mydate': ['2024-01-01', '2024-01-02', '2024-01-03'],
|
||||
'v1': [1.0, 2.0, 3.0],
|
||||
't': [1.0, 2.0, 3.0],
|
||||
}
|
||||
)
|
||||
with pytest.raises(ValueError, match='must be datetime before index configuration'):
|
||||
repo._configure_datetime_index(df, p, {})
|
||||
|
||||
|
||||
def test_create_run_directory_permission_error():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
with patch(
|
||||
'model_manager.utils.repository.data_manager_repository.makedirs',
|
||||
side_effect=PermissionError('no'),
|
||||
):
|
||||
with pytest.raises(PermissionError, match='Permission denied'):
|
||||
repo._create_run_directory('/tmp', 'run', {})
|
||||
|
||||
|
||||
def test_create_run_directory_os_error():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
with patch(
|
||||
'model_manager.utils.repository.data_manager_repository.makedirs',
|
||||
side_effect=OSError('disk'),
|
||||
):
|
||||
with pytest.raises(OSError, match='Failed to create directory'):
|
||||
repo._create_run_directory('/tmp', 'run', {})
|
||||
|
||||
|
||||
def test_generate_report_success(tmp_path):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
val_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
y_train_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
y_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
run_name='testrun',
|
||||
)
|
||||
tmr.equation = {'target_variable': 't'}
|
||||
with (
|
||||
patch.object(repo, '_get_reports_directory', return_value=str(tmp_path)),
|
||||
patch('model_manager.utils.repository.data_manager_repository.Reports') as mrep,
|
||||
):
|
||||
instance = mrep.return_value
|
||||
instance.save_all_sections_html = Mock()
|
||||
out = repo.generate_report(tmr, {})
|
||||
assert out.report_path and out.train_data_path and out.test_data_path
|
||||
if out.equation_path:
|
||||
with open(out.equation_path, encoding='utf-8') as f:
|
||||
json.load(f)
|
||||
|
||||
|
||||
def test_generate_report_adds_target_alias_for_reports(tmp_path):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
val_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
y_train_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
y_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
run_name='testrun',
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(repo, '_get_reports_directory', return_value=str(tmp_path)),
|
||||
patch('model_manager.utils.repository.data_manager_repository.Reports') as mrep,
|
||||
):
|
||||
instance = mrep.return_value
|
||||
instance.save_all_sections_html = Mock()
|
||||
out = repo.generate_report(tmr, {})
|
||||
|
||||
kwargs = mrep.call_args.kwargs
|
||||
reference_data = kwargs['reference_data']
|
||||
current_data = kwargs['current_data']
|
||||
assert 'target' in reference_data.columns
|
||||
assert 'target' in current_data.columns
|
||||
assert reference_data['target'].equals(reference_data['t'])
|
||||
assert current_data['target'].equals(current_data['t'])
|
||||
|
||||
assert out.train_data_path is not None
|
||||
assert out.test_data_path is not None
|
||||
train_csv = pd.read_csv(out.train_data_path)
|
||||
test_csv = pd.read_csv(out.test_data_path)
|
||||
assert 'target' in train_csv.columns
|
||||
assert 'target' in test_csv.columns
|
||||
assert train_csv['target'].equals(train_csv['t'])
|
||||
assert test_csv['target'].equals(test_csv['t'])
|
||||
|
||||
|
||||
def test_generate_report_skips_equation_file_when_not_linear(tmp_path):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_type = 'other'
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
val_data=pd.DataFrame({'v1': [1.0, 2.0], 't': [1.0, 2.0]}),
|
||||
y_train_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
y_pred=pd.DataFrame({'t': [1.0, 2.0]}),
|
||||
run_name='testrun',
|
||||
equation={'k': 'v'},
|
||||
)
|
||||
with (
|
||||
patch.object(repo, '_get_reports_directory', return_value=str(tmp_path)),
|
||||
patch('model_manager.utils.repository.data_manager_repository.Reports'),
|
||||
):
|
||||
out = repo.generate_report(tmr, {})
|
||||
assert out.equation_path is None
|
||||
|
||||
|
||||
def test_generate_report_run_name_missing():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0]}),
|
||||
val_data=pd.DataFrame({'t': [1.0]}),
|
||||
run_name=None,
|
||||
)
|
||||
with pytest.raises(ValueError, match='run_name is not set'):
|
||||
repo.generate_report(tmr, {})
|
||||
|
||||
|
||||
def test_generate_report_requires_predictions():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
tmr = TrainModelResult(
|
||||
params=p,
|
||||
train_data=pd.DataFrame({'t': [1.0]}),
|
||||
val_data=pd.DataFrame({'t': [1.0]}),
|
||||
run_name='testrun',
|
||||
y_train_pred=None,
|
||||
y_pred=None,
|
||||
)
|
||||
with pytest.raises(ValueError, match='y_train_pred or y_pred is not set'):
|
||||
repo.generate_report(tmr, {})
|
||||
|
||||
|
||||
def test_cleanup_run_directory_empty():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
repo.cleanup_run_directory('', {})
|
||||
|
||||
|
||||
def test_cleanup_run_directory_exists(tmp_path):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
d = tmp_path / 'subdir'
|
||||
d.mkdir()
|
||||
repo.cleanup_run_directory(str(d), {})
|
||||
assert not d.exists()
|
||||
|
||||
|
||||
def test_cleanup_run_directory_missing(tmp_path):
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
repo.cleanup_run_directory(str(tmp_path / 'nope'), {})
|
||||
|
||||
|
||||
def test_extract_model_equation_polynomial_poly_names():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_kwargs = {'degree': 2, 'poly_feature_names': ['f1', 'f2']}
|
||||
regr = MagicMock()
|
||||
regr.coef_ = np.array([1.0, 2.0])
|
||||
regr.intercept_ = 3.0
|
||||
wrapper = MagicMock()
|
||||
wrapper.model = MagicMock()
|
||||
wrapper.model.regr = regr
|
||||
eq = repo._extract_model_equation(wrapper.model, p)
|
||||
assert 'equation_string' in eq and eq['degree'] == 2
|
||||
|
||||
|
||||
def test_extract_model_equation_extra_coefficients_ignored():
|
||||
"""More coefficients than feature names: only the first len(names) are used."""
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
regr = MagicMock()
|
||||
regr.coef_ = np.array([1.0, 2.0, 3.0])
|
||||
regr.intercept_ = 0.0
|
||||
wrapper = MagicMock()
|
||||
wrapper.model = MagicMock()
|
||||
wrapper.model.regr = regr
|
||||
eq = repo._extract_model_equation(wrapper.model, p)
|
||||
assert len(eq['coefficients']) == len(p.variable_columns)
|
||||
|
||||
|
||||
def test_extract_model_equation_more_features_than_coefficients():
|
||||
"""Polynomial feature names longer than coef array: extra names get no coefficient entry."""
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
|
||||
p.model_kwargs = {'degree': 2, 'poly_feature_names': ['a', 'b', 'c']}
|
||||
regr = MagicMock()
|
||||
regr.coef_ = np.array([1.0, 2.0])
|
||||
regr.intercept_ = 0.0
|
||||
wrapper = MagicMock()
|
||||
wrapper.model = MagicMock()
|
||||
wrapper.model.regr = regr
|
||||
eq = repo._extract_model_equation(wrapper.model, p)
|
||||
assert list(eq['coefficients'].keys()) == ['a', 'b']
|
||||
|
||||
|
||||
def test_get_reports_directory_path():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
reports_dir = repo._get_reports_directory()
|
||||
assert reports_dir == REPORTS_ROOT
|
||||
163
tests/utils/test_connectors_config.py
Normal file
163
tests/utils/test_connectors_config.py
Normal file
@@ -0,0 +1,163 @@
|
||||
from os import environ
|
||||
from unittest.mock import patch
|
||||
|
||||
from model_manager.utils.connectors_config import (
|
||||
build_minio_config,
|
||||
build_mlflow_config,
|
||||
build_mongodb_config,
|
||||
build_plugin_store_config,
|
||||
build_postgres_config,
|
||||
)
|
||||
|
||||
|
||||
def test_build_mlflow_config_with_env_vars():
|
||||
environ.pop('MLFLOW_URL', None)
|
||||
environ['MLFLOW_URL'] = 'http://test-host:8080'
|
||||
environ['MLFLOW_USERNAME'] = 'test-user'
|
||||
environ['MLFLOW_PASSWORD'] = 'test-pass'
|
||||
|
||||
config = build_mlflow_config()
|
||||
|
||||
assert config['url'] == 'http://test-host:8080'
|
||||
assert config['username'] == 'test-user'
|
||||
assert config['password'] == 'test-pass'
|
||||
|
||||
|
||||
def test_build_mlflow_config_with_defaults():
|
||||
environ.pop('MLFLOW_URL', None)
|
||||
environ.pop('MLFLOW_USERNAME', None)
|
||||
environ.pop('MLFLOW_PASSWORD', None)
|
||||
|
||||
config = build_mlflow_config()
|
||||
|
||||
assert config['url'] == 'http://localhost:5080'
|
||||
assert config['username'] == 'aignosi'
|
||||
assert config['password'] == 'aignosi'
|
||||
|
||||
|
||||
def test_build_postgres_config_with_env_vars():
|
||||
environ['POSTGRES_HOST'] = 'test-host'
|
||||
environ['POSTGRES_PORT'] = '5433'
|
||||
environ['POSTGRES_USER'] = 'test-user'
|
||||
environ['POSTGRES_PASSWORD'] = 'test-pass'
|
||||
environ['POSTGRES_DBNAME'] = 'test-db'
|
||||
environ['POSTGRES_MIN_CONNECTIONS'] = '10'
|
||||
environ['POSTGRES_MAX_CONNECTIONS'] = '30'
|
||||
|
||||
config = build_postgres_config()
|
||||
|
||||
assert config['host'] == 'test-host'
|
||||
assert config['port'] == 5433
|
||||
assert config['user'] == 'test-user'
|
||||
assert config['password'] == 'test-pass'
|
||||
assert config['dbname'] == 'test-db'
|
||||
assert config['min_connections'] == 10
|
||||
assert config['max_connections'] == 30
|
||||
|
||||
|
||||
def test_build_postgres_config_with_defaults():
|
||||
environ.pop('POSTGRES_HOST', None)
|
||||
environ.pop('POSTGRES_PORT', None)
|
||||
environ.pop('POSTGRES_USER', None)
|
||||
environ.pop('POSTGRES_PASSWORD', None)
|
||||
environ.pop('POSTGRES_DBNAME', None)
|
||||
environ.pop('POSTGRES_MIN_CONNECTIONS', None)
|
||||
environ.pop('POSTGRES_MAX_CONNECTIONS', None)
|
||||
|
||||
config = build_postgres_config()
|
||||
|
||||
assert config['host'] == 'localhost'
|
||||
assert config['port'] == 5432
|
||||
assert config['user'] == 'sientia'
|
||||
assert config['password'] == 'sientia'
|
||||
assert config['dbname'] == 'sientia'
|
||||
assert config['min_connections'] == 5
|
||||
assert config['max_connections'] == 20
|
||||
|
||||
|
||||
def test_build_mongo_db_config_with_env_vars():
|
||||
environ['MONGODB_USERNAME'] = 'sientia1'
|
||||
environ['MONGODB_PASSWORD'] = 'sientia1'
|
||||
environ['MONGODB_URL'] = 'localhost:27018'
|
||||
environ['MONGODB_DATABASE'] = 'test_db'
|
||||
environ['MONGODB_TTL_INDEX_HOURS'] = '1'
|
||||
|
||||
assert build_mongodb_config() == {
|
||||
'connection_string': 'mongodb://sientia1:sientia1@localhost:27018',
|
||||
'database_name': 'test_db',
|
||||
'ttl_index_seconds': 3600,
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
|
||||
|
||||
def test_build_mongo_db_config_with_defaults():
|
||||
environ.pop('MONGODB_USERNAME', None)
|
||||
environ.pop('MONGODB_PASSWORD', None)
|
||||
environ.pop('MONGODB_DATABASE', None)
|
||||
environ.pop('MONGODB_URL', None)
|
||||
environ.pop('MONGODB_TTL_INDEX_HOURS', None)
|
||||
assert build_mongodb_config() == {
|
||||
'connection_string': 'mongodb://root:wKZDbMNU1c@localhost:27018',
|
||||
'database_name': 'sientia',
|
||||
'ttl_index_seconds': 3600,
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
|
||||
|
||||
def test_build_minio_config_with_env_vars():
|
||||
environ['MINIO_ENDPOINT_URL'] = 'http://test-minio:9000'
|
||||
environ['MINIO_ACCESS_KEY'] = 'test-access-key'
|
||||
environ['MINIO_SECRET_KEY'] = 'test-secret-key'
|
||||
environ['MINIO_REGION'] = 'eu-west-1'
|
||||
environ['MINIO_SECURE'] = 'true'
|
||||
environ['MINIO_MAX_RETRY_ATTEMPTS'] = '5'
|
||||
environ['MINIO_RETRY_MODE'] = 'standard'
|
||||
environ['MINIO_CONNECT_TIMEOUT'] = '20'
|
||||
environ['MINIO_READ_TIMEOUT'] = '120'
|
||||
environ['MINIO_DEFAULT_BUCKET'] = 'my-bucket'
|
||||
|
||||
config = build_minio_config()
|
||||
|
||||
assert config['endpoint_url'] == 'http://test-minio:9000'
|
||||
assert config['access_key'] == 'test-access-key'
|
||||
assert config['secret_key'] == 'test-secret-key'
|
||||
assert config['region'] == 'eu-west-1'
|
||||
assert config['use_ssl'] is True
|
||||
assert config['max_retry_attempts'] == 5
|
||||
assert config['retry_mode'] == 'standard'
|
||||
assert config['connect_timeout'] == 20
|
||||
assert config['read_timeout'] == 120
|
||||
assert config['default_bucket'] == 'my-bucket'
|
||||
|
||||
|
||||
def test_build_plugin_store_config_cache_ttl_seconds():
|
||||
"""STORE_CACHE_TTL_SECONDS is parsed to int when set."""
|
||||
with patch.dict(environ, {'STORE_CACHE_TTL_SECONDS': '7200'}, clear=False):
|
||||
cfg = build_plugin_store_config()
|
||||
assert cfg['cache_ttl_seconds'] == 7200
|
||||
|
||||
|
||||
def test_build_minio_config_with_defaults():
|
||||
environ.pop('MINIO_ENDPOINT_URL', None)
|
||||
environ.pop('MINIO_ACCESS_KEY', None)
|
||||
environ.pop('MINIO_SECRET_KEY', None)
|
||||
environ.pop('MINIO_REGION', None)
|
||||
environ.pop('MINIO_SECURE', None)
|
||||
environ.pop('MINIO_MAX_RETRY_ATTEMPTS', None)
|
||||
environ.pop('MINIO_RETRY_MODE', None)
|
||||
environ.pop('MINIO_CONNECT_TIMEOUT', None)
|
||||
environ.pop('MINIO_READ_TIMEOUT', None)
|
||||
environ.pop('MINIO_DEFAULT_BUCKET', None)
|
||||
|
||||
config = build_minio_config()
|
||||
|
||||
assert config['endpoint_url'] == 'http://localhost:9000'
|
||||
assert config['access_key'] == 'minioadmin'
|
||||
assert config['secret_key'] == 'minioadmin'
|
||||
assert config['region'] == 'us-east-1'
|
||||
assert config['use_ssl'] is False
|
||||
assert config['max_retry_attempts'] == 3
|
||||
assert config['retry_mode'] == 'adaptive'
|
||||
assert config['connect_timeout'] == 10
|
||||
assert config['read_timeout'] == 60
|
||||
assert config['default_bucket'] == 'model-training'
|
||||
90
tests/utils/test_logger_helper.py
Normal file
90
tests/utils/test_logger_helper.py
Normal file
@@ -0,0 +1,90 @@
|
||||
"""Unit tests for logger_helper module with 100% coverage."""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
@patch('model_manager.utils.logger_helper.SientiaLogger')
|
||||
def test_get_logger_creates_logger_instance(mock_sientia_logger):
|
||||
"""Test get_logger creates a SientiaLogger instance with the given name."""
|
||||
from model_manager.utils.logger_helper import get_logger
|
||||
|
||||
mock_logger_instance = MagicMock()
|
||||
mock_logger_instance.base_logger = MagicMock()
|
||||
mock_sientia_logger.return_value = mock_logger_instance
|
||||
|
||||
result = get_logger('test_module')
|
||||
|
||||
mock_sientia_logger.assert_called_once_with('test_module')
|
||||
assert result is mock_logger_instance
|
||||
|
||||
|
||||
@patch('model_manager.utils.logger_helper.SientiaLogger')
|
||||
def test_get_logger_disables_propagation(mock_sientia_logger):
|
||||
"""Test get_logger disables log propagation."""
|
||||
from model_manager.utils.logger_helper import get_logger
|
||||
|
||||
mock_logger_instance = MagicMock()
|
||||
mock_base_logger = MagicMock()
|
||||
mock_base_logger.propagate = True
|
||||
mock_logger_instance.base_logger = mock_base_logger
|
||||
mock_sientia_logger.return_value = mock_logger_instance
|
||||
|
||||
get_logger('test_module')
|
||||
|
||||
assert mock_base_logger.propagate is False
|
||||
|
||||
|
||||
@patch('model_manager.utils.logger_helper.SientiaLogger')
|
||||
def test_get_logger_with_different_names(mock_sientia_logger):
|
||||
"""Test get_logger works with different logger names."""
|
||||
from model_manager.utils.logger_helper import get_logger
|
||||
|
||||
mock_logger_instance = MagicMock()
|
||||
mock_logger_instance.base_logger = MagicMock()
|
||||
mock_sientia_logger.return_value = mock_logger_instance
|
||||
|
||||
logger1 = get_logger('module1')
|
||||
logger2 = get_logger('module2')
|
||||
logger3 = get_logger('my.nested.module')
|
||||
|
||||
assert mock_sientia_logger.call_count == 3
|
||||
mock_sientia_logger.assert_any_call('module1')
|
||||
mock_sientia_logger.assert_any_call('module2')
|
||||
mock_sientia_logger.assert_any_call('my.nested.module')
|
||||
assert logger1 is mock_logger_instance
|
||||
assert logger2 is mock_logger_instance
|
||||
assert logger3 is mock_logger_instance
|
||||
|
||||
|
||||
@patch('model_manager.utils.logger_helper.SientiaLogger')
|
||||
def test_get_logger_with_empty_name(mock_sientia_logger):
|
||||
"""Test get_logger with empty string name."""
|
||||
from model_manager.utils.logger_helper import get_logger
|
||||
|
||||
mock_logger_instance = MagicMock()
|
||||
mock_logger_instance.base_logger = MagicMock()
|
||||
mock_sientia_logger.return_value = mock_logger_instance
|
||||
|
||||
result = get_logger('')
|
||||
|
||||
mock_sientia_logger.assert_called_once_with('')
|
||||
assert result is mock_logger_instance
|
||||
assert result.base_logger.propagate is False
|
||||
|
||||
|
||||
@patch('model_manager.utils.logger_helper.SientiaLogger')
|
||||
def test_get_logger_returns_configured_logger(mock_sientia_logger):
|
||||
"""Test get_logger returns the configured logger instance."""
|
||||
from model_manager.utils.logger_helper import get_logger
|
||||
|
||||
mock_logger_instance = MagicMock()
|
||||
mock_logger_instance.base_logger = MagicMock()
|
||||
mock_logger_instance.base_logger.propagate = True
|
||||
mock_sientia_logger.return_value = mock_logger_instance
|
||||
|
||||
result = get_logger('test_logger')
|
||||
|
||||
# Verify the logger is returned after configuration
|
||||
assert result is mock_logger_instance
|
||||
# Verify propagation was disabled
|
||||
assert mock_logger_instance.base_logger.propagate is False
|
||||
84
tests/worker/test_prepare_worker.py
Normal file
84
tests/worker/test_prepare_worker.py
Normal file
@@ -0,0 +1,84 @@
|
||||
"""Unit tests for local worker factory."""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def test_build_queue_name_without_runtime_uses_default_suffix():
|
||||
from model_manager.worker.prepare_worker import build_queue_name
|
||||
|
||||
assert build_queue_name('TrainModel') == 'train_model-queue'
|
||||
|
||||
|
||||
def test_prepare_worker_train_queue_uses_train_limits():
|
||||
from model_manager.worker.prepare_worker import prepare_worker
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
fake_worker = MagicMock()
|
||||
fake_client = MagicMock()
|
||||
fake_logger = MagicMock()
|
||||
|
||||
with patch(
|
||||
'model_manager.worker.prepare_worker.Worker', return_value=fake_worker
|
||||
) as worker_class:
|
||||
with patch.dict(
|
||||
'os.environ',
|
||||
{
|
||||
'TRAINMODEL_ACTIVITY_EXECUTOR_MAX_WORKERS': '3',
|
||||
'TRAINMODEL_MAX_CONCURRENT_ACTIVITIES': '6',
|
||||
'TRAINMODEL_MAX_CONCURRENT_WORKFLOW_TASKS': '10',
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
worker = prepare_worker(
|
||||
main_workflow=TrainModel,
|
||||
other_workflows=[],
|
||||
activities=[],
|
||||
temporal_client=fake_client,
|
||||
logger=fake_logger,
|
||||
runtime='model-manager-worker',
|
||||
)
|
||||
|
||||
assert worker is fake_worker
|
||||
worker_class.assert_called_once()
|
||||
kwargs = worker_class.call_args.kwargs
|
||||
assert kwargs['task_queue'] == 'train_model-model-manager-worker-queue'
|
||||
assert kwargs['max_concurrent_activities'] == 6
|
||||
assert kwargs['max_concurrent_workflow_tasks'] == 10
|
||||
assert kwargs['activity_executor']._max_workers == 3
|
||||
kwargs['activity_executor'].shutdown(wait=True, cancel_futures=True)
|
||||
|
||||
|
||||
def test_prepare_worker_cleanup_queue_uses_cleanup_limits():
|
||||
from model_manager.worker.prepare_worker import prepare_worker
|
||||
from model_manager.workflows.cleanup_files import CleanupFiles
|
||||
|
||||
fake_worker = MagicMock()
|
||||
fake_client = MagicMock()
|
||||
fake_logger = MagicMock()
|
||||
|
||||
with patch(
|
||||
'model_manager.worker.prepare_worker.Worker', return_value=fake_worker
|
||||
) as worker_class:
|
||||
with patch.dict(
|
||||
'os.environ',
|
||||
{
|
||||
'CLEANUPFILES_ACTIVITY_EXECUTOR_MAX_WORKERS': '5',
|
||||
'CLEANUPFILES_MAX_CONCURRENT_ACTIVITIES': '7',
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
worker = prepare_worker(
|
||||
main_workflow=CleanupFiles,
|
||||
other_workflows=[],
|
||||
activities=[],
|
||||
temporal_client=fake_client,
|
||||
logger=fake_logger,
|
||||
runtime='model-manager-worker',
|
||||
)
|
||||
|
||||
assert worker is fake_worker
|
||||
kwargs = worker_class.call_args.kwargs
|
||||
assert kwargs['task_queue'] == 'cleanup_files-model-manager-worker-queue'
|
||||
assert kwargs['max_concurrent_activities'] == 7
|
||||
assert kwargs['activity_executor']._max_workers == 5
|
||||
kwargs['activity_executor'].shutdown(wait=True, cancel_futures=True)
|
||||
907
tests/worker/test_worker.py
Normal file
907
tests/worker/test_worker.py
Normal file
@@ -0,0 +1,907 @@
|
||||
"""Unit tests for worker module."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_env_vars():
|
||||
"""Set up test environment variables."""
|
||||
env_vars = {
|
||||
'POD_ID': 'test-pod-123',
|
||||
'HTTP_METRICS_PORT': '9090',
|
||||
'HTTP_SDK_METRICS_PORT': '9091',
|
||||
'TEMPORAL_HOST': 'localhost:7233',
|
||||
'TEMPORAL_NAMESPACE': 'test-namespace',
|
||||
'PROJECT_NAME': 'test-project',
|
||||
'TRAIN_TASK_QUEUE': 'train_model-local_queue',
|
||||
'CLEANUP_TASK_QUEUE': 'cleanup-local_queue',
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'STORE_BASE_URL': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'STORE_OWNER': 'sientia',
|
||||
'STORE_REPO': 'model-library-store',
|
||||
}
|
||||
|
||||
with patch.dict(os.environ, env_vars, clear=False):
|
||||
yield env_vars
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Create a mock logger."""
|
||||
logger = Mock()
|
||||
logger.custom_info = Mock()
|
||||
logger.custom_error = Mock()
|
||||
return logger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_temporal_client():
|
||||
"""Create a mock Temporal client."""
|
||||
client_mock = AsyncMock()
|
||||
client_mock.connect = AsyncMock()
|
||||
return client_mock
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_worker():
|
||||
"""Create a mock Temporal worker."""
|
||||
worker_mock = Mock()
|
||||
worker_mock.run = AsyncMock(return_value=None)
|
||||
return worker_mock
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_notification_handler():
|
||||
"""Create a mock notification handler."""
|
||||
handler = Mock()
|
||||
handler.shutdown = Mock()
|
||||
return handler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_activities():
|
||||
"""Create a mock Activities instance."""
|
||||
activities = AsyncMock()
|
||||
activities.update_experiment_run = Mock()
|
||||
activities.load_model_metadata = Mock()
|
||||
activities.validate_train_params = Mock()
|
||||
activities.train_model = Mock()
|
||||
activities.cleanup_resources = Mock()
|
||||
activities.shutdown = Mock()
|
||||
return activities
|
||||
|
||||
|
||||
def test_pod_id_from_env():
|
||||
"""Test that POD_ID is correctly read from environment."""
|
||||
with patch.dict(os.environ, {'POD_ID': 'pod-test-123'}):
|
||||
# Re-import to get new env value
|
||||
import importlib
|
||||
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
importlib.reload(worker_module)
|
||||
|
||||
assert worker_module.POD_ID == 'pod-test-123'
|
||||
|
||||
|
||||
def test_sdk_metrics_port_default():
|
||||
"""Test that SDK_METRICS_PORT uses default value."""
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
import importlib
|
||||
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
importlib.reload(worker_module)
|
||||
|
||||
assert worker_module.SDK_METRICS_PORT == 9091
|
||||
|
||||
|
||||
def test_sdk_metrics_port_from_env():
|
||||
"""Test that SDK_METRICS_PORT is read from environment."""
|
||||
with patch.dict(os.environ, {'HTTP_SDK_METRICS_PORT': '8888'}):
|
||||
import importlib
|
||||
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
importlib.reload(worker_module)
|
||||
|
||||
assert worker_module.SDK_METRICS_PORT == 8888
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.POD_ID', 'test-pod-123')
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
def test_start_prometheus_server_success(
|
||||
mock_metrics, mock_start_http_server, mock_env_vars, mock_logger
|
||||
):
|
||||
"""Test successful Prometheus server startup."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
|
||||
|
||||
start_prometheus_server(mock_logger, metadata)
|
||||
|
||||
# Verify HTTP server started
|
||||
mock_start_http_server.assert_called_once_with(9090)
|
||||
|
||||
# Verify APP_UP metric was set to 1
|
||||
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
|
||||
mock_app_up.set.assert_called_once_with(1)
|
||||
mock_logger.custom_info.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
def test_start_prometheus_server_custom_port(mock_metrics, mock_start_http_server, mock_logger):
|
||||
"""Test Prometheus server startup with custom port."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
with patch.dict(os.environ, {'HTTP_METRICS_PORT': '8080', 'POD_ID': 'custom-pod'}):
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
metadata: dict[str, str | None] = {'pod_id': 'custom-pod', 'workflow_name': 'train_model'}
|
||||
|
||||
start_prometheus_server(mock_logger, metadata)
|
||||
|
||||
mock_start_http_server.assert_called_once_with(8080)
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.os._exit')
|
||||
def test_start_prometheus_server_failure(
|
||||
mock_exit, mock_metrics, mock_start_http_server, mock_env_vars, mock_logger
|
||||
):
|
||||
"""Test Prometheus server startup failure."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
mock_start_http_server.side_effect = OSError('Port already in use')
|
||||
|
||||
metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
|
||||
|
||||
start_prometheus_server(mock_logger, metadata)
|
||||
|
||||
# Verify exit was called with code 1 e log crítico emitido
|
||||
mock_exit.assert_called_once_with(1)
|
||||
mock_logger.custom_critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', 'model-manager-worker')
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_successful_startup(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_env_vars,
|
||||
mock_logger,
|
||||
mock_temporal_client,
|
||||
mock_worker,
|
||||
mock_notification_handler,
|
||||
mock_activities,
|
||||
):
|
||||
"""Test successful main() execution until workers start."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(
|
||||
side_effect=asyncio.CancelledError()
|
||||
) # Simulate interruption
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
# Run main() and expect it to exit due to CancelledError
|
||||
with pytest.raises(SystemExit) as exc_info:
|
||||
await main()
|
||||
|
||||
assert exc_info.value.code == 1
|
||||
|
||||
# Verify all initialization steps were called
|
||||
mock_get_logger.assert_called_once()
|
||||
mock_start_prometheus.assert_called_once()
|
||||
mock_notification_handler_class.assert_called_once()
|
||||
mock_activities_class.assert_called_once()
|
||||
mock_client_class.connect.assert_called_once()
|
||||
# Agora são criados dois Workers: um para train_model-queue e outro para cleanup-queue
|
||||
assert mock_prepare_worker.call_count == 2
|
||||
|
||||
# Verify cleanup was performed
|
||||
mock_notification_handler.shutdown.assert_called_once()
|
||||
mock_activities.shutdown.assert_called_once()
|
||||
mock_app_up.set.assert_called_with(0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', 'model-manager-worker')
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_handles_exception(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_env_vars,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test main() handles exceptions and performs cleanup."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.shutdown = Mock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=RuntimeError('Worker failed'))
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
# Run main() and expect SystemExit
|
||||
with pytest.raises(SystemExit) as exc_info:
|
||||
await main()
|
||||
|
||||
assert exc_info.value.code == 1
|
||||
|
||||
# Verify error was logged
|
||||
mock_logger.custom_error.assert_called_once()
|
||||
assert 'Worker failed' in str(mock_logger.custom_error.call_args)
|
||||
|
||||
# Verify cleanup was performed
|
||||
mock_notification_handler.shutdown.assert_called_once()
|
||||
mock_activities.shutdown.assert_called_once()
|
||||
mock_app_up.set.assert_called_with(0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', 'model-manager-worker')
|
||||
@patch('model_manager.worker.worker.create_cleanup_schedule')
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_temporal_client_configuration(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_create_cleanup_schedule,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test that Temporal client is configured correctly."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
mock_create_cleanup_schedule.return_value = AsyncMock()
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
'TEMPORAL_HOST': 'temporal.example.com:7233',
|
||||
'TEMPORAL_NAMESPACE': 'production',
|
||||
'TEMPORAL_USE_TLS': 'true',
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'STORE_BASE_URL': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'STORE_OWNER': 'sientia',
|
||||
'STORE_REPO': 'model-library-store',
|
||||
},
|
||||
):
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.shutdown = Mock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
# Run main()
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
|
||||
# Verify Temporal client was configured with correct parameters
|
||||
mock_client_class.connect.assert_called_once_with(
|
||||
target_host='temporal.example.com:7233',
|
||||
namespace='production',
|
||||
runtime=mock_runtime,
|
||||
tls=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', 'model-manager-worker')
|
||||
@patch('model_manager.worker.worker.create_cleanup_schedule')
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_worker_configuration(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_create_cleanup_schedule,
|
||||
mock_env_vars,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test that prepare_worker is configured with correct workflows and activities."""
|
||||
from model_manager.worker.worker import main
|
||||
from model_manager.workflows.cleanup_files import CleanupFiles
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_create_cleanup_schedule.return_value = AsyncMock()
|
||||
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.update_experiment_run = Mock()
|
||||
mock_activities.load_model_metadata = Mock()
|
||||
mock_activities.validate_train_params = Mock()
|
||||
mock_activities.train_model = Mock()
|
||||
mock_activities.cleanup_resources = Mock()
|
||||
mock_activities.cleanup_temp_directories = Mock()
|
||||
mock_activities.shutdown = Mock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
|
||||
assert mock_prepare_worker.call_count == 2
|
||||
|
||||
train_call = mock_prepare_worker.call_args_list[0]
|
||||
assert train_call.kwargs['temporal_client'] is mock_client_instance
|
||||
assert train_call.kwargs['logger'] is mock_logger
|
||||
assert train_call.kwargs['main_workflow'] is TrainModel
|
||||
assert train_call.kwargs['other_workflows'] == []
|
||||
train_activities_list = train_call.kwargs['activities']
|
||||
assert mock_activities.update_experiment_run in train_activities_list
|
||||
assert mock_activities.load_model_metadata in train_activities_list
|
||||
assert mock_activities.validate_train_params in train_activities_list
|
||||
assert mock_activities.train_model in train_activities_list
|
||||
assert mock_activities.cleanup_resources in train_activities_list
|
||||
|
||||
cleanup_call = mock_prepare_worker.call_args_list[1]
|
||||
assert cleanup_call.kwargs['temporal_client'] is mock_client_instance
|
||||
assert cleanup_call.kwargs['logger'] is mock_logger
|
||||
assert cleanup_call.kwargs['main_workflow'] is CleanupFiles
|
||||
assert cleanup_call.kwargs['other_workflows'] == []
|
||||
assert cleanup_call.kwargs['activities'] == [mock_activities.cleanup_temp_directories]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', 'model-manager-worker')
|
||||
@patch('model_manager.worker.worker.create_cleanup_schedule')
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_schedule_creation_failure_does_not_stop_worker(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_create_cleanup_schedule,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test that schedule creation failure does not prevent worker startup."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
# Mock schedule creation to raise an exception (as coroutine)
|
||||
async def mock_schedule_error(*args, **kwargs):
|
||||
raise Exception('Schedule creation failed')
|
||||
|
||||
mock_create_cleanup_schedule.side_effect = mock_schedule_error
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.shutdown = Mock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
|
||||
mock_create_cleanup_schedule.assert_called_once()
|
||||
|
||||
schedule_error_logged = False
|
||||
for call in mock_logger.custom_error.call_args_list:
|
||||
if call[0] and 'Failed to configure cleanup schedule' in call[0][0]:
|
||||
schedule_error_logged = True
|
||||
break
|
||||
assert schedule_error_logged, 'Schedule creation error should be logged'
|
||||
|
||||
assert mock_prepare_worker.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.RUNTIME', None)
|
||||
@patch('model_manager.worker.worker.prepare_worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.ensure_runtime_directories')
|
||||
async def test_main_missing_runtime_uses_single_fallback(
|
||||
mock_ensure_runtime_directories,
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_prepare_worker,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test that main() uses single runtime fallback when RUNTIME is missing."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.shutdown = Mock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
|
||||
mock_prepare_worker.return_value = mock_worker_instance
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'single', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
'cache_ttl_seconds': None,
|
||||
}
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
|
||||
assert mock_prepare_worker.call_count == 2
|
||||
assert mock_prepare_worker.call_args_list[0].kwargs['runtime'] == 'single'
|
||||
assert mock_prepare_worker.call_args_list[1].kwargs['runtime'] == 'single'
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.asyncio.run')
|
||||
def test_main_entrypoint(mock_asyncio_run):
|
||||
"""Test the __main__ entrypoint."""
|
||||
# Import and execute the main block
|
||||
with patch.object(sys, 'argv', ['worker.py']):
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
# Simulate running the module
|
||||
worker_module.main = AsyncMock()
|
||||
|
||||
# This would normally be called by asyncio.run(main())
|
||||
# We just verify the pattern is correct
|
||||
assert callable(worker_module.main)
|
||||
|
||||
|
||||
def test_worker_module_docstring():
|
||||
"""Test that worker module has comprehensive documentation."""
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
assert worker_module.__doc__ is not None
|
||||
assert 'Temporal' in worker_module.__doc__
|
||||
assert 'worker' in worker_module.__doc__
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
def test_start_prometheus_server_prints_success(
|
||||
mock_metrics, mock_start_http_server, capsys, mock_env_vars, mock_logger
|
||||
):
|
||||
"""Test that start_prometheus_server prints success message."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
|
||||
|
||||
start_prometheus_server(mock_logger, metadata)
|
||||
|
||||
# Agora a mensagem é enviada via logger
|
||||
mock_logger.custom_info.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.os._exit')
|
||||
def test_start_prometheus_server_prints_failure(
|
||||
mock_exit, mock_metrics, mock_start_http_server, capsys, mock_env_vars, mock_logger
|
||||
):
|
||||
"""Test that start_prometheus_server prints failure message."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
mock_start_http_server.side_effect = Exception('Test error')
|
||||
|
||||
metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
|
||||
|
||||
start_prometheus_server(mock_logger, metadata)
|
||||
|
||||
# Agora o erro é logado via logger crítico
|
||||
mock_logger.custom_critical.assert_called_once()
|
||||
mock_exit.assert_called_once_with(1)
|
||||
30
tests/workflows/test_cleanup_files.py
Normal file
30
tests/workflows/test_cleanup_files.py
Normal file
@@ -0,0 +1,30 @@
|
||||
"""Unit tests for the CleanupFiles workflow."""
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.cleanup_files.workflow')
|
||||
async def test_cleanup_files_workflow(mock_workflow_module):
|
||||
"""Test the CleanupFiles workflow."""
|
||||
from model_manager.runtime_paths import REPORTS_TEMP_DIR
|
||||
from model_manager.workflows.cleanup_files import CleanupFiles
|
||||
|
||||
# Mock execute_activity_method
|
||||
mock_workflow_module.execute_activity_method = AsyncMock()
|
||||
|
||||
# Instantiate and run the workflow
|
||||
workflow_instance = CleanupFiles()
|
||||
await workflow_instance.run({})
|
||||
|
||||
# Verify that the activities were called with the correct parameters
|
||||
calls = mock_workflow_module.execute_activity_method.call_args_list
|
||||
assert len(calls) == 1
|
||||
|
||||
# Check cleanup_temp_directories call
|
||||
local_call_args = calls[0][0][1]
|
||||
assert local_call_args['temp_path'] == REPORTS_TEMP_DIR
|
||||
assert local_call_args['metadata']['workflow_name'] == 'cleanup_files'
|
||||
assert 'pod_id' in local_call_args['metadata']
|
||||
318
tests/workflows/test_train_model.py
Normal file
318
tests/workflows/test_train_model.py
Normal file
@@ -0,0 +1,318 @@
|
||||
"""Unit tests for TrainModel workflow."""
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from temporalio.exceptions import ApplicationError
|
||||
|
||||
from model_manager.utils.models.experiment_status import ExperimentStatus
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_train_params():
|
||||
"""Minimal mock TrainModelParams."""
|
||||
params = Mock(spec=TrainModelParams)
|
||||
params.experiment_run_id = 123
|
||||
params.bucket_name = 'test-bucket'
|
||||
params.file_name = 'test-file.csv'
|
||||
params.target_variable = 'target'
|
||||
params.variable_columns = ['var1', 'var2']
|
||||
return params
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_input_data():
|
||||
"""Sample workflow input (IDs normalized in run())."""
|
||||
return {
|
||||
'experiment_run_id': 123,
|
||||
'target_variable': 'target',
|
||||
'variable_columns': ['var1', 'var2'],
|
||||
'train_size': 80,
|
||||
'bucket_name': 'test-bucket',
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'model_name': 'Linear Regression',
|
||||
'model_type': 'linear_regression',
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'val_file_name': None,
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_success():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
wf = TrainModel()
|
||||
assert wf._validate_experiment_run_id({'experiment_run_id': 123}) == 123
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_string_numeric():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
wf = TrainModel()
|
||||
assert wf._validate_experiment_run_id({'experiment_run_id': '123'}) == 123
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_missing():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required'):
|
||||
TrainModel()._validate_experiment_run_id({})
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_invalid_type():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
with pytest.raises(ValueError, match='must be an integer or numeric string'):
|
||||
TrainModel()._validate_experiment_run_id({'experiment_run_id': 'not_int'})
|
||||
|
||||
|
||||
def test_extract_error_message_simple():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
assert TrainModel()._extract_error_message(ValueError('x')) == 'x'
|
||||
|
||||
|
||||
def test_extract_error_message_with_cause():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
cause = ValueError('Root')
|
||||
exc = RuntimeError('Outer')
|
||||
exc.__cause__ = cause
|
||||
msg = TrainModel()._extract_error_message(exc)
|
||||
assert 'Outer' in msg and 'Root' in msg
|
||||
|
||||
|
||||
def test_extract_error_message_empty_message():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
out = TrainModel()._extract_error_message(ValueError(''))
|
||||
assert 'ValueError' in out
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_success(mock_wf, sample_input_data, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'experiment_run_id': 123, 'model_metadata': {}},
|
||||
mock_train_params,
|
||||
None,
|
||||
]
|
||||
)
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
out = await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
assert out is mock_train_params
|
||||
assert mock_wf.execute_activity_method.call_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_load_fails(mock_wf, sample_input_data):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[ValueError('load'), None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(ValueError, match='load'):
|
||||
await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_success(mock_wf, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'rid', 'run_dir': '/tmp/r'}
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[tr, None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
out = await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
assert out == tr
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_logs_when_db_update_fails(mock_wf, sample_input_data):
|
||||
"""If persisting ORCHESTRATOR_VALIDATION_ERROR fails, workflow logs a warning."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[ValueError('validation'), RuntimeError('db')],
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(ValueError, match='validation'):
|
||||
await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_logs_when_error_status_persist_fails(mock_wf, mock_train_params):
|
||||
"""If persisting TRAINING_ERROR fails, workflow logs a warning."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[RuntimeError('train'), RuntimeError('db')],
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(RuntimeError, match='train'):
|
||||
await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_failure_updates_db(mock_wf, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[RuntimeError('fail'), None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(RuntimeError, match='fail'):
|
||||
await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
err_call = mock_wf.execute_activity_method.call_args_list[1]
|
||||
assert err_call[0][1]['status'] == ExperimentStatus.TRAINING_ERROR
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_cleanup_resources(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(return_value=None)
|
||||
meta = {'metadata': {'pod_id': 'p'}}
|
||||
await TrainModel()._cleanup_resources('/tmp/x', meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_cleanup_resources_none_skips(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
await TrainModel()._cleanup_resources(None, {'metadata': {}})
|
||||
mock_wf.execute_activity_method.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_success_six_activities(mock_wf, sample_input_data, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'i', 'run_dir': '/tmp/t'}
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
tr,
|
||||
None,
|
||||
None,
|
||||
]
|
||||
)
|
||||
result = await TrainModel().run(sample_input_data)
|
||||
assert result == tr
|
||||
assert mock_wf.execute_activity_method.call_count == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_validation_error(mock_wf, sample_input_data):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[ValueError('bad'), None])
|
||||
with pytest.raises(ValueError, match='bad'):
|
||||
await TrainModel().run(sample_input_data)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_cleanup_failure_does_not_fail_workflow(
|
||||
mock_wf, sample_input_data, mock_train_params
|
||||
):
|
||||
"""After successful training, cleanup failure is logged, workflow still returns result."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'i', 'run_dir': '/tmp/t'}
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
tr,
|
||||
None,
|
||||
RuntimeError('cleanup'),
|
||||
]
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
out = await TrainModel().run(sample_input_data)
|
||||
assert out == tr
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_training_failure_skips_cleanup_activity(
|
||||
mock_wf, sample_input_data, mock_train_params
|
||||
):
|
||||
"""When train_model raises, train_result stays None and cleanup activity is not scheduled."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
RuntimeError('train failed'),
|
||||
]
|
||||
)
|
||||
with pytest.raises(RuntimeError, match='train failed'):
|
||||
await TrainModel().run(sample_input_data)
|
||||
# validate (3) + train activity (1) + TRAINING_ERROR DB update (1); no cleanup (6th) when train_result is unset
|
||||
assert mock_wf.execute_activity_method.call_count == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_missing_experiment_run_id(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.logger = Mock()
|
||||
with pytest.raises(ApplicationError, match='experiment_run_id is required'):
|
||||
await TrainModel().run({})
|
||||
|
||||
|
||||
def test_module_constants():
|
||||
from model_manager.workflows.train_model import (
|
||||
TIMEOUT_DELETE_FILE,
|
||||
TIMEOUT_TRAIN_MODEL,
|
||||
TIMEOUT_VALIDATE_PARAMS,
|
||||
database_retry_policy,
|
||||
network_retry_policy,
|
||||
no_retry_policy,
|
||||
)
|
||||
|
||||
assert isinstance(TIMEOUT_VALIDATE_PARAMS, int)
|
||||
assert no_retry_policy.maximum_attempts == 1
|
||||
assert network_retry_policy.maximum_attempts == 5
|
||||
assert database_retry_policy.maximum_attempts == 5
|
||||
assert TIMEOUT_TRAIN_MODEL == 2700
|
||||
assert TIMEOUT_DELETE_FILE == 120
|
||||
Reference in New Issue
Block a user