Code import - branch release/SIENTIAPDE-1645
This commit is contained in:
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()})
|
||||
Reference in New Issue
Block a user