Code import - branch release/SIENTIAPDE-1645

This commit is contained in:
2026-08-05 13:53:37 +00:00
commit d481e0acff
116 changed files with 92848 additions and 0 deletions

View File

View 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)

View 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()

View 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__()

View 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()})