SIENTIAPDE-1253: Refactor training workflow and activities to raise exceptions on failure
This commit refactors the training workflow and associated activities to raise exceptions on failure instead of returning success/failure dictionaries. This allows the Temporal workflow to handle errors more effectively and ensures that the workflow stops when a critical error occurs. Key changes: - The train_model workflow is introduced to orchestrate the entire training process, including parameter validation, data download, model training, and model saving. - The validate_train_params activity is added to validate and convert training parameters. - The train_model and save_model activities are updated to raise exceptions on failure. - The ExperimentStatus enum is updated to include a new status for orchestrator validation errors. - The tests are updated to reflect the new exception-based error handling. - The activities now return the TrainModelResult directly instead of a dictionary.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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'
|
||||
|
||||
480
model_manager/workflows/train_model.py
Normal file
480
model_manager/workflows/train_model.py
Normal file
@@ -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),
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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
|
||||
|
||||
672
tests/workflows/test_train_model.py
Normal file
672
tests/workflows/test_train_model.py
Normal file
@@ -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'
|
||||
Reference in New Issue
Block a user