SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
@@ -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''
|
||||
@@ -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']
|
||||
@@ -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'
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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}
|
||||
@@ -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'
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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'
|
||||
Reference in New Issue
Block a user