diff --git a/model_manager/activities/mlflow.py b/model_manager/activities/mlflow.py index 0458c05..1762b04 100644 --- a/model_manager/activities/mlflow.py +++ b/model_manager/activities/mlflow.py @@ -13,6 +13,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.temporal.activities.base import BaseActivity from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ + from model_manager.utils.models.train_model_result import TrainModelResult from model_manager.utils.repository.model_repository import MLFlowRepository @@ -346,19 +347,14 @@ class MLFlow(BaseActivity): raise e @activity.defn(name='save_model') - async def save_model(self, input_data: dict[str, Any]) -> dict[str, Any]: + async def save_model(self, input_data: dict[str, Any]) -> TrainModelResult: """ - Save a trained ML model and its artifacts to MLflow with comprehensive error handling. + Save a trained ML model and its artifacts to MLflow. This activity orchestrates the complete model saving pipeline: 1. Generates the next run name for the experiment 2. Creates and organizes artifacts (reports, data files) 3. Logs model, parameters, metrics, and artifacts to MLflow - 4. Returns success/failure status with results or error message - - The activity does NOT raise exceptions on failure - it catches all errors, - sends notifications, and returns a failure status. This allows the workflow - to handle the error gracefully and update the database accordingly. Args: input_data: Configuration for model saving operation @@ -367,23 +363,17 @@ class MLFlow(BaseActivity): - train_result (TrainModelResult): Training result with model and metrics Returns: - dict: Save result with the following structure: - { - 'success': bool, # True if saving succeeded, False otherwise - 'result': TrainModelResult | None, # Updated result if success=True - 'error_message': str | None # Error message if success=False - } + TrainModelResult: Updated training result with run_name and artifacts + + Raises: + Exception: If model saving fails (after sending notification) Example: - # Successful save result = await save_model({ 'metadata': {'workflow_id': 'save-123', 'experiment_run_id': 456}, 'train_result': TrainModelResult(...) }) - # Returns: {'success': True, 'result': TrainModelResult(...), 'error_message': None} - - # Failed save - # Returns: {'success': False, 'result': None, 'error_message': 'Error details...'} + # Returns: TrainModelResult with run_name and artifacts """ metadata = input_data.get('metadata', {}) train_result = input_data['train_result'] @@ -418,11 +408,7 @@ class MLFlow(BaseActivity): metadata, ) - return { - 'success': True, - 'result': train_result, - 'error_message': None, - } + return train_result except Exception as e: # noqa: BLE001 error_msg = f'Error saving model - Experiment: {train_result.params.experiment_name if train_result and train_result.params else "unknown"}, Error: {str(e)}' @@ -441,10 +427,5 @@ class MLFlow(BaseActivity): # Log error with metadata self.error(trace, metadata=metadata) - # Return failure result (do NOT raise exception) - # This allows workflow to update database with error status - return { - 'success': False, - 'result': None, - 'error_message': str(e), - } + # Re-raise exception to stop workflow + raise diff --git a/model_manager/activities/training.py b/model_manager/activities/training.py index 8b9dee7..0040c37 100644 --- a/model_manager/activities/training.py +++ b/model_manager/activities/training.py @@ -19,6 +19,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.temporal.activities.base import BaseActivity 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 @@ -51,20 +52,84 @@ class Training(BaseActivity): BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True) self.training_repository = TrainingRepository(logger) - @activity.defn(name='train_model') - async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]: + @activity.defn(name='validate_train_params') + async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams: """ - Train a machine learning model with comprehensive error handling. + Validate and convert training parameters from dict to TrainModelParams. + + This activity validates the input training parameters and converts them + to a TrainModelParams object. + + Args: + input_data: Training parameters and metadata at the same level + Required keys: + - metadata (dict): Workflow execution metadata + - All TrainModelParams fields (experiment_run_id, target_variable, etc.) + + Returns: + TrainModelParams: Validated and converted training parameters + + Raises: + ValueError, TypeError, KeyError: If validation fails (after sending notification) + + Example: + result = await validate_train_params({ + 'metadata': {'workflow_id': 'train-123'}, + 'experiment_run_id': 456, + 'target_variable': 'price', + 'variable_columns': ['feature1', 'feature2'], + 'train_size': 80, + # ... other required fields at same level + }) + # Returns: TrainModelParams(...) + """ + metadata = input_data.get('metadata', {}) + + try: + self.info('Validating training parameters', metadata) + + # Convert input_data directly to TrainModelParams (this validates all fields) + # The from_dict method will extract only the fields it needs + train_params = TrainModelParams.from_dict(input_data) + + self.info( + f'Training parameters validated successfully - ' + f'Target: {train_params.target_variable}, ' + f'Experiment: {train_params.experiment_name}', + metadata, + ) + + return train_params + + except (ValueError, TypeError, KeyError) as e: + error_msg = f'Error validating training parameters: {str(e)}' + trace = traceback.format_exc() + + # Send notification (MongoDB) + self.send_notification( + metadata=metadata, + notification_id='VALIDATE_TRAIN_PARAMS_ERROR', + message=error_msg, + block='validate_train_params', + level=NotificationLevel.ERROR, + attachment_content=trace, + ) + + # Log error with metadata + self.error(trace, metadata=metadata) + + # Re-raise exception to stop workflow + raise + + @activity.defn(name='train_model') + async def train_model(self, input_data: dict[str, Any]) -> TrainModelResult: + """ + Train a machine learning model. This activity orchestrates the complete ML training pipeline: 1. Validates input parameters 2. Trains the model using TrainingRepository 3. Performs post-training calculations - 4. Returns success/failure status with results or error message - - The activity does NOT raise exceptions on failure - it catches all errors, - sends notifications, and returns a failure status. This allows the workflow - to handle the error gracefully and update the database accordingly. Args: input_data: Configuration for model training operation @@ -74,24 +139,19 @@ class Training(BaseActivity): - train_params (TrainModelParams): Training parameters object Returns: - dict: Training result with the following structure: - { - 'success': bool, # True if training succeeded, False otherwise - 'result': TrainModelResult | None, # Training result if success=True - 'error_message': str | None # Error message if success=False - } + TrainModelResult: Training result with model, metrics, and data + + Raises: + ValueError: If input validation fails + Exception: If training fails (after sending notification) Example: - # Successful training result = await train_model({ 'metadata': {'workflow_id': 'train-123', 'experiment_run_id': 456}, 'uploaded_file': BytesIO(csv_data), - 'train_params': TrainModelParams(...) # Already converted object + 'train_params': TrainModelParams(...) }) - # Returns: {'success': True, 'result': TrainModelResult(...), 'error_message': None} - - # Failed training - # Returns: {'success': False, 'result': None, 'error_message': 'Error details...'} + # Returns: TrainModelResult(...) """ metadata = input_data.get('metadata', {}) uploaded_file = input_data['uploaded_file'] @@ -127,11 +187,7 @@ class Training(BaseActivity): metadata, ) - return { - 'success': True, - 'result': final_result, - 'error_message': None, - } + return final_result except Exception as e: # noqa: BLE001 target = ( @@ -155,10 +211,5 @@ class Training(BaseActivity): # Log error with metadata self.error(trace, metadata=metadata) - # Return failure result (do NOT raise exception) - # This allows workflow to update database with error status - return { - 'success': False, - 'result': None, - 'error_message': str(e), - } + # Re-raise exception to stop workflow + raise diff --git a/model_manager/utils/models/experiment_status.py b/model_manager/utils/models/experiment_status.py index 139f2de..2cfcbca 100644 --- a/model_manager/utils/models/experiment_status.py +++ b/model_manager/utils/models/experiment_status.py @@ -23,6 +23,7 @@ class ExperimentStatus(str, Enum): FILE_DELETE_ERROR: Cleanup failed due to file system or MinIO errors. """ + ORCHESTRATOR_VALIDATION_ERROR = 'ORCHESTRATOR_VALIDATION_ERROR' MAGE_WAITING_PROC = 'MAGE_WAITING_PROC' TRAINING_SUCCESS = 'TRAINING_SUCCESS' TRAINING_ERROR = 'TRAINING_ERROR' diff --git a/model_manager/workflows/train_model.py b/model_manager/workflows/train_model.py new file mode 100644 index 0000000..c3a4065 --- /dev/null +++ b/model_manager/workflows/train_model.py @@ -0,0 +1,480 @@ +""" +Train Model Workflow for ML model training pipeline. + +This workflow orchestrates the complete ML model training process, including: +- Parameter validation and conversion +- Data download from MinIO +- Model training +- Model saving to MLFlow +- Experiment tracking and status updates +""" + +from temporalio import workflow + +with workflow.unsafe.imports_passed_through(): + import shutil + from datetime import timedelta + from typing import Any + + from sientia_do.temporal.policies import retry_policy + + from model_manager.activities.activities import Activities + from model_manager.activities.experiment_tracking import UpdateType + 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 + + +@workflow.defn(name='train_model') +class TrainModel: + """ + Complete ML model training workflow. + + This workflow implements the full training pipeline from parameter validation + through model training and saving to MLFlow. It provides comprehensive error + handling with database status updates at each stage. + + The workflow ensures: + - Proper parameter validation before training starts + - Status tracking in database for monitoring + - Error handling with detailed error messages + - Cleanup and proper resource management + """ + + @workflow.run + async def run(self, input_data: dict[str, Any]) -> None: + """ + Execute the complete model training workflow. + + This method orchestrates all steps of the training pipeline: + 1. Validate and convert training parameters + 2. Download training data from MinIO + 3. Train the model + 4. Save model to MLFlow + 5. Update experiment tracking status + + Args: + input_data: Complete configuration for the training workflow + All TrainModelParams fields at the same level: + - experiment_run_id (int): Unique identifier for the experiment run (REQUIRED) + - target_variable (str): Target variable to predict + - variable_columns (list[str]): Feature columns + - train_size (int): Training data percentage + - ... (all other TrainModelParams fields) + + Raises: + ValueError: If experiment_run_id is missing or invalid + """ + # CRITICAL: Validate experiment_run_id first + # Without it, we cannot update database status, so fail immediately + experiment_run_id = self._validate_experiment_run_id(input_data) + + metadata = { + 'metadata': { + 'experiment_run_id': experiment_run_id, + 'workflow_name': 'train_model', + } + } + + # Step 1: Validate and convert training parameters + train_params = await self._validate_training_parameters( + input_data, experiment_run_id, metadata + ) + + # Step 2 & 3: Download from MinIO and Train model + # The _download_and_train_model method handles both steps: + # - Downloads file from MinIO (returns BytesIO) + # - Trains model with the downloaded file + # Any error (download OR training) = TRAINING_ERROR + train_result = await self._download_and_train_model( + train_params=train_params, + experiment_run_id=experiment_run_id, + metadata=metadata, + ) + + # Step 4: Save model to MLFlow + saved_result = await self._save_model_to_mlflow( + train_result=train_result, + experiment_run_id=experiment_run_id, + metadata=metadata, + ) + + # Step 5: Cleanup resources and delete file from MinIO + await self._cleanup_resources( + saved_result=saved_result, + experiment_run_id=experiment_run_id, + metadata=metadata, + ) + + def _validate_experiment_run_id(self, input_data: dict[str, Any]) -> int: + """ + Validate experiment_run_id from input data. + + This method ensures that experiment_run_id is present and valid. + Without a valid experiment_run_id, we cannot update database status, + so this validation must happen before any other operation. + + Args: + input_data: Input data dictionary containing experiment_run_id + + Returns: + int: Validated experiment_run_id + + Raises: + ValueError: If experiment_run_id is missing or not an integer + """ + experiment_run_id = input_data.get('experiment_run_id') + + if experiment_run_id is None: + error_msg = 'experiment_run_id is required but was not provided' + workflow.logger.error(error_msg) + raise ValueError(error_msg) + + if not isinstance(experiment_run_id, int): + error_msg = ( + f'experiment_run_id must be an integer, got {type(experiment_run_id).__name__}' + ) + workflow.logger.error(error_msg) + raise ValueError(error_msg) + + return experiment_run_id + + async def _validate_training_parameters( + self, + input_data: dict[str, Any], + experiment_run_id: int, + metadata: dict[str, Any], + ) -> TrainModelParams: + """ + Validate and convert training parameters from dict to TrainModelParams. + + This method calls the validate_train_params activity to convert and validate + the input parameters. On success, updates DB status to MAGE_WAITING_PROC. + On error, updates DB status to ORCHESTRATOR_VALIDATION_ERROR. + + Args: + input_data: Input data dictionary containing all training parameters + experiment_run_id: Validated experiment run ID + metadata: Workflow execution metadata + + Returns: + TrainModelParams: Validated training parameters object + + Raises: + Exception: If validation fails (after updating DB status) + """ + validation_input = { + **metadata, + **input_data, # All training params at same level + } + + try: + # Execute validation activity + train_params = await workflow.execute_activity_method( + Activities.validate_train_params, + validation_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=30), + ) + + # Validation succeeded: Update status + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS, + status=ExperimentStatus.MAGE_WAITING_PROC, + ) + + return train_params + + except Exception as e: + # Validation failed: Update status with error + error_message = str(e) + + workflow.logger.error( + f'Validation failed for experiment {experiment_run_id}: {error_message}' + ) + + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS_WITH_ERROR, + status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR, + error_message=error_message, + ) + + # Re-raise exception to stop workflow + raise + + async def _download_and_train_model( + self, + train_params: TrainModelParams, + experiment_run_id: int, + metadata: dict[str, Any], + ) -> TrainModelResult: + """ + Download file from MinIO and train model. + + This method orchestrates the download and training steps using proper + resource management with try/catch/finally. On success, updates DB status + to TRAINING_SUCCESS. On error, updates DB status to TRAINING_ERROR. + + Args: + train_params: TrainModelParams object with training configuration + experiment_run_id: Validated experiment run ID + metadata: Workflow execution metadata + + Returns: + TrainModelResult: Training result from train_model activity + + Raises: + Exception: If download or training fails (after updating DB status) + """ + uploaded_file = None + + try: + # Step 1: Download file from MinIO + download_input = { + **metadata, + 'bucket_name': train_params.bucket_name, + 'file_name': train_params.file_name, + } + + uploaded_file = await workflow.execute_activity_method( + Activities.fetch_file_from_minio, + download_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60), + ) + + # Step 2: Train model with downloaded file + train_input = { + **metadata, + 'uploaded_file': uploaded_file, + 'train_params': train_params, + } + + train_result = await workflow.execute_activity_method( + Activities.train_model, + train_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=300), + ) + + # Training succeeded: Update status + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS, + status=ExperimentStatus.TRAINING_SUCCESS, + ) + + return train_result + + except Exception as e: + # Download or training failed: Update status with error + error_message = str(e) + + workflow.logger.error( + f'Download or training failed for experiment {experiment_run_id}: {error_message}' + ) + + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS_WITH_ERROR, + status=ExperimentStatus.TRAINING_ERROR, + error_message=error_message, + ) + + # Re-raise exception to stop workflow + raise + + finally: + # Ensure BytesIO is closed + if uploaded_file and hasattr(uploaded_file, 'close'): + uploaded_file.close() + + async def _save_model_to_mlflow( + self, + train_result: TrainModelResult, + experiment_run_id: int, + metadata: dict[str, Any], + ) -> TrainModelResult: + """ + Save trained model to MLFlow. + + This method calls the save_model activity to save the trained model and + its artifacts to MLFlow. On success, updates DB status to MLFLOW_SENT + with MODEL_SAVED type and run_name. On error, updates DB status to + MLFLOW_SEND_ERROR. + + Args: + train_result: TrainModelResult from training step + experiment_run_id: Validated experiment run ID + metadata: Workflow execution metadata + + Returns: + TrainModelResult: Updated training result with MLFlow run name + + Raises: + Exception: If model saving fails (after updating DB status) + """ + try: + # Save model to MLFlow + save_input = { + **metadata, + 'train_result': train_result, + } + + saved_result = await workflow.execute_activity_method( + Activities.save_model, + save_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=120), + ) + + # Model saved successfully: Update status to MLFLOW_SENT with run_name + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.MODEL_SAVED, + status=ExperimentStatus.MLFLOW_SENT, + run_name=saved_result.run_name, + ) + + return saved_result + + except Exception as e: + # Model saving failed: Update status with error + error_message = str(e) + + workflow.logger.error( + f'Model saving failed for experiment {experiment_run_id}: {error_message}' + ) + + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS_WITH_ERROR, + status=ExperimentStatus.MLFLOW_SEND_ERROR, + error_message=error_message, + ) + + # Re-raise exception to stop workflow + raise + + async def _cleanup_resources( + self, + saved_result: TrainModelResult, + experiment_run_id: int, + metadata: dict[str, Any], + ) -> None: + """ + Cleanup resources and delete file from MinIO. + + This method removes the temporary run directory and deletes the training + file from MinIO. On success, updates DB status to FILE_DELETED. + On error, updates DB status to FILE_DELETE_ERROR. + + Args: + saved_result: TrainModelResult with run_dir and params information + experiment_run_id: Validated experiment run ID + metadata: Workflow execution metadata + + Raises: + Exception: If cleanup fails (after updating DB status) + """ + try: + # Step 1: Remove temporary run directory + if hasattr(saved_result, 'run_dir') and saved_result.run_dir: + shutil.rmtree(saved_result.run_dir) + workflow.logger.info( + f'Removed run directory: {saved_result.run_dir} for experiment {experiment_run_id}' + ) + + # Step 2: Delete file from MinIO + delete_input = { + **metadata, + 'bucket_name': saved_result.params.bucket_name, + 'file_name': saved_result.params.file_name, + } + + await workflow.execute_activity_method( + Activities.delete_file_from_minio, + delete_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=30), + ) + + # Cleanup succeeded: Update status to FILE_DELETED + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS, + status=ExperimentStatus.FILE_DELETED, + ) + + except Exception as e: + # Cleanup failed: Update status with error + error_message = str(e) + + workflow.logger.error( + f'Cleanup failed for experiment {experiment_run_id}: {error_message}' + ) + + await self._update_experiment_run( + metadata=metadata, + experiment_run_id=experiment_run_id, + update_type=UpdateType.STATUS_WITH_ERROR, + status=ExperimentStatus.FILE_DELETE_ERROR, + error_message=error_message, + ) + + # Re-raise exception to stop workflow + raise + + async def _update_experiment_run( + self, + metadata: dict[str, Any], + experiment_run_id: int, + update_type: UpdateType, + status: ExperimentStatus, + error_message: str | None = None, + run_name: str | None = None, + ) -> None: + """ + Update experiment run status in the database. + + This is a helper method to simplify calls to the update_experiment_run activity. + It handles both success and error status updates. + + Args: + metadata: Workflow execution metadata + experiment_run_id: Unique identifier for the experiment run + update_type: Type of update (STATUS, STATUS_WITH_ERROR, or MODEL_SAVED) + status: Status to set in the database + error_message: Error message (required if update_type is STATUS_WITH_ERROR) + run_name: MLFlow run name (required if update_type is MODEL_SAVED) + """ + update_input = { + **metadata, + 'experiment_run_id': experiment_run_id, + 'update_type': update_type, + 'status': status, + } + + # Add error_message only if provided + if error_message is not None: + update_input['error_message'] = error_message + + # Add run_name only if provided + if run_name is not None: + update_input['run_name'] = run_name + + await workflow.execute_activity_method( + Activities.update_experiment_run, + update_input, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=30), + ) diff --git a/tests/activities/test_mlflow.py b/tests/activities/test_mlflow.py index 779f9d4..7a0c1cb 100644 --- a/tests/activities/test_mlflow.py +++ b/tests/activities/test_mlflow.py @@ -1,6 +1,7 @@ from unittest.mock import ANY, MagicMock, patch import numpy as np +import pytest from pytest import fixture, mark from sientia_do.notifications.models import NotificationLevel from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ @@ -333,10 +334,9 @@ async def test_save_model_success(mlflow): 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 - assert response['success'] is True - assert response['result'] == train_result - assert response['error_message'] is None + # 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' @@ -362,13 +362,9 @@ async def test_save_model_get_next_run_name_error(mlflow): 'train_result': train_result, } - # Call the method - response = await mlflow.save_model(input_data) - - # Verify error handling - assert response['success'] is False - assert response['result'] is None - assert 'MLflow connection error' in response['error_message'] + # 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( @@ -404,13 +400,9 @@ async def test_save_model_generate_artifacts_error(mlflow): 'train_result': train_result, } - # Call the method - response = await mlflow.save_model(input_data) - - # Verify error handling - assert response['success'] is False - assert response['result'] is None - assert 'Reports directory does not exist' in response['error_message'] + # 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( @@ -447,13 +439,9 @@ async def test_save_model_save_run_error(mlflow): 'train_result': train_result, } - # Call the method - response = await mlflow.save_model(input_data) - - # Verify error handling - assert response['success'] is False - assert response['result'] is None - assert 'One or more metrics (MSE, R2, MAE) are None' in response['error_message'] + # 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( @@ -491,10 +479,9 @@ async def test_save_model_missing_metadata(mlflow): # Call the method response = await mlflow.save_model(input_data) - # Verify it still works (metadata defaults to {}) - assert response['success'] is True - assert response['result'] == train_result - assert response['error_message'] is None + # Verify it still works (metadata defaults to {}) - returns TrainModelResult directly + assert response == train_result + assert response.run_name == 'test_experiment-1' @mark.asyncio @@ -540,10 +527,8 @@ async def test_save_model_complete_flow(mlflow): mlflow.model_monitoring_repository.generate_artifacts.assert_called_once() mlflow.model_monitoring_repository.save_run.assert_called_once_with(updated_result) - # Verify response - assert response['success'] is True - assert response['result'] == updated_result - assert response['error_message'] is None - assert updated_result.run_name == 'production_model-5' - assert updated_result.run_dir is not None - assert updated_result.report_path is not None + # 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 diff --git a/tests/activities/test_training.py b/tests/activities/test_training.py index df2affc..3a2b582 100644 --- a/tests/activities/test_training.py +++ b/tests/activities/test_training.py @@ -3,6 +3,7 @@ from io import BytesIO from unittest.mock import MagicMock, patch +import pytest from pytest import mark from model_manager.activities.training import Training @@ -74,10 +75,11 @@ async def test_train_model_success(mock_training_repository_class): # Execute result = await training.train_model(input_data) - # Assertions - assert result['success'] is True - assert result['result'] == mock_final_result - assert result['error_message'] is None + # 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() @@ -124,11 +126,10 @@ async def test_train_model_invalid_file_type(mock_training_repository_class): 'train_params': train_params, } - result = await training.train_model(input_data) + # Should raise ValueError + with pytest.raises(ValueError, match='uploaded_file must be BytesIO'): + await training.train_model(input_data) - assert result['success'] is False - assert result['result'] is None - assert 'uploaded_file must be BytesIO' in result['error_message'] # Verify notification was sent (via BaseActivity) notification_handler.send_notification.assert_called_once() @@ -176,11 +177,10 @@ async def test_train_model_training_error(mock_training_repository_class): 'train_params': train_params, } - result = await training.train_model(input_data) + # Should raise ValueError + with pytest.raises(ValueError, match='Training data is empty'): + await training.train_model(input_data) - assert result['success'] is False - assert result['result'] is None - assert 'Training data is empty' in result['error_message'] # Verify notification was sent (via BaseActivity) notification_handler.send_notification.assert_called_once() @@ -227,16 +227,13 @@ async def test_train_model_sends_notification_on_error(mock_training_repository_ 'train_params': train_params, } - result = await training.train_model(input_data) + # 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() - # Verify result - the important part is that error was caught and returned - assert result['success'] is False - assert result['result'] is None - assert 'Database connection failed' in result['error_message'] - @mark.asyncio @patch('model_manager.activities.training.TrainingRepository') @@ -283,11 +280,10 @@ async def test_train_model_after_calculation_error(mock_training_repository_clas 'train_params': train_params, } - result = await training.train_model(input_data) + # Should raise Exception + with pytest.raises(Exception, match='Metric calculation failed'): + await training.train_model(input_data) - assert result['success'] is False - assert result['result'] is None - assert 'Metric calculation failed' in result['error_message'] # Verify notification was sent (via BaseActivity) notification_handler.send_notification.assert_called_once() @@ -316,11 +312,188 @@ async def test_train_model_invalid_train_params_type(mock_training_repository_cl }, # This is a dict, not TrainModelParams } - result = await training.train_model(input_data) + # 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'], + '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': [], + } + + 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'] + 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) - assert result['success'] is False - assert result['result'] is None - assert 'train_params must be TrainModelParams' in result['error_message'] - assert 'dict' in result['error_message'] - # Verify notification was sent (via BaseActivity) 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'], + '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}, + 'upp_lim': {'sensor1': 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' diff --git a/tests/utils/models/test_experiment_status.py b/tests/utils/models/test_experiment_status.py index 238b06f..4ec5ae4 100644 --- a/tests/utils/models/test_experiment_status.py +++ b/tests/utils/models/test_experiment_status.py @@ -15,8 +15,8 @@ def test_experiment_status_values(): def test_experiment_status_count(): - """Test that enum has exactly 7 status values.""" - assert len(ExperimentStatus) == 7 + """Test that enum has exactly 8 status values.""" + assert len(ExperimentStatus) == 8 def test_experiment_status_is_string(): @@ -40,7 +40,7 @@ def test_experiment_status_membership(): def test_experiment_status_iteration(): """Test that enum can be iterated.""" statuses = list(ExperimentStatus) - assert len(statuses) == 7 + assert len(statuses) == 8 assert ExperimentStatus.MAGE_WAITING_PROC in statuses assert ExperimentStatus.TRAINING_SUCCESS in statuses assert ExperimentStatus.TRAINING_ERROR in statuses diff --git a/tests/workflows/test_train_model.py b/tests/workflows/test_train_model.py new file mode 100644 index 0000000..8db760a --- /dev/null +++ b/tests/workflows/test_train_model.py @@ -0,0 +1,672 @@ +"""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.shutil') +@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) +async def test_run_success_complete_flow( + workflow_mock: AsyncMock, + mock_shutil: MagicMock, + 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 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, # delete_file_from_minio + None, # update_experiment_run (FILE_DELETED) + ] + ) + + # Execute workflow + await train_model_workflow.run(input_data) + + # Verify all activity calls + assert workflow_mock.execute_activity_method.call_count == 9 + # Verify shutil.rmtree was called + mock_shutil.rmtree.assert_called_once_with('test_run_dir') + + +@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'}} + + 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'}} + + 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'}} + + 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'}} + + # 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'}} + + # 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.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.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) +@patch('model_manager.workflows.train_model.shutil') +async def test_cleanup_resources_success( + mock_shutil: MagicMock, + 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.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 shutil.rmtree was called + mock_shutil.rmtree.assert_called_once_with('test_run_dir') + + # Verify activities were called + assert workflow_mock.execute_activity_method.call_count == 2 + + +@mark.asyncio +@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) +@patch('model_manager.workflows.train_model.shutil') +async def test_cleanup_resources_delete_error( + mock_shutil: MagicMock, + 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.execute_activity_method = AsyncMock( + side_effect=[ + 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 + assert workflow_mock.execute_activity_method.call_count == 2 + + +@mark.asyncio +@patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) +@patch('model_manager.workflows.train_model.shutil') +async def test_cleanup_resources_without_run_dir( + mock_shutil: MagicMock, + 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 shutil.rmtree was NOT called + mock_shutil.rmtree.assert_not_called() + + # Verify activities were still called + 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'