SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.

This commit is contained in:
Bruno Domingues
2025-10-22 15:37:56 -03:00
parent f2a1c88ff3
commit 5789a13023
31 changed files with 37878 additions and 6097 deletions

View File

@@ -1,535 +0,0 @@
from unittest.mock import ANY, MagicMock, patch
from pytest import mark
from model_manager.activities.activities import Activities
from model_manager.activities.experiment_tracking import ExperimentTracking
from model_manager.activities.mlflow import MLFlow
from model_manager.activities.training import Training
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___init__(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
assert isinstance(activities, Activities)
assert isinstance(activities, ExperimentTracking)
assert isinstance(activities, MLFlow)
assert isinstance(activities, Training)
mock_experiment_tracking_init.assert_called_once_with(
ANY,
host=postgres_config['host'],
port=postgres_config['port'],
user=postgres_config['user'],
password=postgres_config['password'],
dbname=postgres_config['dbname'],
min_connections=postgres_config['min_connections'],
max_connections=postgres_config['max_connections'],
logger=logger,
notification_handler=notification_handler,
)
mock_mlflow_init.assert_called_once_with(
ANY,
mlflow_host=mlflow_config['host'],
mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'],
logger=logger,
notification_handler=notification_handler,
)
mock_minio_init.assert_called_once_with(
ANY,
endpoint_url=minio_config['endpoint_url'],
access_key=minio_config['access_key'],
secret_key=minio_config['secret_key'],
region=minio_config['region'],
use_ssl=minio_config['use_ssl'],
max_retry_attempts=minio_config['max_retry_attempts'],
retry_mode=minio_config['retry_mode'],
connect_timeout=minio_config['connect_timeout'],
read_timeout=minio_config['read_timeout'],
logger=logger,
notification_handler=notification_handler,
)
mock_training_init.assert_called_once_with(
ANY, logger=logger, notification_handler=notification_handler
)
@mark.asyncio
@patch('model_manager.activities.activities.ExperimentTracking', return_value=MagicMock())
@patch('model_manager.activities.activities.MLFlow', return_value=MagicMock())
async def test_shutdown(_mock_mlflow_init, mock_experiment_tracking_init):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
await activities.shutdown()
mock_experiment_tracking_init.close.assert_called_once()
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___with_engine(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ calls parent destructor when engine attribute exists."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Add engine attribute to simulate Postgres initialization
activities.engine = MagicMock()
# Create a mock __del__ that will be detected by hasattr
mock_parent_del = MagicMock()
# Patch both the class and the instance to ensure super().__del__ exists and is callable
with patch.object(ExperimentTracking, '__del__', mock_parent_del, create=True):
# Trigger __del__
activities.__del__()
# Verify parent __del__ was called
mock_parent_del.assert_called_once()
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___without_engine(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ does not call parent destructor when engine attribute is missing."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Ensure engine attribute does NOT exist
if hasattr(activities, 'engine'):
delattr(activities, 'engine')
# Mock super().__del__ to track if it's called
with patch.object(ExperimentTracking, '__del__', MagicMock()) as mock_parent_del:
# Trigger __del__
activities.__del__()
# Verify parent __del__ was NOT called
mock_parent_del.assert_not_called()
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___handles_exception_gracefully(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ handles exceptions from parent destructor gracefully."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Add engine attribute
activities.engine = MagicMock()
# Mock super().__del__ to raise an exception
mock_parent_del = MagicMock(side_effect=RuntimeError('Cleanup failed'))
with patch.object(ExperimentTracking, '__del__', mock_parent_del):
# Trigger __del__ - should not raise exception
try:
activities.__del__()
# Test passes if no exception is raised
except Exception as e:
# Test fails if exception propagates
raise AssertionError(f'__del__ should not raise exception, but raised: {e}') from e
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___when_parent_has_no_del(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ handles case when parent class has no __del__ method."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Add engine attribute
activities.engine = MagicMock()
# Remove __del__ from parent to simulate it not existing
with patch.object(ExperimentTracking, '__del__', create=False):
# Trigger __del__ - should not raise exception
try:
activities.__del__()
# Test passes if no exception is raised
except Exception as e:
# Test fails if exception propagates
raise AssertionError(
f'__del__ should handle missing parent __del__, but raised: {e}'
) from e
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___calls_super_successfully(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ successfully calls super().__del__() when it exists - covers line 118."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
# Mock all parent __init__ methods to return None
mock_experiment_tracking_init.return_value = None
mock_mlflow_init.return_value = None
mock_minio_init.return_value = None
mock_training_init.return_value = None
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Add engine attribute to simulate Postgres initialization
activities.engine = MagicMock()
# Track if super().__del__() was actually called
super_del_called = []
def mock_super_del(self):
"""Mock parent __del__ that tracks when it's called."""
super_del_called.append(True)
# Patch ExperimentTracking.__del__ to exist and be callable
with patch.object(ExperimentTracking, '__del__', mock_super_del, create=True):
# Trigger __del__ - this should execute line 118: super().__del__()
activities.__del__()
# Verify that super().__del__() was actually called (line 118 executed)
assert len(super_del_called) == 1, 'super().__del__() should have been called once'
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___del___when_super_has_no_del_method(
mock_training_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
"""Test __del__ handles case when hasattr(super(), '__del__') returns False - covers line 118 false branch."""
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
# Mock all __init__ methods to return None
mock_experiment_tracking_init.return_value = None
mock_mlflow_init.return_value = None
mock_minio_init.return_value = None
mock_training_init.return_value = None
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
# Add engine attribute to pass the first hasattr check (line 115)
activities.engine = MagicMock()
# Create a mock class without __del__ method to simulate super() not having __del__
class MockSuperWithoutDel:
"""Mock class that explicitly does not have __del__ method."""
pass
# Patch super() to return an instance that doesn't have __del__
mock_super_instance = MockSuperWithoutDel()
with patch('builtins.super', return_value=mock_super_instance):
# Trigger __del__ - should handle the case when hasattr(super(), '__del__') is False
try:
activities.__del__()
# Test passes - the false branch of line 118 was executed without error
except Exception as e:
# Test fails if exception propagates
raise AssertionError(
f'__del__ should handle super() without __del__ method, but raised: {e}'
) from e

View File

@@ -1,564 +0,0 @@
"""Unit tests for ExperimentTracking activity."""
from unittest.mock import AsyncMock, MagicMock, patch
from pytest import mark, raises
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_status_success(mock_postgres_init):
"""Test successful status update."""
mock_postgres_init.return_value = None
# Create instance
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Mock methods
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
# Test data
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS,
'status': 'TRAINING_SUCCESS',
}
# Execute
await tracking.update_experiment_run(input_data)
# Assertions
tracking.info.assert_called()
tracking.execute_query.assert_called_once()
call_args = tracking.execute_query.call_args
assert 'UPDATE experiment_run' in call_args[0][0]
assert 'SET status = %s' in call_args[0][0]
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_status_with_error_success(mock_postgres_init):
"""Test successful status update with error message."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS_WITH_ERROR,
'status': 'TRAINING_ERROR',
'error_message': 'Model training failed due to insufficient data',
}
await tracking.update_experiment_run(input_data)
tracking.info.assert_called()
tracking.execute_query.assert_called_once()
call_args = tracking.execute_query.call_args
assert 'UPDATE experiment_run' in call_args[0][0]
assert 'SET status = %s, error_message = %s' in call_args[0][0]
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_error_message_truncation(mock_postgres_init):
"""Test that error messages longer than 1024 chars are truncated."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
# Create error message longer than 1024 characters
long_error = 'A' * 2000
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS_WITH_ERROR,
'status': 'TRAINING_ERROR',
'error_message': long_error,
}
await tracking.update_experiment_run(input_data)
# Check that error message was truncated to 1024 chars
call_args = tracking.execute_query.call_args
query_params = call_args[0][1]
assert len(query_params[1]) == 1024 # error_message is second param
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_model_saved_success(mock_postgres_init):
"""Test successful model saved update."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': UpdateType.MODEL_SAVED,
'run_name': 'experiment-model-123',
'status': 'MLFLOW_SENT',
}
await tracking.update_experiment_run(input_data)
tracking.info.assert_called()
tracking.execute_query.assert_called_once()
call_args = tracking.execute_query.call_args
assert 'UPDATE experiment_run' in call_args[0][0]
assert 'SET run_name = %s, status = %s' in call_args[0][0]
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_missing_status_raises_error(mock_postgres_init):
"""Test that missing status parameter raises ValueError."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS,
# Missing 'status' parameter
}
with raises(RuntimeError, match='Error updating experiment run'):
await tracking.update_experiment_run(input_data)
tracking.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_missing_error_message_raises_error(mock_postgres_init):
"""Test that missing error_message parameter raises ValueError."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS_WITH_ERROR,
'status': 'TRAINING_ERROR',
# Missing 'error_message' parameter
}
with raises(RuntimeError, match='Error updating experiment run'):
await tracking.update_experiment_run(input_data)
tracking.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_missing_run_name_raises_error(mock_postgres_init):
"""Test that missing run_name parameter raises ValueError."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': UpdateType.MODEL_SAVED,
# Missing 'run_name' parameter
}
with raises(RuntimeError, match='Error updating experiment run'):
await tracking.update_experiment_run(input_data)
tracking.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_invalid_update_type_raises_error(mock_postgres_init):
"""Test that invalid update_type raises ValueError."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {},
'experiment_run_id': 456,
'update_type': 'INVALID_TYPE',
'status': 'TRAINING_SUCCESS',
}
with raises(RuntimeError, match='Error updating experiment run'):
await tracking.update_experiment_run(input_data)
tracking.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_no_rows_updated_raises_error(mock_postgres_init):
"""Test that zero rows updated raises ValueError."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(return_value={'rowcount': 0})
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {},
'experiment_run_id': 999, # Non-existent ID
'update_type': UpdateType.STATUS,
'status': 'TRAINING_SUCCESS',
}
with raises(RuntimeError, match='Error updating experiment run'):
await tracking.update_experiment_run(input_data)
tracking.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
async def test_update_experiment_run_sends_notification_on_error(mock_postgres_init):
"""Test that notification is sent when update fails."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
tracking.info = MagicMock()
tracking.execute_query = AsyncMock(side_effect=Exception('Database error'))
tracking.send_notification = MagicMock()
tracking.error = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'update_type': UpdateType.STATUS,
'status': 'TRAINING_SUCCESS',
}
with raises(RuntimeError):
await tracking.update_experiment_run(input_data)
# Verify notification was sent
tracking.send_notification.assert_called_once()
call_args = tracking.send_notification.call_args
assert call_args[1]['notification_id'] == 'UPDATE_EXPERIMENT_RUN_ERROR'
assert call_args[1]['metadata'] == {'workflow_id': 'test-123'}
def test_update_type_enum_values():
"""Test UpdateType enum has correct values."""
assert UpdateType.STATUS == 'status'
assert UpdateType.STATUS_WITH_ERROR == 'status_with_error'
assert UpdateType.MODEL_SAVED == 'model_saved'
# Tests for __del__ method - 100% coverage
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
def test___del___with_engine_and_parent_del_exists(mock_postgres_init):
"""Test __del__ calls parent destructor when engine exists and parent has __del__."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Add engine attribute to simulate Postgres initialization
tracking.engine = MagicMock()
# Track if parent __del__ was called
parent_del_called = []
def mock_parent_del(self):
"""Mock parent __del__ that tracks when it's called."""
parent_del_called.append(True)
# Patch parent class to have __del__ method
with patch.object(type(tracking).__bases__[0], '__del__', mock_parent_del, create=True):
# Trigger __del__ - should call parent __del__ (line 103)
tracking.__del__()
# Verify parent __del__ was called (line 103 executed)
assert len(parent_del_called) == 1, 'Parent __del__ should have been called once'
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
def test___del___without_engine(mock_postgres_init):
"""Test __del__ does not call parent destructor when engine attribute is missing."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Ensure engine attribute does not exist
if hasattr(tracking, 'engine'):
delattr(tracking, 'engine')
# Mock parent __del__ to track if it's called
mock_parent_del = MagicMock()
with patch.object(type(tracking).__bases__[0], '__del__', mock_parent_del, create=True):
# Trigger __del__ - should NOT call parent __del__ (line 100 is False)
tracking.__del__()
# Verify parent __del__ was NOT called
mock_parent_del.assert_not_called()
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
def test___del___when_parent_has_no_del_method(mock_postgres_init):
"""Test __del__ handles case when parent class has no __del__ method - covers line 102 false branch."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Add engine attribute to pass the first hasattr check (line 100)
tracking.engine = MagicMock()
# Create a mock class without __del__ method
class MockSuperWithoutDel:
"""Mock class that explicitly does not have __del__ method."""
pass
# Patch super() to return an instance without __del__
mock_super_instance = MockSuperWithoutDel()
with patch('builtins.super', return_value=mock_super_instance):
# Trigger __del__ - should handle the case when hasattr(super(), '__del__') is False (line 102)
try:
tracking.__del__()
# Test passes - the false branch of line 102 was executed without error
except Exception as e:
raise AssertionError(
f'__del__ should handle super() without __del__ method, but raised: {e}'
) from e
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
def test___del___handles_exception_from_parent_del(mock_postgres_init):
"""Test __del__ handles exceptions from parent destructor gracefully - covers line 104."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Add engine attribute
tracking.engine = MagicMock()
# Mock parent __del__ to raise an exception
mock_parent_del = MagicMock(side_effect=RuntimeError('Cleanup failed'))
with patch.object(type(tracking).__bases__[0], '__del__', mock_parent_del, create=True):
# Trigger __del__ - should catch exception and not propagate it (line 104-106)
try:
tracking.__del__()
# Test passes if no exception is raised
except Exception as e:
raise AssertionError(
f'__del__ should handle exceptions gracefully, but raised: {e}'
) from e
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
def test___del___handles_attribute_error_from_parent_del(mock_postgres_init):
"""Test __del__ handles AttributeError from parent destructor - covers line 104 exception handling."""
mock_postgres_init.return_value = None
tracking = ExperimentTracking(
host='localhost',
port=5432,
user='test',
password='test',
dbname='test',
min_connections=1,
max_connections=10,
logger=MagicMock(),
notification_handler=MagicMock(),
)
# Add engine attribute
tracking.engine = MagicMock()
# Mock parent __del__ to raise AttributeError
mock_parent_del = MagicMock(side_effect=AttributeError('engine not found'))
with patch.object(type(tracking).__bases__[0], '__del__', mock_parent_del, create=True):
# Trigger __del__ - should catch AttributeError and not propagate it
try:
tracking.__del__()
# Test passes if no exception is raised
except Exception as e:
raise AssertionError(
f'__del__ should handle AttributeError gracefully, but raised: {e}'
) from e

View File

@@ -1,328 +0,0 @@
from io import BytesIO
from unittest.mock import MagicMock, patch
from pytest import fixture, mark, raises
from model_manager.activities.minio import MinIO
@patch('model_manager.activities.minio.boto3.client')
def test___init__(mock_boto3_client):
"""Test MinIO initialization with correct configuration."""
mock_client = MagicMock()
mock_boto3_client.return_value = mock_client
logger = MagicMock()
notification_handler = MagicMock()
minio = MinIO(
endpoint_url='http://localhost:9000',
access_key='minioadmin',
secret_key='minioadmin',
region='us-east-1',
use_ssl=False,
max_retry_attempts=3,
retry_mode='adaptive',
connect_timeout=10,
read_timeout=60,
logger=logger,
notification_handler=notification_handler,
)
assert minio.endpoint_url == 'http://localhost:9000'
assert minio.access_key == 'minioadmin'
assert minio.secret_key == 'minioadmin'
assert minio.region == 'us-east-1'
assert minio.use_ssl is False
assert minio.max_retry_attempts == 3
assert minio.retry_mode == 'adaptive'
assert minio.connect_timeout == 10
assert minio.read_timeout == 60
# Verify boto3 client was created with correct parameters
mock_boto3_client.assert_called_once()
call_kwargs = mock_boto3_client.call_args[1]
assert call_kwargs['endpoint_url'] == 'http://localhost:9000'
assert call_kwargs['aws_access_key_id'] == 'minioadmin'
assert call_kwargs['aws_secret_access_key'] == 'minioadmin'
assert call_kwargs['use_ssl'] is False
@patch('model_manager.activities.minio.boto3.client')
def test___init___failure(mock_boto3_client):
"""Test MinIO initialization failure handling."""
mock_boto3_client.side_effect = Exception('Connection failed')
logger = MagicMock()
notification_handler = MagicMock()
with raises(ConnectionError, match='Failed to initialize MinIO client'):
MinIO(
endpoint_url='http://localhost:9000',
access_key='minioadmin',
secret_key='minioadmin',
region='us-east-1',
use_ssl=False,
max_retry_attempts=3,
retry_mode='adaptive',
connect_timeout=10,
read_timeout=60,
logger=logger,
notification_handler=notification_handler,
)
@fixture
@patch('model_manager.activities.minio.boto3.client')
def minio(mock_boto3_client):
"""Fixture to create a MinIO instance for testing."""
mock_client = MagicMock()
mock_boto3_client.return_value = mock_client
logger = MagicMock()
notification_handler = MagicMock()
minio_instance = MinIO(
endpoint_url='http://localhost:9000',
access_key='minioadmin',
secret_key='minioadmin',
region='us-east-1',
use_ssl=False,
max_retry_attempts=3,
retry_mode='adaptive',
connect_timeout=10,
read_timeout=60,
logger=logger,
notification_handler=notification_handler,
)
minio_instance.send_notification = MagicMock()
minio_instance.minio_client = mock_client
return minio_instance
metadata = {
'metadata': {
'workflow_name': 'test_workflow',
'model_name': 'test_model',
'model_id': 'test_model_id',
}
}
@mark.asyncio
async def test_fetch_file_from_minio_success(minio):
"""Test successful file fetch from MinIO."""
# Arrange
test_content = b'test file content'
mock_response = {'Body': MagicMock()}
mock_response['Body'].__enter__ = MagicMock(
return_value=MagicMock(read=MagicMock(return_value=test_content))
)
mock_response['Body'].__exit__ = MagicMock(return_value=None)
minio.minio_client.get_object.return_value = mock_response
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'test-file.txt',
}
# Act
result = await minio.fetch_file_from_minio(input_data)
# Assert
assert isinstance(result, BytesIO)
result.seek(0)
assert result.read() == test_content
minio.minio_client.get_object.assert_called_once_with(Bucket='test-bucket', Key='test-file.txt')
@mark.asyncio
async def test_fetch_file_from_minio_file_not_found(minio):
"""Test file fetch when file doesn't exist."""
# Arrange
minio.minio_client.get_object.side_effect = Exception(
'NoSuchKey: The specified key does not exist'
)
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'nonexistent.txt',
}
# Act & Assert
with raises(OSError, match='Error fetching file from MinIO'):
await minio.fetch_file_from_minio(input_data)
# Verify notification was sent
minio.send_notification.assert_called_once()
call_kwargs = minio.send_notification.call_args[1]
assert call_kwargs['notification_id'] == 'FETCH_FILE_FROM_MINIO_ERROR'
assert call_kwargs['block'] == 'fetch_file_from_minio'
@mark.asyncio
async def test_fetch_file_from_minio_network_error(minio):
"""Test file fetch with network error."""
# Arrange
minio.minio_client.get_object.side_effect = Exception('Network timeout')
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'test-file.txt',
}
# Act & Assert
with raises(OSError, match='Error fetching file from MinIO'):
await minio.fetch_file_from_minio(input_data)
minio.send_notification.assert_called_once()
@mark.asyncio
async def test_delete_file_from_minio_success(minio):
"""Test successful file deletion from MinIO."""
# Arrange
minio.minio_client.delete_object.return_value = None
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'test-file.txt',
}
# Act
result = await minio.delete_file_from_minio(input_data)
# Assert
assert result is None
minio.minio_client.delete_object.assert_called_once_with(
Bucket='test-bucket', Key='test-file.txt'
)
@mark.asyncio
async def test_delete_file_from_minio_idempotent(minio):
"""Test that delete is idempotent (no error if file doesn't exist)."""
# Arrange
# MinIO delete_object is idempotent - no error if file doesn't exist
minio.minio_client.delete_object.return_value = None
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'nonexistent.txt',
}
# Act
result = await minio.delete_file_from_minio(input_data)
# Assert
assert result is None
minio.minio_client.delete_object.assert_called_once()
@mark.asyncio
async def test_delete_file_from_minio_access_denied(minio):
"""Test file deletion with access denied error."""
# Arrange
minio.minio_client.delete_object.side_effect = Exception('AccessDenied: Access Denied')
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'test-file.txt',
}
# Act & Assert
with raises(OSError, match='Error deleting file from MinIO'):
await minio.delete_file_from_minio(input_data)
# Verify notification was sent
minio.send_notification.assert_called_once()
call_kwargs = minio.send_notification.call_args[1]
assert call_kwargs['notification_id'] == 'DELETE_FILE_FROM_MINIO_ERROR'
assert call_kwargs['block'] == 'delete_file_from_minio'
@mark.asyncio
async def test_delete_file_from_minio_network_error(minio):
"""Test file deletion with network error."""
# Arrange
minio.minio_client.delete_object.side_effect = Exception('Connection timeout')
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'test-file.txt',
}
# Act & Assert
with raises(OSError, match='Error deleting file from MinIO'):
await minio.delete_file_from_minio(input_data)
minio.send_notification.assert_called_once()
@mark.asyncio
async def test_fetch_file_from_minio_large_file(minio):
"""Test fetching a large file from MinIO."""
# Arrange
# Simulate a 10MB file
large_content = b'x' * (10 * 1024 * 1024)
mock_response = {'Body': MagicMock()}
mock_response['Body'].__enter__ = MagicMock(
return_value=MagicMock(read=MagicMock(return_value=large_content))
)
mock_response['Body'].__exit__ = MagicMock(return_value=None)
minio.minio_client.get_object.return_value = mock_response
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'large-file.bin',
}
# Act
result = await minio.fetch_file_from_minio(input_data)
# Assert
assert isinstance(result, BytesIO)
result.seek(0)
assert len(result.read()) == 10 * 1024 * 1024
@mark.asyncio
async def test_fetch_file_from_minio_empty_file(minio):
"""Test fetching an empty file from MinIO."""
# Arrange
empty_content = b''
mock_response = {'Body': MagicMock()}
mock_response['Body'].__enter__ = MagicMock(
return_value=MagicMock(read=MagicMock(return_value=empty_content))
)
mock_response['Body'].__exit__ = MagicMock(return_value=None)
minio.minio_client.get_object.return_value = mock_response
input_data = {
'metadata': metadata['metadata'],
'bucket_name': 'test-bucket',
'file_name': 'empty-file.txt',
}
# Act
result = await minio.fetch_file_from_minio(input_data)
# Assert
assert isinstance(result, BytesIO)
result.seek(0)
assert result.read() == b''

View File

@@ -1,447 +0,0 @@
from unittest.mock import ANY, MagicMock, patch
import pytest
from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel
from model_manager.activities.mlflow import MLFlow
@patch('model_manager.activities.mlflow.MLFlowRepository')
def test___init__(mock_mlflow_repository):
mlflow = MLFlow(
mlflow_host='http://localhost',
mlflow_port=5000,
mlflow_username='admin',
mlflow_password='admin',
logger=MagicMock(),
notification_handler=MagicMock(),
)
assert mlflow.mlflow_host == 'http://localhost'
assert mlflow.mlflow_port == 5000
assert mlflow.mlflow_username == 'admin'
assert mlflow.mlflow_password == 'admin'
mock_mlflow_repository.assert_called_once_with('http://localhost:5000', 'admin', 'admin', ANY)
@fixture
@patch('model_manager.activities.mlflow.MLFlowRepository')
def mlflow(mock_mlflow_repository):
mlflow = MLFlow(
mlflow_host='http://localhost:5000',
mlflow_port=5000,
mlflow_username='admin',
mlflow_password='admin',
logger=MagicMock(),
notification_handler=MagicMock(),
)
mlflow.send_notification = MagicMock()
return mlflow
metadata = {
'metadata': {
'model_id': 'test_model',
'model_name': 'test_model',
'workflow_name': 'test_workflow',
'schema_name': 'test_schedule',
},
}
@mark.asyncio
async def test_save_model_success(mlflow):
"""Test save_model successfully saves model and artifacts to MLflow."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
# Mock train result
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'test_experiment'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
train_result.run_name = None # Will be set by get_next_run_name
# Mock repository methods
mlflow.model_monitoring_repository.get_next_run_name.return_value = 'test_experiment-1'
mlflow.model_monitoring_repository.generate_artifacts.return_value = train_result
mlflow.model_monitoring_repository.save_run.return_value = None
input_data = {
**metadata,
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify repository methods were called
mlflow.model_monitoring_repository.get_next_run_name.assert_called_once_with('test_experiment')
mlflow.model_monitoring_repository.generate_artifacts.assert_called_once_with(train_result)
mlflow.model_monitoring_repository.save_run.assert_called_once_with(train_result)
# Verify response - now returns TrainModelResult directly
assert response == train_result
assert response.run_name == 'test_experiment-1'
assert train_result.run_name == 'test_experiment-1'
@mark.asyncio
async def test_save_model_get_next_run_name_error(mlflow):
"""Test save_model handles error during get_next_run_name."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'test_experiment'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
# Mock error in get_next_run_name
mlflow.model_monitoring_repository.get_next_run_name.side_effect = Exception(
'MLflow connection error'
)
input_data = {
**metadata,
'train_result': train_result,
}
# Call the method - should raise exception
with pytest.raises(Exception, match='MLflow connection error'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='SAVE_MODEL_ERROR',
message=ANY,
block='save_model',
level=NotificationLevel.ERROR,
attachment_content=ANY,
)
@mark.asyncio
async def test_save_model_generate_artifacts_error(mlflow):
"""Test save_model handles error during generate_artifacts."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'test_experiment'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
# Mock successful get_next_run_name but error in generate_artifacts
mlflow.model_monitoring_repository.get_next_run_name.return_value = 'test_experiment-1'
mlflow.model_monitoring_repository.generate_artifacts.side_effect = FileNotFoundError(
'Reports directory does not exist'
)
input_data = {
**metadata,
'train_result': train_result,
}
# Call the method - should raise exception
with pytest.raises(FileNotFoundError, match='Reports directory does not exist'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='SAVE_MODEL_ERROR',
message=ANY,
block='save_model',
level=NotificationLevel.ERROR,
attachment_content=ANY,
)
@mark.asyncio
async def test_save_model_save_run_error(mlflow):
"""Test save_model handles error during save_run."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'test_experiment'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
# Mock successful get_next_run_name and generate_artifacts but error in save_run
mlflow.model_monitoring_repository.get_next_run_name.return_value = 'test_experiment-1'
mlflow.model_monitoring_repository.generate_artifacts.return_value = train_result
mlflow.model_monitoring_repository.save_run.side_effect = ValueError(
'One or more metrics (MSE, R2, MAE) are None'
)
input_data = {
**metadata,
'train_result': train_result,
}
# Call the method - should raise exception
with pytest.raises(ValueError, match=r'One or more metrics \(MSE, R2, MAE\) are None'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='SAVE_MODEL_ERROR',
message=ANY,
block='save_model',
level=NotificationLevel.ERROR,
attachment_content=ANY,
)
@mark.asyncio
async def test_save_model_missing_metadata(mlflow):
"""Test save_model handles missing metadata gracefully."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'test_experiment'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
# Mock repository methods
mlflow.model_monitoring_repository.get_next_run_name.return_value = 'test_experiment-1'
mlflow.model_monitoring_repository.generate_artifacts.return_value = train_result
mlflow.model_monitoring_repository.save_run.return_value = None
# Input data without metadata
input_data = {
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify it still works (metadata defaults to {}) - returns TrainModelResult directly
assert response == train_result
assert response.run_name == 'test_experiment-1'
@mark.asyncio
async def test_save_model_complete_flow(mlflow):
"""Test save_model complete flow with all steps."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.experiment_name = 'production_model'
train_result = MagicMock(spec=TrainModelResult)
train_result.params = params
train_result.run_name = None
train_result.run_dir = None
train_result.report_path = None
# Mock complete flow
mlflow.model_monitoring_repository.get_next_run_name.return_value = 'production_model-5'
# After generate_artifacts, paths should be set
updated_result = MagicMock(spec=TrainModelResult)
updated_result.params = params
updated_result.run_name = 'production_model-5'
updated_result.run_dir = '/reports/production_model-5_20231010'
updated_result.report_path = '/reports/production_model-5_20231010/report.html'
updated_result.train_data_path = '/reports/production_model-5_20231010/train_data.csv'
updated_result.test_data_path = '/reports/production_model-5_20231010/test_data.csv'
mlflow.model_monitoring_repository.generate_artifacts.return_value = updated_result
mlflow.model_monitoring_repository.save_run.return_value = None
input_data = {
**metadata,
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify complete flow
mlflow.model_monitoring_repository.get_next_run_name.assert_called_once_with('production_model')
mlflow.model_monitoring_repository.generate_artifacts.assert_called_once()
mlflow.model_monitoring_repository.save_run.assert_called_once_with(updated_result)
# Verify response - returns TrainModelResult directly
assert response == updated_result
assert response.run_name == 'production_model-5'
assert response.run_dir is not None
assert response.report_path is not None
# ============================================================================
# Tests for cleanup_run_directory
# ============================================================================
@mark.asyncio
async def test_cleanup_run_directory_success(mlflow):
"""Test successful cleanup of run directory."""
input_data = {
**metadata,
'run_dir': 'test_run_dir',
}
# Patch os and shutil inside the activity method
with (
patch('os.path.exists', return_value=True) as mock_exists,
patch('shutil.rmtree') as mock_rmtree,
):
# Call the method
await mlflow.cleanup_run_directory(input_data)
# Verify directory existence was checked
mock_exists.assert_called_once_with('test_run_dir')
# Verify shutil.rmtree was called
mock_rmtree.assert_called_once_with('test_run_dir')
@mark.asyncio
async def test_cleanup_run_directory_already_deleted(mlflow):
"""Test cleanup when directory is already deleted (idempotent)."""
input_data = {
**metadata,
'run_dir': 'already_deleted_dir',
}
# Patch os and shutil inside the activity method
with (
patch('os.path.exists', return_value=False) as mock_exists,
patch('shutil.rmtree') as mock_rmtree,
):
# Call the method - should not raise error
await mlflow.cleanup_run_directory(input_data)
# Verify directory existence was checked
mock_exists.assert_called_once_with('already_deleted_dir')
# Verify shutil.rmtree was NOT called
mock_rmtree.assert_not_called()
@mark.asyncio
async def test_cleanup_run_directory_no_run_dir(mlflow):
"""Test cleanup when no run_dir is provided."""
input_data = {
**metadata,
# No run_dir key
}
# Call the method - should not raise error
await mlflow.cleanup_run_directory(input_data)
# Should complete without errors
@mark.asyncio
async def test_cleanup_run_directory_empty_run_dir(mlflow):
"""Test cleanup when run_dir is empty string."""
input_data = {
**metadata,
'run_dir': '',
}
# Call the method - should not raise error
await mlflow.cleanup_run_directory(input_data)
# Should complete without errors
@mark.asyncio
async def test_cleanup_run_directory_none_run_dir(mlflow):
"""Test cleanup when run_dir is None."""
input_data = {
**metadata,
'run_dir': None,
}
# Call the method - should not raise error
await mlflow.cleanup_run_directory(input_data)
# Should complete without errors
@mark.asyncio
async def test_cleanup_run_directory_error(mlflow):
"""Test cleanup handles errors correctly."""
input_data = {
**metadata,
'run_dir': 'error_dir',
}
# Patch os and shutil inside the activity method
with (
patch('os.path.exists', return_value=True),
patch('shutil.rmtree', side_effect=PermissionError('Permission denied')),
):
# Call the method - should raise exception
with pytest.raises(PermissionError, match='Permission denied'):
await mlflow.cleanup_run_directory(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='CLEANUP_RUN_DIRECTORY_ERROR',
message=ANY,
block='cleanup_run_directory',
level=NotificationLevel.ERROR,
attachment_content=ANY,
)
@mark.asyncio
async def test_cleanup_run_directory_missing_metadata(mlflow):
"""Test cleanup handles missing metadata gracefully."""
input_data = {
'run_dir': 'no_metadata_dir',
}
# Patch os and shutil inside the activity method
with patch('os.path.exists', return_value=True), patch('shutil.rmtree') as mock_rmtree:
# Call the method
await mlflow.cleanup_run_directory(input_data)
# Verify directory was deleted
mock_rmtree.assert_called_once_with('no_metadata_dir')
@mark.asyncio
async def test_cleanup_run_directory_oserror(mlflow):
"""Test cleanup handles OSError correctly."""
input_data = {
**metadata,
'run_dir': 'os_error_dir',
}
# Patch os and shutil inside the activity method
with (
patch('os.path.exists', return_value=True),
patch('shutil.rmtree', side_effect=OSError('Directory not empty')),
):
# Call the method - should raise exception
with pytest.raises(OSError, match='Directory not empty'):
await mlflow.cleanup_run_directory(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once()
call_args = mlflow.send_notification.call_args[1]
assert call_args['notification_id'] == 'CLEANUP_RUN_DIRECTORY_ERROR'
assert call_args['level'] == NotificationLevel.ERROR
assert 'Directory not empty' in call_args['message']

View File

@@ -1,499 +0,0 @@
"""Unit tests for Training activity."""
from io import BytesIO
from unittest.mock import MagicMock, patch
import pytest
from pytest import mark
from model_manager.activities.training import Training
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_success(mock_training_repository_class):
"""Test successful model training."""
# Create mock repository instance
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
# Create mock train result
mock_train_result = MagicMock(spec=TrainModelResult)
mock_train_result.mse_val = 0.5
mock_train_result.mae_val = 0.3
mock_train_result.r2_val = 0.95
mock_final_result = MagicMock(spec=TrainModelResult)
mock_final_result.mse_val = 0.5
mock_final_result.mae_val = 0.3
mock_final_result.r2_val = 0.95
# Setup repository mocks
mock_repository.train.return_value = mock_train_result
mock_repository.after_train_calculation.return_value = mock_final_result
# Create Training instance
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
# Mock inherited methods
training.info = MagicMock()
# Test data
uploaded_file = BytesIO(b'test,data\n1,2\n3,4')
train_params = TrainModelParams(
experiment_run_id=123,
target_variable='price',
variable_columns=['feature1', 'feature2'],
train_size=80,
shuffle=True,
use_scaler=True,
include_ar=False,
bucket_name='test-bucket',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0, 'feature2': 0.0},
upp_lim={'feature1': 100.0, 'feature2': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
input_data = {
'metadata': {'workflow_id': 'test-123'},
'uploaded_file': uploaded_file,
'train_params': train_params,
}
# Execute
result = await training.train_model(input_data)
# Assertions - now returns TrainModelResult directly
assert result == mock_final_result
assert result.mse_val == 0.5
assert result.mae_val == 0.3
assert result.r2_val == 0.95
# Verify repository calls
mock_repository.train.assert_called_once()
mock_repository.after_train_calculation.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_invalid_file_type(mock_training_repository_class):
"""Test training with invalid file type."""
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
# Invalid file type (string instead of BytesIO)
train_params = TrainModelParams(
experiment_run_id=123,
target_variable='price',
variable_columns=['feature1'],
train_size=80,
shuffle=True,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
input_data = {
'metadata': {},
'uploaded_file': 'not_a_bytesio',
'train_params': train_params,
}
# Should raise ValueError
with pytest.raises(ValueError, match='uploaded_file must be BytesIO'):
await training.train_model(input_data)
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_training_error(mock_training_repository_class):
"""Test training failure during model training."""
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
# Setup repository to raise error
mock_repository.train.side_effect = ValueError('Training data is empty')
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
uploaded_file = BytesIO(b'test,data\n')
train_params = TrainModelParams(
experiment_run_id=456,
target_variable='price',
variable_columns=['feature1'],
train_size=80,
shuffle=True,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
input_data = {
'metadata': {'workflow_id': 'test-456'},
'uploaded_file': uploaded_file,
'train_params': train_params,
}
# Should raise ValueError
with pytest.raises(ValueError, match='Training data is empty'):
await training.train_model(input_data)
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_sends_notification_on_error(mock_training_repository_class):
"""Test that notification is sent when training fails."""
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
mock_repository.train.side_effect = Exception('Database connection failed')
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
uploaded_file = BytesIO(b'test,data\n1,2')
train_params = TrainModelParams(
experiment_run_id=789,
target_variable='price',
variable_columns=['feature1'],
train_size=80,
shuffle=True,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
input_data = {
'metadata': {'workflow_id': 'test-789', 'experiment_run_id': 789},
'uploaded_file': uploaded_file,
'train_params': train_params,
}
# Should raise Exception
with pytest.raises(Exception, match='Database connection failed'):
await training.train_model(input_data)
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_after_calculation_error(mock_training_repository_class):
"""Test training failure during post-training calculations."""
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
# Train succeeds but after_calculation fails
mock_train_result = MagicMock(spec=TrainModelResult)
mock_repository.train.return_value = mock_train_result
mock_repository.after_train_calculation.side_effect = Exception('Metric calculation failed')
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
uploaded_file = BytesIO(b'test,data\n1,2\n3,4')
train_params = TrainModelParams(
experiment_run_id=999,
target_variable='price',
variable_columns=['feature1'],
train_size=80,
shuffle=True,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
input_data = {
'metadata': {},
'uploaded_file': uploaded_file,
'train_params': train_params,
}
# Should raise Exception
with pytest.raises(Exception, match='Metric calculation failed'):
await training.train_model(input_data)
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
async def test_train_model_invalid_train_params_type(mock_training_repository_class):
"""Test training with invalid train_params type (dict instead of TrainModelParams)."""
mock_repository = MagicMock()
mock_training_repository_class.return_value = mock_repository
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
uploaded_file = BytesIO(b'test,data\n1,2\n3,4')
# Invalid train_params type (dict instead of TrainModelParams object)
input_data = {
'metadata': {'workflow_id': 'test-invalid'},
'uploaded_file': uploaded_file,
'train_params': {
'experiment_run_id': 123,
'target_variable': 'price',
'variable_columns': ['feature1'],
}, # This is a dict, not TrainModelParams
}
# Should raise ValueError
with pytest.raises(ValueError, match='train_params must be TrainModelParams.*dict'):
await training.train_model(input_data)
# ============================================================================
# Tests for validate_train_params
# ============================================================================
@mark.asyncio
async def test_validate_train_params_success():
"""Test successful validation of training parameters."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'target_variable': 'price',
'variable_columns': ['feature1', 'feature2', 'price'], # target_variable must be in list
'train_size': 80,
'shuffle': True,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'test-bucket',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0, 'feature2': 0.0, 'price': 0.0},
'upp_lim': {'feature1': 100.0, 'feature2': 100.0, 'price': 1000.0},
'window': 10,
'experiment_name': 'test_experiment',
'removed_intervals': [],
}
result = await training.validate_train_params(input_data)
assert isinstance(result, TrainModelParams)
assert result.experiment_run_id == 456
assert result.target_variable == 'price'
assert result.variable_columns == ['feature1', 'feature2', 'price']
assert result.train_size == 80
assert result.experiment_name == 'test_experiment'
assert training.info.call_count == 2
@mark.asyncio
async def test_validate_train_params_missing_required_field():
"""Test validation fails when required field is missing."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'variable_columns': ['feature1'],
'train_size': 80,
'shuffle': True,
'use_scaler': False,
'include_ar': False,
'bucket_name': 'test',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0},
'upp_lim': {'feature1': 100.0},
'window': 10,
'experiment_name': 'test',
'removed_intervals': [],
}
with pytest.raises(ValueError, match='target_variable'):
await training.validate_train_params(input_data)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_invalid_type():
"""Test validation fails when field has invalid type."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-invalid'},
'experiment_run_id': 456,
'target_variable': 'price',
'variable_columns': ['feature1'],
'train_size': 'invalid',
'shuffle': True,
'use_scaler': False,
'include_ar': False,
'bucket_name': 'test',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0},
'upp_lim': {'feature1': 100.0},
'window': 10,
'experiment_name': 'test',
'removed_intervals': [],
}
with pytest.raises((ValueError, TypeError)):
await training.validate_train_params(input_data)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_empty_input():
"""Test validation fails with empty input."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {'metadata': {}}
with pytest.raises(ValueError):
await training.validate_train_params(input_data)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_without_metadata():
"""Test validation works even without metadata key."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
input_data = {
'experiment_run_id': 789,
'target_variable': 'temperature',
'variable_columns': ['sensor1', 'temperature'], # target_variable must be in list
'train_size': 75,
'shuffle': False,
'use_scaler': True,
'include_ar': True,
'bucket_name': 'sensors',
'file_name': 'data.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 2,
'lag_val': 2,
'rem_static_win': True,
'low_lim': {'sensor1': -50.0, 'temperature': -50.0},
'upp_lim': {'sensor1': 150.0, 'temperature': 150.0},
'window': 20,
'experiment_name': 'sensor_experiment',
'removed_intervals': [],
}
result = await training.validate_train_params(input_data)
assert isinstance(result, TrainModelParams)
assert result.experiment_run_id == 789
assert result.target_variable == 'temperature'
assert result.experiment_name == 'sensor_experiment'

View File

@@ -16,9 +16,8 @@ def test_init_with_all_credentials(mock_set_tracking_uri):
tracking_uri = 'http://mlflow.example.com'
username = 'test_user'
password = 'test_pass'
logger = MagicMock()
ModelServing(tracking_uri=tracking_uri, username=username, password=password, logger=logger)
ModelServing(tracking_uri=tracking_uri, username=username, password=password)
mock_set_tracking_uri.assert_called_once_with(tracking_uri)
import os

View File

@@ -1,497 +0,0 @@
"""Unit tests for TrainModelParams class."""
import pytest
from model_manager.utils.models.train_model_params import TrainModelParams
@pytest.fixture
def valid_params_dict():
"""Create valid parameters dictionary for testing."""
return {
'variable_columns': ['var1', 'var2', 'var3'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {'var1': 0.0, 'var2': 0.0, 'var3': 0.0},
'upp_lim': {'var1': 100.0, 'var2': 100.0, 'var3': 100.0},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'test-bucket',
'file_name': 'test-file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'Test experiment name',
'removed_intervals': [],
}
def test_train_model_params_creation_with_valid_params(valid_params_dict):
"""Test creating TrainModelParams with all valid parameters."""
params = TrainModelParams(**valid_params_dict)
assert params.variable_columns == ['var1', 'var2', 'var3']
assert params.lag_train == 5
assert params.lag_val == 3
assert params.target_variable == 'target'
assert params.rem_static_win is True
assert params.low_lim == {'var1': 0.0, 'var2': 0.0, 'var3': 0.0}
assert params.upp_lim == {'var1': 100.0, 'var2': 100.0, 'var3': 100.0}
assert params.window == 10
assert params.use_scaler is True
assert params.include_ar is False
assert params.bucket_name == 'test-bucket'
assert params.file_name == 'test-file.csv'
assert params.line_separator == '\n'
assert params.decimal_separator == '.'
assert params.train_size == 80
assert params.shuffle is True
assert params.experiment_run_id == 123
assert params.experiment_name == 'Test experiment name'
assert params.removed_intervals == []
def test_train_model_params_from_dict_creation(valid_params_dict):
"""Test creating TrainModelParams using from_dict method."""
params = TrainModelParams.from_dict(valid_params_dict)
assert params.variable_columns == ['var1', 'var2', 'var3']
assert params.lag_train == 5
assert params.experiment_run_id == 123
def test_train_model_params_variable_columns_none_raises_error(valid_params_dict):
"""Test that None variable_columns raises ValueError."""
valid_params_dict['variable_columns'] = None
with pytest.raises(ValueError, match='variable_columns is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_variable_columns_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for variable_columns raises TypeError."""
valid_params_dict['variable_columns'] = 'not a list'
with pytest.raises(TypeError, match='variable_columns must be of type list'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_lag_train_none_raises_error(valid_params_dict):
"""Test that None lag_train raises ValueError."""
valid_params_dict['lag_train'] = None
with pytest.raises(ValueError, match='lag_train is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_lag_train_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for lag_train raises TypeError."""
valid_params_dict['lag_train'] = '5'
with pytest.raises(TypeError, match='lag_train must be of type int'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_target_variable_none_raises_error(valid_params_dict):
"""Test that None target_variable raises ValueError."""
valid_params_dict['target_variable'] = None
with pytest.raises(ValueError, match='target_variable is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_target_variable_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for target_variable raises TypeError."""
valid_params_dict['target_variable'] = 123
with pytest.raises(TypeError, match='target_variable must be of type str'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_boolean_fields(valid_params_dict):
"""Test boolean fields validation."""
# Test rem_static_win
valid_params_dict['rem_static_win'] = None
with pytest.raises(ValueError, match='rem_static_win is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
valid_params_dict['rem_static_win'] = 'true'
with pytest.raises(TypeError, match='rem_static_win must be of type bool'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_dict_fields(valid_params_dict):
"""Test dict fields validation."""
# Test low_lim
valid_params_dict['low_lim'] = None
with pytest.raises(ValueError, match='low_lim is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
valid_params_dict['low_lim'] = {'var1': 0.0}
valid_params_dict['upp_lim'] = 'not a dict'
with pytest.raises(TypeError, match='upp_lim must be of type dict'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_bucket_name_none_raises_error(valid_params_dict):
"""Test that None bucket_name raises ValueError."""
valid_params_dict['bucket_name'] = None
with pytest.raises(ValueError, match='bucket_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_file_name_none_raises_error(valid_params_dict):
"""Test that None file_name raises ValueError."""
valid_params_dict['file_name'] = None
with pytest.raises(ValueError, match='file_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_experiment_run_id_none_raises_error(valid_params_dict):
"""Test that None experiment_run_id raises ValueError."""
valid_params_dict['experiment_run_id'] = None
with pytest.raises(ValueError, match='experiment_run_id is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_experiment_name_none_raises_error(valid_params_dict):
"""Test that None experiment_name raises ValueError."""
valid_params_dict['experiment_name'] = None
with pytest.raises(ValueError, match='experiment_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_removed_intervals_can_be_none(valid_params_dict):
"""Test that removed_intervals can be None (uses _check_type not _check_none)."""
valid_params_dict['removed_intervals'] = None
params = TrainModelParams.from_dict(valid_params_dict)
assert params.removed_intervals is None
def test_train_model_params_removed_intervals_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for removed_intervals raises TypeError."""
valid_params_dict['removed_intervals'] = 'not a list'
with pytest.raises(TypeError, match='removed_intervals must be of type list'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_removed_intervals_with_values(valid_params_dict):
"""Test removed_intervals with actual interval values."""
valid_params_dict['removed_intervals'] = [
('2023-01-01', '2023-01-10'),
('2023-02-01', '2023-02-05'),
]
params = TrainModelParams.from_dict(valid_params_dict)
assert len(params.removed_intervals) == 2
assert params.removed_intervals[0] == ('2023-01-01', '2023-01-10')
def test_train_model_params_all_fields_count():
"""Test that TrainModelParams has exactly 19 required fields."""
import inspect
sig = inspect.signature(TrainModelParams.__init__)
# Subtract 1 for 'self'
param_count = len(sig.parameters) - 1
assert param_count == 19
def test_train_model_params_with_minimal_valid_data():
"""Test creating params with minimal valid data."""
params = TrainModelParams(
variable_columns=['x'],
lag_train=1,
lag_val=1,
target_variable='y',
rem_static_win=False,
low_lim={},
upp_lim={},
window=1,
use_scaler=False,
include_ar=False,
bucket_name='bucket',
file_name='file.csv',
line_separator='\n',
decimal_separator='.',
train_size=50,
shuffle=False,
experiment_run_id=1,
experiment_name='name',
removed_intervals=[],
)
assert params.variable_columns == ['x']
assert params.lag_train == 1
assert params.experiment_run_id == 1
def test_train_model_params_check_none_method():
"""Test _check_none method behavior."""
params_dict = {
'variable_columns': ['var1'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {},
'upp_lim': {},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'bucket',
'file_name': 'file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'exp',
'removed_intervals': [],
}
params = TrainModelParams(**params_dict)
# Test that _check_none is a private method
assert hasattr(params, '_check_none')
assert callable(params._check_none)
def test_train_model_params_check_type_method():
"""Test _check_type method behavior."""
params_dict = {
'variable_columns': ['var1'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {},
'upp_lim': {},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'bucket',
'file_name': 'file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'exp',
'removed_intervals': [],
}
params = TrainModelParams(**params_dict)
# Test that _check_type is a private method
assert hasattr(params, '_check_type')
assert callable(params._check_type)
# ============================================================================
# Tests for validate_business_rules method
# ============================================================================
def test_validate_business_rules_success(valid_params_dict):
"""Test that valid params pass business rules validation."""
# Ensure target_variable is in variable_columns
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
# Should not raise any exception
params.validate_business_rules()
def test_validate_business_rules_train_size_too_low(valid_params_dict):
"""Test that train_size < 1 raises ValueError."""
valid_params_dict['train_size'] = 0
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='train_size must be between 1 and 99'):
params.validate_business_rules()
def test_validate_business_rules_train_size_too_high(valid_params_dict):
"""Test that train_size > 99 raises ValueError."""
valid_params_dict['train_size'] = 100
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='train_size must be between 1 and 99'):
params.validate_business_rules()
def test_validate_business_rules_empty_variable_columns(valid_params_dict):
"""Test that empty variable_columns raises ValueError."""
valid_params_dict['variable_columns'] = []
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='variable_columns cannot be empty'):
params.validate_business_rules()
def test_validate_business_rules_lag_train_zero(valid_params_dict):
"""Test that lag_train = 0 raises ValueError."""
valid_params_dict['lag_train'] = 0
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='lag_train must be positive'):
params.validate_business_rules()
def test_validate_business_rules_lag_train_negative(valid_params_dict):
"""Test that lag_train < 0 raises ValueError."""
valid_params_dict['lag_train'] = -1
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='lag_train must be positive'):
params.validate_business_rules()
def test_validate_business_rules_lag_val_zero(valid_params_dict):
"""Test that lag_val = 0 raises ValueError."""
valid_params_dict['lag_val'] = 0
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='lag_val must be positive'):
params.validate_business_rules()
def test_validate_business_rules_window_zero(valid_params_dict):
"""Test that window = 0 raises ValueError."""
valid_params_dict['window'] = 0
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='window must be positive'):
params.validate_business_rules()
def test_validate_business_rules_low_lim_upp_lim_keys_mismatch(valid_params_dict):
"""Test that mismatched keys in low_lim and upp_lim raises ValueError."""
valid_params_dict['low_lim'] = {'var1': 0.0, 'var2': 0.0}
valid_params_dict['upp_lim'] = {'var1': 100.0, 'var3': 100.0} # var3 instead of var2
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='low_lim and upp_lim must have the same keys'):
params.validate_business_rules()
def test_validate_business_rules_low_lim_greater_than_upp_lim(valid_params_dict):
"""Test that low_lim >= upp_lim raises ValueError."""
valid_params_dict['low_lim'] = {'var1': 100.0, 'var2': 0.0, 'var3': 0.0}
valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0}
valid_params_dict['target_variable'] = 'var2'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'):
params.validate_business_rules()
def test_validate_business_rules_low_lim_equal_to_upp_lim(valid_params_dict):
"""Test that low_lim == upp_lim raises ValueError."""
valid_params_dict['low_lim'] = {'var1': 50.0, 'var2': 0.0, 'var3': 0.0}
valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0}
valid_params_dict['target_variable'] = 'var2'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'):
params.validate_business_rules()
def test_validate_business_rules_target_not_in_variable_columns(valid_params_dict):
"""Test that target_variable not in variable_columns raises ValueError."""
valid_params_dict['target_variable'] = 'nonexistent_var'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(
ValueError, match='target_variable "nonexistent_var" must be in variable_columns'
):
params.validate_business_rules()
def test_validate_business_rules_empty_bucket_name(valid_params_dict):
"""Test that empty bucket_name raises ValueError."""
valid_params_dict['bucket_name'] = ' '
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='bucket_name cannot be empty or whitespace'):
params.validate_business_rules()
def test_validate_business_rules_empty_file_name(valid_params_dict):
"""Test that empty file_name raises ValueError."""
valid_params_dict['file_name'] = ''
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='file_name cannot be empty or whitespace'):
params.validate_business_rules()
def test_validate_business_rules_empty_experiment_name(valid_params_dict):
"""Test that empty experiment_name raises ValueError."""
valid_params_dict['experiment_name'] = ' '
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
with pytest.raises(ValueError, match='experiment_name cannot be empty or whitespace'):
params.validate_business_rules()
def test_validate_business_rules_all_valid_edge_cases(valid_params_dict):
"""Test that edge case valid values pass validation."""
valid_params_dict['train_size'] = 1 # Minimum valid
valid_params_dict['lag_train'] = 1 # Minimum valid
valid_params_dict['lag_val'] = 1 # Minimum valid
valid_params_dict['window'] = 1 # Minimum valid
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
# Should not raise any exception
params.validate_business_rules()
def test_validate_business_rules_train_size_99(valid_params_dict):
"""Test that train_size = 99 (maximum valid) passes validation."""
valid_params_dict['train_size'] = 99
valid_params_dict['target_variable'] = 'var1'
params = TrainModelParams.from_dict(valid_params_dict)
# Should not raise any exception
params.validate_business_rules()

View File

@@ -1,708 +0,0 @@
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
from pandas import DataFrame
from model_manager.utils.repository.model_repository import MLFlowRepository
@pytest.fixture
def mlflow_repository():
with patch(
'model_manager.utils.repository.model_repository.ModelServing', autospec=True
) as mock_model_serving:
mock_instance = mock_model_serving.return_value
mock_instance.get_transformed_data = MagicMock()
repo = MLFlowRepository(
host='http://localhost:5000', username='admin', password='admin', logger=MagicMock()
)
return repo
metadata = {
'metadata': {
'model_id': 'test_model',
'model_name': 'test_model',
'workflow_name': 'test_workflow',
'schema_name': 'test_schedule',
},
}
# ========== Tests for Model Artifact Generation Methods ==========
def test_get_next_run_name(mlflow_repository):
"""Test get_next_run_name generates correct run name based on existing runs."""
mlflow_repository.model_serving.search_runs_by_name.return_value = [
MagicMock(),
MagicMock(),
MagicMock(),
]
result = mlflow_repository.get_next_run_name('test_experiment')
mlflow_repository.model_serving.search_runs_by_name.assert_called_once_with(
experiment_names=['test_experiment'], order_by=['start_time desc']
)
assert result == 'test_experiment-4'
def test_get_next_run_name_first_run(mlflow_repository):
"""Test get_next_run_name for first run (no existing runs)."""
mlflow_repository.model_serving.search_runs_by_name.return_value = []
result = mlflow_repository.get_next_run_name('test_experiment')
assert result == 'test_experiment-1'
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_artifacts_success(mock_path, mlflow_repository):
"""Test generate_artifacts successfully creates all artifacts."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
# Mock data
params = MagicMock(spec=TrainModelParams)
params.target_variable = 'target'
params.variable_columns = ['feat1', 'feat2']
params.experiment_name = 'test_exp'
data = MagicMock(spec=TrainModelResult)
data.run_name = 'test_run-1'
data.params = params
data.x_train = DataFrame({'feat1': [1, 2], 'feat2': [3, 4]})
data.y_train = DataFrame({'target': [5, 6]})
data.x_test = DataFrame({'feat1': [7, 8], 'feat2': [9, 10]})
data.y_test = DataFrame({'target': [11, 12]})
data.regr = MagicMock()
data.regr.predict = MagicMock(return_value=np.array([5.1, 6.1]))
data.y_pred = np.array([11.1, 12.1])
# Mock path operations
mock_path.exists.return_value = True
mock_path.join.side_effect = lambda *args: '/'.join(args)
# Mock private methods
mlflow_repository._get_reports_directory = MagicMock(return_value='/reports')
mlflow_repository._create_run_directory = MagicMock(return_value='/reports/test_run-1_20231010')
mlflow_repository._setup_run_directory = MagicMock()
mlflow_repository._generate_report = MagicMock(return_value=data)
result = mlflow_repository.generate_artifacts(data)
# Assertions
mlflow_repository._get_reports_directory.assert_called_once()
mlflow_repository._create_run_directory.assert_called_once_with('/reports', 'test_run-1')
mlflow_repository._setup_run_directory.assert_called_once()
mlflow_repository._generate_report.assert_called_once()
assert result == data
def test_generate_artifacts_missing_run_name(mlflow_repository):
"""Test generate_artifacts raises ValueError when run_name is not set."""
from model_manager.utils.models.train_model_result import TrainModelResult
data = MagicMock(spec=TrainModelResult)
data.run_name = None
with pytest.raises(ValueError) as exc_info:
mlflow_repository.generate_artifacts(data)
assert 'run_name must be set before generating artifacts' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_artifacts_reports_directory_not_exists(mock_path, mlflow_repository):
"""Test generate_artifacts raises FileNotFoundError when reports directory doesn't exist."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.target_variable = 'target'
data = MagicMock(spec=TrainModelResult)
data.run_name = 'test_run-1'
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame({'target': [2]})
data.x_test = DataFrame({'feat1': [3]})
data.y_test = DataFrame({'target': [4]})
data.regr = MagicMock()
data.y_pred = np.array([4.1])
mlflow_repository._get_reports_directory = MagicMock(return_value='/reports')
mock_path.exists.return_value = False
with pytest.raises(FileNotFoundError) as exc_info:
mlflow_repository.generate_artifacts(data)
assert 'Reports directory does not exist' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_artifacts_header_file_not_exists(mock_path, mlflow_repository):
"""Test generate_artifacts raises FileNotFoundError when header.html doesn't exist."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.target_variable = 'target'
data = MagicMock(spec=TrainModelResult)
data.run_name = 'test_run-1'
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame({'target': [2]})
data.x_test = DataFrame({'feat1': [3]})
data.y_test = DataFrame({'target': [4]})
data.regr = MagicMock()
data.regr.predict = MagicMock(return_value=np.array([2.1]))
data.y_pred = np.array([4.1])
mlflow_repository._get_reports_directory = MagicMock(return_value='/reports')
mlflow_repository._create_run_directory = MagicMock(return_value='/reports/test_run-1_20231010')
# First call returns True (reports dir exists), second returns False (header.html doesn't exist)
mock_path.exists.side_effect = [True, False]
mock_path.join.side_effect = lambda *args: '/'.join(args)
with pytest.raises(FileNotFoundError) as exc_info:
mlflow_repository.generate_artifacts(data)
assert 'Header file does not exist' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_success(mock_path, mlflow_repository):
"""Test save_run successfully logs all parameters, metrics, models, and artifacts."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.train_size = 80
params.removed_intervals = [(1, 10), (20, 30)]
params.experiment_name = 'test_exp'
params.target_variable = 'target'
params.variable_columns = ['feat1', 'feat2']
params.lag_train = 5
params.lag_val = 3
params.window = 10
params.low_lim = 0.0
params.upp_lim = 1.0
params.include_ar = True
data = MagicMock(spec=TrainModelResult)
data.run_name = 'test_run-1'
data.params = params
data.report_path = '/reports/report.html'
data.train_data_path = '/reports/train.csv'
data.test_data_path = '/reports/test.csv'
data.mse_val = 0.123
data.r2_val = 0.987
data.mae_val = 0.456
data.scaler_dict = {'scaler': 'minmax'}
data.process_data = MagicMock()
data.regr = MagicMock()
mock_path.exists.return_value = True
mlflow_repository.save_run(data)
# Verify experiment was set
mlflow_repository.model_serving.set_experiment.assert_called_once_with('test_exp')
# Verify parameters were logged
assert mlflow_repository.model_serving.log_param.call_count == 13
# Verify metrics were logged
mlflow_repository.model_serving.log_metric.assert_any_call('MSE', 0.123)
mlflow_repository.model_serving.log_metric.assert_any_call('R2', 0.987)
mlflow_repository.model_serving.log_metric.assert_any_call('MAE', 0.456)
# Verify models were logged
mlflow_repository.model_serving.log_model.assert_any_call(data.process_data, 'data_model')
mlflow_repository.model_serving.log_model.assert_any_call(data.regr, 'prediction_model')
# Verify artifacts were logged
mlflow_repository.model_serving.log_artifact.assert_any_call('/reports/report.html')
mlflow_repository.model_serving.log_artifact.assert_any_call('/reports/train.csv')
mlflow_repository.model_serving.log_artifact.assert_any_call('/reports/test.csv')
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_missing_report_path(mock_path, mlflow_repository):
"""Test save_run raises ValueError when report_path is missing."""
from model_manager.utils.models.train_model_result import TrainModelResult
data = MagicMock(spec=TrainModelResult)
data.report_path = None
with pytest.raises(ValueError) as exc_info:
mlflow_repository.save_run(data)
assert 'Report file does not exist' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_missing_metrics(mock_path, mlflow_repository):
"""Test save_run raises ValueError when metrics are None."""
from model_manager.utils.models.train_model_result import TrainModelResult
data = MagicMock(spec=TrainModelResult)
data.report_path = '/reports/report.html'
data.train_data_path = '/reports/train.csv'
data.test_data_path = '/reports/test.csv'
data.mse_val = None
data.r2_val = 0.987
data.mae_val = 0.456
mock_path.exists.return_value = True
with pytest.raises(ValueError) as exc_info:
mlflow_repository.save_run(data)
assert 'One or more metrics (MSE, R2, MAE) are None' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_mlflow_error(mock_path, mlflow_repository):
"""Test save_run handles MLflow errors gracefully."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.train_size = 80
params.removed_intervals = []
params.experiment_name = 'test_exp'
data = MagicMock(spec=TrainModelResult)
data.run_name = 'test_run-1'
data.params = params
data.report_path = '/reports/report.html'
data.train_data_path = '/reports/train.csv'
data.test_data_path = '/reports/test.csv'
data.mse_val = 0.123
data.r2_val = 0.987
data.mae_val = 0.456
mock_path.exists.return_value = True
mlflow_repository.model_serving.set_experiment.side_effect = Exception(
'MLflow connection error'
)
with pytest.raises(RuntimeError) as exc_info:
mlflow_repository.save_run(data)
assert 'Failed to save run' in str(exc_info.value)
assert 'MLflow connection error' in str(exc_info.value)
# ========== Additional Tests for 100% Coverage ==========
def test_init_artifacts_data_empty_x_train(mlflow_repository):
"""Test _init_artifacts_data raises ValueError when x_train is empty."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame() # Empty DataFrame
data.y_train = DataFrame({'target': [1]})
data.x_test = DataFrame({'feat1': [1]})
data.y_test = DataFrame({'target': [1]})
with pytest.raises(ValueError) as exc_info:
mlflow_repository._init_artifacts_data(data)
assert 'Training features (x_train) are empty' in str(exc_info.value)
def test_init_artifacts_data_empty_y_train(mlflow_repository):
"""Test _init_artifacts_data raises ValueError when y_train is empty."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame() # Empty DataFrame
data.x_test = DataFrame({'feat1': [1]})
data.y_test = DataFrame({'target': [1]})
with pytest.raises(ValueError) as exc_info:
mlflow_repository._init_artifacts_data(data)
assert 'Training target (y_train) is empty' in str(exc_info.value)
def test_init_artifacts_data_empty_x_test(mlflow_repository):
"""Test _init_artifacts_data raises ValueError when x_test is empty."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame({'target': [1]})
data.x_test = DataFrame() # Empty DataFrame
data.y_test = DataFrame({'target': [1]})
with pytest.raises(ValueError) as exc_info:
mlflow_repository._init_artifacts_data(data)
assert 'Test features (x_test) are empty' in str(exc_info.value)
def test_init_artifacts_data_empty_y_test(mlflow_repository):
"""Test _init_artifacts_data raises ValueError when y_test is empty."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame({'target': [1]})
data.x_test = DataFrame({'feat1': [1]})
data.y_test = DataFrame() # Empty DataFrame
with pytest.raises(ValueError) as exc_info:
mlflow_repository._init_artifacts_data(data)
assert 'Test target (y_test) is empty' in str(exc_info.value)
def test_init_artifacts_data_none_y_pred(mlflow_repository):
"""Test _init_artifacts_data raises ValueError when y_pred is None."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame({'feat1': [1]})
data.y_train = DataFrame({'target': [1]})
data.x_test = DataFrame({'feat1': [1]})
data.y_test = DataFrame({'target': [1]})
data.y_pred = None
with pytest.raises(ValueError) as exc_info:
mlflow_repository._init_artifacts_data(data)
assert 'Test predictions (y_pred) are None' in str(exc_info.value)
def test_init_artifacts_data_success(mlflow_repository):
"""Test _init_artifacts_data successfully prepares data."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.target_variable = 'target'
data = MagicMock(spec=TrainModelResult)
data.params = params
data.x_train = DataFrame({'feat1': [1, 2]})
data.y_train = DataFrame({'target': [3, 4]})
data.x_test = DataFrame({'feat1': [5, 6]})
data.y_test = DataFrame({'target': [7, 8]})
data.regr = MagicMock()
data.regr.predict = MagicMock(return_value=np.array([3.1, 4.1]))
data.y_pred = np.array([7.1, 8.1])
reference_data, current_data = mlflow_repository._init_artifacts_data(data)
assert 'target' in reference_data.columns
assert 'prediction' in reference_data.columns
assert 'target' in current_data.columns
assert 'prediction' in current_data.columns
assert len(reference_data) == 2
assert len(current_data) == 2
@patch('model_manager.utils.repository.model_repository.makedirs')
@patch('model_manager.utils.repository.model_repository.path')
def test_create_run_directory_success(mock_path, mock_makedirs, mlflow_repository):
"""Test _create_run_directory successfully creates directory."""
mock_path.join.return_value = '/reports/test_run_20231010_123456_123456'
result = mlflow_repository._create_run_directory('/reports', 'test_run')
mock_makedirs.assert_called_once_with('/reports/test_run_20231010_123456_123456', exist_ok=True)
assert result == '/reports/test_run_20231010_123456_123456'
@patch('model_manager.utils.repository.model_repository.makedirs')
@patch('model_manager.utils.repository.model_repository.path')
def test_create_run_directory_permission_error(mock_path, mock_makedirs, mlflow_repository):
"""Test _create_run_directory handles PermissionError."""
mock_path.join.return_value = '/reports/test_run_20231010'
mock_makedirs.side_effect = PermissionError('Permission denied')
with pytest.raises(PermissionError) as exc_info:
mlflow_repository._create_run_directory('/reports', 'test_run')
assert 'Permission denied when creating directory' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.makedirs')
@patch('model_manager.utils.repository.model_repository.path')
def test_create_run_directory_os_error(mock_path, mock_makedirs, mlflow_repository):
"""Test _create_run_directory handles OSError."""
mock_path.join.return_value = '/reports/test_run_20231010'
mock_makedirs.side_effect = OSError('Disk full')
with pytest.raises(OSError) as exc_info:
mlflow_repository._create_run_directory('/reports', 'test_run')
assert 'Failed to create directory' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.shutil')
@patch('model_manager.utils.repository.model_repository.path')
def test_setup_run_directory_success(mock_path, mock_shutil, mlflow_repository):
"""Test _setup_run_directory successfully sets up directory."""
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_open = MagicMock()
with patch('builtins.open', mock_open):
mlflow_repository._setup_run_directory('/run_dir', '/reports/header.html')
assert mock_open.call_count == 3 # 3 empty files
mock_shutil.copy.assert_called_once_with('/reports/header.html', '/run_dir/header.html')
@patch('model_manager.utils.repository.model_repository.shutil')
@patch('model_manager.utils.repository.model_repository.path')
def test_setup_run_directory_file_not_found(mock_path, mock_shutil, mlflow_repository):
"""Test _setup_run_directory handles FileNotFoundError."""
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_shutil.copy.side_effect = FileNotFoundError('Header not found')
mock_open = MagicMock()
with patch('builtins.open', mock_open):
with pytest.raises(FileNotFoundError) as exc_info:
mlflow_repository._setup_run_directory('/run_dir', '/reports/header.html')
assert 'Header file not found' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.shutil')
@patch('model_manager.utils.repository.model_repository.path')
def test_setup_run_directory_permission_error(mock_path, mock_shutil, mlflow_repository):
"""Test _setup_run_directory handles PermissionError."""
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_open = MagicMock()
mock_open.side_effect = PermissionError('Permission denied')
with patch('builtins.open', mock_open):
with pytest.raises(PermissionError) as exc_info:
mlflow_repository._setup_run_directory('/run_dir', '/reports/header.html')
assert 'Permission denied when setting up directory' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.shutil')
@patch('model_manager.utils.repository.model_repository.path')
def test_setup_run_directory_os_error(mock_path, mock_shutil, mlflow_repository):
"""Test _setup_run_directory handles OSError."""
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_open = MagicMock()
mock_open.side_effect = OSError('Disk error')
with patch('builtins.open', mock_open):
with pytest.raises(OSError) as exc_info:
mlflow_repository._setup_run_directory('/run_dir', '/reports/header.html')
assert 'Failed to setup run directory' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.Reports')
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_report_success(mock_path, mock_reports, mlflow_repository):
"""Test _generate_report successfully generates all reports."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.variable_columns = ['feat1', 'feat2']
data = MagicMock(spec=TrainModelResult)
data.params = params
data.run_dir = '/run_dir'
reference_data = DataFrame(
{'feat1': [1.0], 'feat2': [2.0], 'target': [3.0], 'prediction': [3.1]}
)
current_data = DataFrame({'feat1': [4.0], 'feat2': [5.0], 'target': [6.0], 'prediction': [6.1]})
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_report_instance = MagicMock()
mock_reports.return_value = mock_report_instance
# Mock DataFrame.to_csv to avoid actual file writing
with patch.object(DataFrame, 'to_csv'):
result = mlflow_repository._generate_report(reference_data, current_data, data)
mock_reports.assert_called_once()
mock_report_instance.add_data_quality_section.assert_called_once()
mock_report_instance.add_data_drift_section.assert_called_once()
mock_report_instance.add_regression_section.assert_called_once()
mock_report_instance.save_all_sections_html.assert_called_once()
assert result == data
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_report_value_error(mock_path, mlflow_repository):
"""Test _generate_report handles ValueError from data conversion."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
data = MagicMock(spec=TrainModelResult)
data.params = params
data.run_dir = '/run_dir'
# DataFrame with non-numeric data
reference_data = DataFrame({'feat1': ['a', 'b']})
current_data = DataFrame({'feat1': ['c', 'd']})
with pytest.raises(ValueError) as exc_info:
mlflow_repository._generate_report(reference_data, current_data, data)
assert 'Failed to convert data to float64' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.Reports')
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_report_permission_error(mock_path, mock_reports, mlflow_repository):
"""Test _generate_report handles PermissionError."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.variable_columns = ['feat1']
data = MagicMock(spec=TrainModelResult)
data.params = params
data.run_dir = '/run_dir'
reference_data = DataFrame({'feat1': [1.0]})
current_data = DataFrame({'feat1': [2.0]})
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_report_instance = MagicMock()
mock_reports.return_value = mock_report_instance
mock_report_instance.save_all_sections_html.side_effect = PermissionError('Permission denied')
with pytest.raises(PermissionError) as exc_info:
mlflow_repository._generate_report(reference_data, current_data, data)
assert 'Permission denied when writing report files' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.Reports')
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_report_os_error(mock_path, mock_reports, mlflow_repository):
"""Test _generate_report handles OSError."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.variable_columns = ['feat1']
data = MagicMock(spec=TrainModelResult)
data.params = params
data.run_dir = '/run_dir'
reference_data = DataFrame({'feat1': [1.0]})
current_data = DataFrame({'feat1': [2.0]})
mock_path.join.side_effect = lambda *args: '/'.join(args)
mock_report_instance = MagicMock()
mock_reports.return_value = mock_report_instance
mock_report_instance.save_all_sections_html.side_effect = OSError('Disk error')
with pytest.raises(OSError) as exc_info:
mlflow_repository._generate_report(reference_data, current_data, data)
assert 'Failed to generate report' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.Reports')
@patch('model_manager.utils.repository.model_repository.path')
def test_generate_report_run_dir_none(mock_path, mock_reports, mlflow_repository):
"""Test _generate_report raises ValueError when run_dir is None."""
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
params = MagicMock(spec=TrainModelParams)
params.variable_columns = ['feat1']
data = MagicMock(spec=TrainModelResult)
data.params = params
data.run_dir = None # Not set
reference_data = DataFrame({'feat1': [1.0]})
current_data = DataFrame({'feat1': [2.0]})
mock_report_instance = MagicMock()
mock_reports.return_value = mock_report_instance
with pytest.raises(ValueError) as exc_info:
mlflow_repository._generate_report(reference_data, current_data, data)
assert 'run_dir is not set after directory creation' in str(exc_info.value)
def test_get_reports_directory(mlflow_repository):
"""Test _get_reports_directory returns correct path."""
result = mlflow_repository._get_reports_directory()
assert result.endswith('model_manager/reports')
assert 'model_manager' in result
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_missing_train_data_path(mock_path, mlflow_repository):
"""Test save_run raises ValueError when train_data_path is missing."""
from model_manager.utils.models.train_model_result import TrainModelResult
data = MagicMock(spec=TrainModelResult)
data.report_path = '/reports/report.html'
data.train_data_path = None
mock_path.exists.return_value = True
with pytest.raises(ValueError) as exc_info:
mlflow_repository.save_run(data)
assert 'Training data file does not exist' in str(exc_info.value)
@patch('model_manager.utils.repository.model_repository.path')
def test_save_run_missing_test_data_path(mock_path, mlflow_repository):
"""Test save_run raises ValueError when test_data_path is missing."""
from model_manager.utils.models.train_model_result import TrainModelResult
data = MagicMock(spec=TrainModelResult)
data.report_path = '/reports/report.html'
data.train_data_path = '/reports/train.csv'
data.test_data_path = None
mock_path.exists.return_value = True
with pytest.raises(ValueError) as exc_info:
mlflow_repository.save_run(data)
assert 'Test data file does not exist' in str(exc_info.value)

View File

@@ -1,324 +0,0 @@
"""Unit tests for TrainingRepository."""
from io import BytesIO
from unittest.mock import MagicMock, patch
import numpy as np
import pandas as pd
from pytest import fixture, raises
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.utils.repository.training_repository import TrainingRepository
@fixture
def logger():
"""Create a mock logger."""
return MagicMock()
@fixture
def training_repository(logger):
"""Create a TrainingRepository instance."""
return TrainingRepository(logger)
@fixture
def train_params():
"""Create sample training parameters."""
return TrainModelParams(
variable_columns=['feature1', 'feature2'],
lag_train=1,
lag_val=1,
target_variable='target',
rem_static_win=False,
low_lim={'feature1': 0.0, 'feature2': 0.0},
upp_lim={'feature1': 100.0, 'feature2': 100.0},
window=10,
use_scaler=True,
include_ar=False,
bucket_name='test-bucket',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
train_size=80,
shuffle=True,
experiment_run_id=123,
experiment_name='test_experiment',
removed_intervals=[],
)
@fixture
def sample_csv_data():
"""Create sample CSV data."""
csv_content = """feature1,feature2,target
1.0,2.0,10.0
2.0,3.0,15.0
3.0,4.0,20.0
4.0,5.0,25.0
5.0,6.0,30.0
6.0,7.0,35.0
7.0,8.0,40.0
8.0,9.0,45.0
9.0,10.0,50.0
10.0,11.0,55.0
"""
return BytesIO(csv_content.encode())
@patch('model_manager.utils.repository.training_repository.load_data')
@patch('model_manager.utils.repository.training_repository.DataPreprocessor')
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.LinearRegressionModel')
def test_train_success(
mock_linear_model,
mock_split,
mock_preprocessor_class,
mock_load_data,
training_repository,
train_params,
sample_csv_data,
):
"""Test successful model training."""
# Setup mocks
mock_data = pd.DataFrame(
{'feature1': [1, 2, 3, 4, 5], 'feature2': [2, 3, 4, 5, 6], 'target': [10, 15, 20, 25, 30]}
)
mock_load_data.return_value = mock_data
mock_preprocessor = MagicMock()
mock_preprocessor_class.return_value = mock_preprocessor
mock_preprocessor.transform.return_value = mock_data
x_train = pd.DataFrame({'feature1': [1, 2, 3], 'feature2': [2, 3, 4]})
x_test = pd.DataFrame({'feature1': [4, 5], 'feature2': [5, 6]})
y_train = pd.Series([10, 15, 20], name='target')
y_test = pd.Series([25, 30], name='target')
mock_split.return_value = (x_train, x_test, y_train, y_test)
mock_model = MagicMock()
mock_linear_model.return_value = mock_model
mock_scaler = MagicMock()
mock_preprocessor.get_scaler.return_value = mock_scaler
# Execute
result = training_repository.train(sample_csv_data, train_params)
# Assertions
assert isinstance(result, TrainModelResult)
assert result.params == train_params
assert result.process_data == mock_preprocessor
assert result.regr == mock_model
mock_load_data.assert_called_once()
mock_preprocessor.fit.assert_called_once()
mock_model.fit.assert_called_once()
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_empty_data_after_transform(
mock_load_data, training_repository, train_params, sample_csv_data
):
"""Test training with empty data after transformation."""
mock_data = pd.DataFrame({'feature1': [], 'feature2': [], 'target': []})
mock_load_data.return_value = mock_data
with patch.object(training_repository, 'init_data_preprocessor') as mock_init:
mock_preprocessor = MagicMock()
mock_init.return_value = mock_preprocessor
mock_preprocessor.transform.return_value = pd.DataFrame()
with raises(ValueError, match='Data view is empty after transformation'):
training_repository.train(sample_csv_data, train_params)
def test_init_scaler_dict_with_minmax_scaler(training_repository, train_params):
"""Test scaler dict initialization with MinMaxScaler."""
mock_preprocessor = MagicMock()
mock_scaler = MagicMock()
mock_scaler.x_min = [0.0, 1.0]
mock_scaler.x_max = [10.0, 11.0]
mock_scaler.y_min = 5.0
mock_scaler.y_max = 50.0
mock_preprocessor.get_scaler.return_value = mock_scaler
# Patch isinstance to return True for MinMaxScaler
with patch(
'model_manager.utils.repository.training_repository.isinstance',
side_effect=lambda obj, cls: cls.__name__ == 'MinMaxScaler',
):
result = training_repository.init_scaler_dict(mock_preprocessor, train_params)
assert result is not None
assert 'feature1' in result
assert 'feature2' in result
assert 'target' in result
assert result['feature1'] == {'min': 0.0, 'max': 10.0}
assert result['feature2'] == {'min': 1.0, 'max': 11.0}
assert result['target'] == {'min': 5.0, 'max': 50.0}
def test_init_scaler_dict_with_z_scaler(training_repository, train_params):
"""Test scaler dict initialization with Z_Scaler."""
mock_preprocessor = MagicMock()
mock_scaler = MagicMock()
mock_scaler.create_dict.return_value = {'mean': 5.0, 'std': 2.0}
mock_preprocessor.get_scaler.return_value = mock_scaler
# Patch isinstance to return True for Z_Scaler
with patch(
'model_manager.utils.repository.training_repository.isinstance',
side_effect=lambda obj, cls: cls.__name__ == 'Z_Scaler',
):
result = training_repository.init_scaler_dict(mock_preprocessor, train_params)
assert result == {'mean': 5.0, 'std': 2.0}
mock_scaler.create_dict.assert_called_once()
def test_init_scaler_dict_without_scaler(training_repository):
"""Test scaler dict initialization when use_scaler is False."""
train_params_no_scaler = TrainModelParams(
variable_columns=['feature1'],
lag_train=1,
lag_val=1,
target_variable='target',
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
train_size=80,
shuffle=True,
experiment_run_id=123,
experiment_name='test',
removed_intervals=[],
)
mock_preprocessor = MagicMock()
result = training_repository.init_scaler_dict(mock_preprocessor, train_params_no_scaler)
assert result == {}
@patch('model_manager.utils.repository.training_repository.mse')
@patch('model_manager.utils.repository.training_repository.mae')
@patch('model_manager.utils.repository.training_repository.r2')
def test_after_train_calculation_with_scaler(
mock_r2, mock_mae, mock_mse, training_repository, train_params
):
"""Test post-training calculations with scaler."""
# Setup mock train result
mock_train_result = MagicMock(spec=TrainModelResult)
mock_train_result.params = train_params
mock_train_result.x_train = pd.DataFrame({'feature1': [1, 2, 3], 'feature2': [2, 3, 4]})
mock_train_result.x_test = pd.DataFrame({'feature1': [4, 5], 'feature2': [5, 6]})
mock_train_result.y_train = pd.Series([10, 15, 20], name='target')
mock_train_result.y_test = pd.Series([25, 30], name='target')
mock_regr = MagicMock()
mock_regr.predict.return_value = np.array([24.5, 29.5])
mock_train_result.regr = mock_regr
mock_scaler = MagicMock()
mock_scaler.denormalize_single_input.side_effect = lambda x, col: x
mock_scaler.denormalize_predictions.side_effect = lambda x, col: x
mock_process_data = MagicMock()
mock_process_data.get_scaler.return_value = mock_scaler
mock_train_result.process_data = mock_process_data
# Setup metric mocks
mock_mse.return_value = 0.5
mock_mae.return_value = 0.3
mock_r2.return_value = 0.95
# Execute
result = training_repository.after_train_calculation(train_params, mock_train_result)
# Assertions
assert result == mock_train_result
assert result.mse_val == 0.5
assert result.mae_val == 0.3
assert result.r2_val == 0.95
assert result.y_pred is not None
mock_regr.predict.assert_called_once()
mock_mse.assert_called_once()
mock_mae.assert_called_once()
mock_r2.assert_called_once()
@patch('model_manager.utils.repository.training_repository.mse')
@patch('model_manager.utils.repository.training_repository.mae')
@patch('model_manager.utils.repository.training_repository.r2')
def test_after_train_calculation_without_scaler(mock_r2, mock_mae, mock_mse, training_repository):
"""Test post-training calculations without scaler."""
train_params_no_scaler = TrainModelParams(
variable_columns=['feature1'],
lag_train=1,
lag_val=1,
target_variable='target',
rem_static_win=False,
low_lim={'feature1': 0.0},
upp_lim={'feature1': 100.0},
window=10,
use_scaler=False,
include_ar=False,
bucket_name='test',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
train_size=80,
shuffle=True,
experiment_run_id=123,
experiment_name='test',
removed_intervals=[],
)
mock_train_result = MagicMock(spec=TrainModelResult)
mock_train_result.params = train_params_no_scaler
mock_train_result.x_train = pd.DataFrame({'feature1': [1, 2, 3]})
mock_train_result.x_test = pd.DataFrame({'feature1': [4, 5]})
mock_train_result.y_train = pd.Series([10, 15, 20], name='target')
mock_train_result.y_test = pd.Series([25, 30], name='target')
mock_regr = MagicMock()
mock_regr.predict.return_value = np.array([24.5, 29.5])
mock_train_result.regr = mock_regr
# Setup metric mocks
mock_mse.return_value = 0.5
mock_mae.return_value = 0.3
mock_r2.return_value = 0.95
# Execute
result = training_repository.after_train_calculation(train_params_no_scaler, mock_train_result)
# Assertions
assert result.mse_val == 0.5
assert result.mae_val == 0.3
assert result.r2_val == 0.95
@patch('model_manager.utils.repository.training_repository.DataPreprocessor')
def test_init_data_preprocessor(mock_preprocessor_class, training_repository, train_params):
"""Test DataPreprocessor initialization."""
mock_preprocessor = MagicMock()
mock_preprocessor_class.return_value = mock_preprocessor
result = training_repository.init_data_preprocessor(train_params)
assert result == mock_preprocessor
mock_preprocessor_class.assert_called_once()
call_kwargs = mock_preprocessor_class.call_args[1]
assert call_kwargs['target_variable'] == 'target'
assert call_kwargs['input_columns'] == ['feature1', 'feature2']
assert call_kwargs['low_lim'] == {'feature1': 0.0, 'feature2': 0.0}
assert call_kwargs['upp_lim'] == {'feature1': 100.0, 'feature2': 100.0}

View File

@@ -10,8 +10,7 @@ from model_manager.utils.connectors_config import (
def test_build_mlflow_config_with_env_vars():
# Arrange
environ['MLFLOW_HOST'] = 'http://test-host'
environ['MLFLOW_PORT'] = '8080'
environ['MLFLOW_URL'] = 'http://test-host:8080'
environ['MLFLOW_USERNAME'] = 'test-user'
environ['MLFLOW_PASSWORD'] = 'test-pass'
@@ -19,8 +18,7 @@ def test_build_mlflow_config_with_env_vars():
config = build_mlflow_config()
# Assert
assert config['host'] == 'http://test-host'
assert config['port'] == 8080
assert config['url'] == 'http://test-host:8080'
assert config['username'] == 'test-user'
assert config['password'] == 'test-pass'
@@ -28,8 +26,7 @@ def test_build_mlflow_config_with_env_vars():
def test_build_mlflow_config_with_defaults():
# Arrange
# Clear any existing env vars
environ.pop('MLFLOW_HOST', None)
environ.pop('MLFLOW_PORT', None)
environ.pop('MLFLOW_URL', None)
environ.pop('MLFLOW_USERNAME', None)
environ.pop('MLFLOW_PASSWORD', None)
@@ -37,8 +34,7 @@ def test_build_mlflow_config_with_defaults():
config = build_mlflow_config()
# Assert
assert config['host'] == 'http://localhost'
assert config['port'] == 5080
assert config['url'] == 'http://localhost:5080'
assert config['username'] == 'aignosi'
assert config['password'] == 'aignosi'

View File

@@ -1,494 +0,0 @@
from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest
@pytest.fixture
def mock_env_vars(monkeypatch):
"""Fixture to set up environment variables for tests."""
monkeypatch.setenv('POD_ID', 'test-pod-123')
monkeypatch.setenv('TEMPORAL_HOST', 'test-temporal:7233')
monkeypatch.setenv('TEMPORAL_NAMESPACE', 'test-namespace')
monkeypatch.setenv('HTTP_METRICS_PORT', '9090')
monkeypatch.setenv('HTTP_SDK_METRICS_PORT', '9091')
monkeypatch.setenv('PROJECT_NAME', 'test-project')
monkeypatch.setenv('POSTGRES_HOST', 'localhost')
monkeypatch.setenv('POSTGRES_PORT', '5432')
monkeypatch.setenv('POSTGRES_USER', 'test')
monkeypatch.setenv('POSTGRES_PASSWORD', 'test')
monkeypatch.setenv('POSTGRES_DBNAME', 'test')
monkeypatch.setenv('MLFLOW_HOST', 'http://localhost')
monkeypatch.setenv('MLFLOW_PORT', '5000')
monkeypatch.setenv('MLFLOW_USERNAME', 'test')
monkeypatch.setenv('MLFLOW_PASSWORD', 'test')
monkeypatch.setenv('MINIO_ENDPOINT_URL', 'http://localhost:9000')
monkeypatch.setenv('MINIO_ACCESS_KEY', 'test')
monkeypatch.setenv('MINIO_SECRET_KEY', 'test')
monkeypatch.setenv('MONGODB_USERNAME', 'test')
monkeypatch.setenv('MONGODB_PASSWORD', 'test')
monkeypatch.setenv('MONGODB_URL', 'localhost:27017')
monkeypatch.setenv('MONGODB_DATABASE_NAME', 'test')
@patch('model_manager.worker.worker.start_http_server')
@patch('model_manager.worker.worker.metrics')
def test_start_prometheus_server_success(mock_metrics, mock_start_http_server, mock_env_vars):
"""Test successful Prometheus server startup."""
# Import after patching to ensure mocks are in place
from model_manager.worker.worker import start_prometheus_server
# Act
start_prometheus_server()
# Assert
mock_start_http_server.assert_called_once_with(9090)
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(1)
@patch('model_manager.worker.worker.start_http_server')
@patch('model_manager.worker.worker.os._exit')
def test_start_prometheus_server_failure(mock_exit, mock_start_http_server, mock_env_vars):
"""Test Prometheus server startup failure."""
# Arrange
mock_start_http_server.side_effect = Exception('Port already in use')
# Import after patching
from model_manager.worker.worker import start_prometheus_server
# Act
start_prometheus_server()
# Assert
mock_start_http_server.assert_called_once_with(9090)
mock_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_success(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test successful main function execution."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete successfully
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert
mock_start_prometheus.assert_called_once()
mock_notification_handler.assert_called_once()
mock_activities.assert_called_once()
mock_client_connect.assert_called_once_with(
target_host='test-temporal:7233',
namespace='test-namespace',
runtime=ANY,
)
assert mock_worker.call_count == 1 # Only one worker created
mock_gather.assert_called_once()
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
async def test_main_exception_handling(
mock_metrics,
mock_sys_exit,
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_env_vars,
):
"""Test main function exception handling and cleanup."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to raise an exception
mock_gather.side_effect = Exception('Worker failed')
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - Verify cleanup was performed
mock_logger.custom_error.assert_called_once()
mock_handler.shutdown.assert_called_once()
mock_activities_instance.shutdown.assert_called_once()
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(0)
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_creates_only_one_worker(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main creates only one worker with correct configurations."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify only one worker was created
assert mock_worker.call_count == 1
# Verify worker (train_model-queue)
first_call = mock_worker.call_args_list[0]
assert first_call[1]['task_queue'] == 'train_model-queue'
assert 'TrainModel' in str(first_call[1]['workflows'])
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_initializes_activities_with_configs(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main initializes Activities with correct configurations."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify Activities was initialized with correct parameters
mock_activities.assert_called_once()
call_kwargs = mock_activities.call_args[1]
assert 'postgres_config' in call_kwargs
assert 'mlflow_config' in call_kwargs
assert 'minio_config' in call_kwargs
assert call_kwargs['logger'] == mock_logger
assert call_kwargs['notification_handler'] == mock_handler
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_uses_environment_variables(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main uses environment variables correctly."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify environment variables were used
mock_client_connect.assert_called_once_with(
target_host='test-temporal:7233',
namespace='test-namespace',
runtime=ANY,
)
mock_notification_handler.assert_called_once()
notification_call_kwargs = mock_notification_handler.call_args[1]
assert notification_call_kwargs['project_name'] == 'test-project'
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_cleanup_with_none_notification_handler(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test cleanup when notification_handler is None (line 194 branch False)."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
# Return None for notification_handler
mock_notification_handler.return_value = None
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - notification_handler.shutdown() should NOT be called (line 194 False)
# Since notification_handler is None, we can't call shutdown on it
mock_activities_instance.shutdown.assert_called_once()
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_cleanup_with_falsy_activities(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test cleanup when activities evaluates to False (line 196 branch False)."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
# Create a falsy activities object (empty list, 0, False, etc.)
# Using an object that evaluates to False but doesn't cause AttributeError
class FalsyActivities:
def __bool__(self):
return False
def __getattr__(self, name):
# Return mock methods to avoid AttributeError during worker creation
return MagicMock()
falsy_activities = FalsyActivities()
mock_activities.return_value = falsy_activities
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - notification_handler.shutdown() is called, but activities.shutdown() is NOT
mock_handler.shutdown.assert_called_once()
# activities is falsy, so shutdown should NOT be called
mock_sys_exit.assert_called_once_with(1)

View File

@@ -1,678 +0,0 @@
"""Unit tests for TrainModel workflow."""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from pytest import fixture, mark
from model_manager.utils.models.experiment_status import ExperimentStatus
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.workflows.train_model import TrainModel
@fixture
def train_model_workflow() -> TrainModel:
"""Fixture for TrainModel workflow instance."""
return TrainModel()
@fixture
def mock_train_params():
"""Fixture for mock TrainModelParams."""
return TrainModelParams(
experiment_run_id=123,
target_variable='price',
variable_columns=['feature1', 'feature2'],
train_size=80,
shuffle=True,
use_scaler=True,
include_ar=False,
bucket_name='test-bucket',
file_name='test.csv',
line_separator='\n',
decimal_separator='.',
lag_train=1,
lag_val=1,
rem_static_win=False,
low_lim={'feature1': 0.0, 'feature2': 0.0},
upp_lim={'feature1': 100.0, 'feature2': 100.0},
window=10,
experiment_name='test_experiment',
removed_intervals=[],
)
@fixture
def mock_train_result(mock_train_params):
"""Fixture for mock TrainModelResult."""
result = MagicMock(spec=TrainModelResult)
result.params = mock_train_params
result.run_name = 'test_experiment-1'
result.run_dir = 'test_run_dir' # Relative path instead of /tmp
return result
# ============================================================================
# Tests for run() - Complete workflow
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_run_success_complete_flow(
workflow_mock: AsyncMock,
train_model_workflow: TrainModel,
mock_train_params,
mock_train_result,
):
"""Test successful complete workflow execution."""
input_data = {
'experiment_run_id': 123,
'target_variable': 'price',
'variable_columns': ['feature1', 'feature2'],
'train_size': 80,
'shuffle': True,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'test-bucket',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0, 'feature2': 0.0},
'upp_lim': {'feature1': 100.0, 'feature2': 100.0},
'window': 10,
'experiment_name': 'test_experiment',
'removed_intervals': [],
}
# Mock workflow.logger to avoid RuntimeWarning about unawaited coroutines
workflow_mock.logger = MagicMock()
# Mock activity responses
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update_experiment_run (MAGE_WAITING_PROC)
b'file_content', # fetch_file_from_minio
mock_train_result, # train_model
None, # update_experiment_run (TRAINING_SUCCESS)
mock_train_result, # save_model
None, # update_experiment_run (MLFLOW_SENT with run_name)
None, # cleanup_run_directory
None, # delete_file_from_minio
None, # update_experiment_run (FILE_DELETED)
]
)
# Execute workflow
await train_model_workflow.run(input_data)
# Verify all activity calls (now 10 instead of 9 due to cleanup_run_directory)
assert workflow_mock.execute_activity_method.call_count == 10
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_run_missing_experiment_run_id(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test workflow fails when experiment_run_id is missing."""
input_data = {
'target_variable': 'price',
'variable_columns': ['feature1'],
}
# Mock workflow.logger to avoid NotInWorkflowEventLoopError
workflow_mock.logger = MagicMock()
with pytest.raises(ValueError, match='experiment_run_id is required'):
await train_model_workflow.run(input_data)
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_run_invalid_experiment_run_id_type(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test workflow fails when experiment_run_id has invalid type."""
input_data = {
'experiment_run_id': 'invalid', # Should be int
'target_variable': 'price',
}
# Mock workflow.logger to avoid NotInWorkflowEventLoopError
workflow_mock.logger = MagicMock()
with pytest.raises(ValueError, match='experiment_run_id must be an integer'):
await train_model_workflow.run(input_data)
# ============================================================================
# Tests for _validate_experiment_run_id()
# ============================================================================
@patch('model_manager.workflows.train_model.workflow')
def test_validate_experiment_run_id_success(
workflow_mock: MagicMock, train_model_workflow: TrainModel
):
"""Test successful experiment_run_id validation."""
workflow_mock.logger = MagicMock()
input_data = {'experiment_run_id': 456}
result = train_model_workflow._validate_experiment_run_id(input_data)
assert result == 456
@patch('model_manager.workflows.train_model.workflow')
def test_validate_experiment_run_id_missing(
workflow_mock: MagicMock, train_model_workflow: TrainModel
):
"""Test validation fails when experiment_run_id is missing."""
workflow_mock.logger = MagicMock()
input_data = {}
with pytest.raises(ValueError, match='experiment_run_id is required'):
train_model_workflow._validate_experiment_run_id(input_data)
@patch('model_manager.workflows.train_model.workflow')
def test_validate_experiment_run_id_none(
workflow_mock: MagicMock, train_model_workflow: TrainModel
):
"""Test validation fails when experiment_run_id is None."""
workflow_mock.logger = MagicMock()
input_data = {'experiment_run_id': None}
with pytest.raises(ValueError, match='experiment_run_id is required'):
train_model_workflow._validate_experiment_run_id(input_data)
@patch('model_manager.workflows.train_model.workflow')
def test_validate_experiment_run_id_invalid_type(
workflow_mock: MagicMock, train_model_workflow: TrainModel
):
"""Test validation fails when experiment_run_id is not an integer."""
workflow_mock.logger = MagicMock()
input_data = {'experiment_run_id': 'not_an_int'}
with pytest.raises(ValueError, match='experiment_run_id must be an integer'):
train_model_workflow._validate_experiment_run_id(input_data)
# ============================================================================
# Tests for _validate_training_parameters()
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_validate_training_parameters_success(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params
):
"""Test successful parameter validation."""
input_data = {
'experiment_run_id': 123,
'target_variable': 'price',
'variable_columns': ['feature1'],
'train_size': 80,
'shuffle': True,
'use_scaler': False,
'include_ar': False,
'bucket_name': 'test',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0},
'upp_lim': {'feature1': 100.0},
'window': 10,
'experiment_name': 'test',
'removed_intervals': [],
}
metadata = {
'metadata': {
'experiment_run_id': 123,
'workflow_name': 'train_model',
}
}
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update_experiment_run
]
)
result = await train_model_workflow._validate_training_parameters(input_data, 123, metadata)
assert result == mock_train_params
assert workflow_mock.execute_activity_method.call_count == 2
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_validate_training_parameters_validation_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test parameter validation handles errors correctly."""
input_data = {'experiment_run_id': 123}
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
# Mock workflow.logger to avoid RuntimeWarning about unawaited coroutines
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
ValueError('Missing required field'), # validate_train_params fails
None, # update_experiment_run with error
]
)
with pytest.raises(ValueError, match='Missing required field'):
await train_model_workflow._validate_training_parameters(input_data, 123, metadata)
# Verify error status was updated
assert workflow_mock.execute_activity_method.call_count == 2
# ============================================================================
# Tests for _download_and_train_model()
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_success(
workflow_mock: AsyncMock,
train_model_workflow: TrainModel,
mock_train_params,
mock_train_result,
):
"""Test successful download and training."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
b'file_content', # fetch_file_from_minio
mock_train_result, # train_model
None, # update_experiment_run
]
)
result = await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
assert result == mock_train_result
assert workflow_mock.execute_activity_method.call_count == 3
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_download_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params
):
"""Test download error is handled correctly."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
# Mock workflow.logger to avoid RuntimeWarning about unawaited coroutines
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
Exception('MinIO connection failed'), # fetch_file_from_minio fails
None, # update_experiment_run with error
]
)
with pytest.raises(Exception, match='MinIO connection failed'):
await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
# Verify error status was updated
assert workflow_mock.execute_activity_method.call_count == 2
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_training_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params
):
"""Test training error is handled correctly."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
# Mock workflow.logger to avoid RuntimeWarning about unawaited coroutines
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
b'file_content', # fetch_file_from_minio succeeds
Exception('Training failed'), # train_model fails
None, # update_experiment_run with error
]
)
with pytest.raises(Exception, match='Training failed'):
await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
# Verify error status was updated
assert workflow_mock.execute_activity_method.call_count == 3
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_closes_bytesio_on_success(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result
):
"""Test that BytesIO is closed in finally block on success."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
# Create a mock BytesIO with close method
mock_file = MagicMock()
mock_file.close = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_file, # fetch_file_from_minio returns BytesIO
mock_train_result, # train_model succeeds
None, # update_experiment_run
]
)
result = await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
# Verify BytesIO.close() was called in finally block
mock_file.close.assert_called_once()
assert result == mock_train_result
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_closes_bytesio_on_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params
):
"""Test that BytesIO is closed in finally block even on error."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
# Create a mock BytesIO with close method
mock_file = MagicMock()
mock_file.close = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_file, # fetch_file_from_minio returns BytesIO
Exception('Training failed'), # train_model fails
None, # update_experiment_run with error
]
)
with pytest.raises(Exception, match='Training failed'):
await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
# Verify BytesIO.close() was called in finally block even after exception
mock_file.close.assert_called_once()
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_download_and_train_model_handles_file_without_close(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result
):
"""Test that workflow handles file objects without close method gracefully."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
# Create a mock file without close method
mock_file = MagicMock(spec=[])
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_file, # fetch_file_from_minio returns object without close
mock_train_result, # train_model succeeds
None, # update_experiment_run
]
)
# Should not raise error even if file doesn't have close method
result = await train_model_workflow._download_and_train_model(
train_params=mock_train_params, experiment_run_id=123, metadata=metadata
)
assert result == mock_train_result
# ============================================================================
# Tests for _save_model_to_mlflow()
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_save_model_to_mlflow_success(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result
):
"""Test successful model saving to MLFlow."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
mock_train_result, # save_model
None, # update_experiment_run with MODEL_SAVED
]
)
result = await train_model_workflow._save_model_to_mlflow(
train_result=mock_train_result, experiment_run_id=123, metadata=metadata
)
assert result == mock_train_result
assert workflow_mock.execute_activity_method.call_count == 2
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_save_model_to_mlflow_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result
):
"""Test MLFlow save error is handled correctly."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
Exception('MLFlow connection failed'), # save_model fails
None, # update_experiment_run with error
]
)
with pytest.raises(Exception, match='MLFlow connection failed'):
await train_model_workflow._save_model_to_mlflow(
train_result=mock_train_result, experiment_run_id=123, metadata=metadata
)
# Verify error status was updated
assert workflow_mock.execute_activity_method.call_count == 2
# ============================================================================
# Tests for _cleanup_resources()
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_cleanup_resources_success(
workflow_mock: AsyncMock,
train_model_workflow: TrainModel,
mock_train_result,
):
"""Test successful resource cleanup."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
None, # cleanup_run_directory
None, # delete_file_from_minio
None, # update_experiment_run
]
)
await train_model_workflow._cleanup_resources(
saved_result=mock_train_result, experiment_run_id=123, metadata=metadata
)
# Verify activities were called (cleanup_run_directory + delete_file_from_minio + update_experiment_run)
assert workflow_mock.execute_activity_method.call_count == 3
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_cleanup_resources_delete_error(
workflow_mock: AsyncMock,
train_model_workflow: TrainModel,
mock_train_result,
):
"""Test cleanup handles delete errors correctly."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
workflow_mock.logger = MagicMock()
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
None, # cleanup_run_directory succeeds
Exception('MinIO delete failed'), # delete_file_from_minio fails
None, # update_experiment_run with error
]
)
with pytest.raises(Exception, match='MinIO delete failed'):
await train_model_workflow._cleanup_resources(
saved_result=mock_train_result, experiment_run_id=123, metadata=metadata
)
# Verify error status was updated (cleanup_run_directory + delete_file_from_minio + update_experiment_run)
assert workflow_mock.execute_activity_method.call_count == 3
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_cleanup_resources_without_run_dir(
workflow_mock: AsyncMock,
train_model_workflow: TrainModel,
mock_train_result,
):
"""Test cleanup works when run_dir is not set."""
metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}}
# Mock result without run_dir
mock_train_result.run_dir = None
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
None, # delete_file_from_minio
None, # update_experiment_run
]
)
await train_model_workflow._cleanup_resources(
saved_result=mock_train_result, experiment_run_id=123, metadata=metadata
)
# Verify cleanup_run_directory was NOT called (no run_dir)
# Only delete_file_from_minio + update_experiment_run
assert workflow_mock.execute_activity_method.call_count == 2
# ============================================================================
# Tests for _update_experiment_run()
# ============================================================================
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_update_experiment_run_status_only(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test updating experiment run with status only."""
from model_manager.activities.experiment_tracking import UpdateType
metadata = {'metadata': {'experiment_run_id': 123}}
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
await train_model_workflow._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.STATUS,
status=ExperimentStatus.TRAINING_SUCCESS,
)
workflow_mock.execute_activity_method.assert_called_once()
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_update_experiment_run_with_error(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test updating experiment run with error message."""
from model_manager.activities.experiment_tracking import UpdateType
metadata = {'metadata': {'experiment_run_id': 123}}
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
await train_model_workflow._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.TRAINING_ERROR,
error_message='Training failed',
)
workflow_mock.execute_activity_method.assert_called_once()
call_args = workflow_mock.execute_activity_method.call_args[0][1]
assert call_args['error_message'] == 'Training failed'
@mark.asyncio
@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock)
async def test_update_experiment_run_with_run_name(
workflow_mock: AsyncMock, train_model_workflow: TrainModel
):
"""Test updating experiment run with run_name."""
from model_manager.activities.experiment_tracking import UpdateType
metadata = {'metadata': {'experiment_run_id': 123}}
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
await train_model_workflow._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.MODEL_SAVED,
status=ExperimentStatus.MLFLOW_SENT,
run_name='test_experiment-1',
)
workflow_mock.execute_activity_method.assert_called_once()
call_args = workflow_mock.execute_activity_method.call_args[0][1]
assert call_args['run_name'] == 'test_experiment-1'