SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.
This commit is contained in:
@@ -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