SIENTIAPDE-1241: Fixed formatting and linted code
This commit is contained in:
@@ -4,10 +4,6 @@ This module tests the Prometheus metrics configuration used for
|
||||
monitoring and observability in the Sientia DataOps Model Manager.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_app_up_metric_exists():
|
||||
"""Test that APP_UP metric is properly defined."""
|
||||
@@ -15,7 +11,9 @@ def test_app_up_metric_exists():
|
||||
|
||||
assert APP_UP is not None
|
||||
assert APP_UP._name == 'app_up'
|
||||
assert APP_UP._documentation == 'Indicates if the application is running (1) or shutting down (0)'
|
||||
assert (
|
||||
APP_UP._documentation == 'Indicates if the application is running (1) or shutting down (0)'
|
||||
)
|
||||
|
||||
|
||||
def test_app_up_metric_has_pod_id_label():
|
||||
@@ -75,11 +73,11 @@ def test_app_up_metric_multiple_pods():
|
||||
|
||||
def test_app_up_metric_default_value():
|
||||
"""Test that APP_UP metric starts with no value set."""
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
# Create a new label that hasn't been used yet
|
||||
import uuid
|
||||
|
||||
from model_manager.metrics import APP_UP
|
||||
|
||||
unique_pod = f'test-pod-{uuid.uuid4()}'
|
||||
|
||||
# The metric should exist but not have a value until set
|
||||
@@ -204,8 +202,8 @@ def test_app_up_metric_thread_safety():
|
||||
|
||||
def test_prometheus_client_gauge_import():
|
||||
"""Test that Gauge is properly imported from prometheus_client."""
|
||||
from model_manager.metrics import Gauge
|
||||
from prometheus_client import Gauge as PrometheusGauge
|
||||
|
||||
assert Gauge is PrometheusGauge
|
||||
from model_manager.metrics import Gauge
|
||||
|
||||
assert Gauge is PrometheusGauge
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Unit tests for StorageRepository class."""
|
||||
|
||||
from io import BytesIO
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from botocore.exceptions import ClientError
|
||||
@@ -462,4 +462,3 @@ def test_fetch_file_logs_file_size(mock_boto3, mock_logger, storage_config):
|
||||
# Verify logging includes file size
|
||||
log_calls = [str(call) for call in mock_logger.info.call_args_list]
|
||||
assert any('12345 bytes' in str(call) for call in log_calls)
|
||||
|
||||
|
||||
@@ -496,12 +496,14 @@ class TestTrain:
|
||||
):
|
||||
"""Test basic training workflow."""
|
||||
# Mock load_data to return a DataFrame
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
# Mock split_train_test to return train/test splits
|
||||
@@ -540,12 +542,14 @@ class TestTrain:
|
||||
"""Test training with scaler enabled."""
|
||||
sample_params.use_scaler = True
|
||||
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
@@ -567,12 +571,14 @@ class TestTrain:
|
||||
"""Test training with shuffle enabled."""
|
||||
sample_params.shuffle = True
|
||||
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
@@ -581,7 +587,7 @@ class TestTrain:
|
||||
y_test = pd.Series([25, 30], name='target')
|
||||
mock_split_train_test.return_value = (x_train, x_test, y_train, y_test)
|
||||
|
||||
result = training_repo.train(sample_csv_data, sample_params)
|
||||
training_repo.train(sample_csv_data, sample_params)
|
||||
|
||||
# Verify split was called with shuffle=True
|
||||
call_kwargs = mock_split_train_test.call_args[1]
|
||||
@@ -595,12 +601,14 @@ class TestTrain:
|
||||
"""Test training with different train size."""
|
||||
sample_params.train_size = 70
|
||||
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
@@ -609,7 +617,7 @@ class TestTrain:
|
||||
y_test = pd.Series([25, 30], name='target')
|
||||
mock_split_train_test.return_value = (x_train, x_test, y_train, y_test)
|
||||
|
||||
result = training_repo.train(sample_csv_data, sample_params)
|
||||
training_repo.train(sample_csv_data, sample_params)
|
||||
|
||||
# Verify split was called with train_size=0.7
|
||||
call_kwargs = mock_split_train_test.call_args[1]
|
||||
@@ -622,12 +630,14 @@ class TestTrain:
|
||||
):
|
||||
"""Test that ValueError is raised when transformed data is empty."""
|
||||
# Mock load_data to return empty DataFrame
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [],
|
||||
'var2': [],
|
||||
'var3': [],
|
||||
'target': [],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [],
|
||||
'var2': [],
|
||||
'var3': [],
|
||||
'target': [],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
with pytest.raises(ValueError, match='Data view is empty after transformation'):
|
||||
@@ -645,12 +655,14 @@ class TestTrain:
|
||||
sample_csv_data,
|
||||
):
|
||||
"""Test that training success is logged."""
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
@@ -676,12 +688,14 @@ class TestTrain:
|
||||
sample_params.line_separator = ';'
|
||||
sample_params.decimal_separator = ','
|
||||
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
@@ -701,12 +715,14 @@ class TestTrain:
|
||||
self, mock_load_data, mock_split_train_test, training_repo, sample_params, sample_csv_data
|
||||
):
|
||||
"""Test that TrainModelResult contains all expected fields."""
|
||||
mock_df = pd.DataFrame({
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
})
|
||||
mock_df = pd.DataFrame(
|
||||
{
|
||||
'var1': [1, 2, 3, 4, 5],
|
||||
'var2': [2, 3, 4, 5, 6],
|
||||
'var3': [3, 4, 5, 6, 7],
|
||||
'target': [10, 15, 20, 25, 30],
|
||||
}
|
||||
)
|
||||
mock_load_data.return_value = mock_df
|
||||
|
||||
x_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [2, 3, 4], 'var3': [3, 4, 5]})
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -144,7 +144,9 @@ def test_start_prometheus_server_custom_port(mock_metrics, mock_start_http_serve
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.os._exit')
|
||||
def test_start_prometheus_server_failure(mock_exit, mock_metrics, mock_start_http_server, mock_env_vars):
|
||||
def test_start_prometheus_server_failure(
|
||||
mock_exit, mock_metrics, mock_start_http_server, mock_env_vars
|
||||
):
|
||||
"""Test Prometheus server startup failure."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
@@ -194,7 +196,10 @@ async def test_main_successful_startup(
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {'connection_string': 'mongodb://test', 'database_name': 'test_db'}
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
@@ -209,7 +214,9 @@ async def test_main_successful_startup(
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError()) # Simulate interruption
|
||||
mock_worker_instance.run = AsyncMock(
|
||||
side_effect=asyncio.CancelledError()
|
||||
) # Simulate interruption
|
||||
mock_worker_class.return_value = mock_worker_instance
|
||||
|
||||
mock_app_up = Mock()
|
||||
@@ -269,7 +276,10 @@ async def test_main_handles_exception(
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {'connection_string': 'mongodb://test', 'database_name': 'test_db'}
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
@@ -342,10 +352,16 @@ async def test_main_temporal_client_configuration(
|
||||
"""Test that Temporal client is configured correctly."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
with patch.dict(os.environ, {'TEMPORAL_HOST': 'temporal.example.com:7233', 'TEMPORAL_NAMESPACE': 'production'}):
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{'TEMPORAL_HOST': 'temporal.example.com:7233', 'TEMPORAL_NAMESPACE': 'production'},
|
||||
):
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {'connection_string': 'mongodb://test', 'database_name': 'test_db'}
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
@@ -415,7 +431,10 @@ async def test_main_worker_configuration(
|
||||
|
||||
# Setup mocks
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {'connection_string': 'mongodb://test', 'database_name': 'test_db'}
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
@@ -472,8 +491,6 @@ def test_main_entrypoint(mock_asyncio_run):
|
||||
"""Test the __main__ entrypoint."""
|
||||
# Import and execute the main block
|
||||
with patch.object(sys, 'argv', ['worker.py']):
|
||||
import importlib
|
||||
|
||||
import model_manager.worker.worker as worker_module
|
||||
|
||||
# Simulate running the module
|
||||
@@ -495,7 +512,9 @@ def test_worker_module_docstring():
|
||||
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
def test_start_prometheus_server_prints_success(mock_metrics, mock_start_http_server, capsys, mock_env_vars):
|
||||
def test_start_prometheus_server_prints_success(
|
||||
mock_metrics, mock_start_http_server, capsys, mock_env_vars
|
||||
):
|
||||
"""Test that start_prometheus_server prints success message."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
@@ -511,7 +530,9 @@ def test_start_prometheus_server_prints_success(mock_metrics, mock_start_http_se
|
||||
@patch('model_manager.worker.worker.start_http_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
@patch('model_manager.worker.worker.os._exit')
|
||||
def test_start_prometheus_server_prints_failure(mock_exit, mock_metrics, mock_start_http_server, capsys, mock_env_vars):
|
||||
def test_start_prometheus_server_prints_failure(
|
||||
mock_exit, mock_metrics, mock_start_http_server, capsys, mock_env_vars
|
||||
):
|
||||
"""Test that start_prometheus_server prints failure message."""
|
||||
from model_manager.worker.worker import start_prometheus_server
|
||||
|
||||
@@ -522,5 +543,3 @@ def test_start_prometheus_server_prints_failure(mock_exit, mock_metrics, mock_st
|
||||
captured = capsys.readouterr()
|
||||
assert 'Failed to start Prometheus server' in captured.out
|
||||
assert 'Test error' in captured.out
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""Unit tests for TrainModel workflow."""
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
from model_manager.utils.exceptions import ModelTrainingError
|
||||
from model_manager.utils.models.experiment_status import ExperimentStatus
|
||||
@@ -116,7 +117,7 @@ def test_extract_error_message_with_cause():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
workflow_instance = TrainModel()
|
||||
|
||||
|
||||
# Create exception chain
|
||||
cause = ValueError('Root cause')
|
||||
exc = RuntimeError('Outer error')
|
||||
@@ -146,7 +147,7 @@ def test_extract_error_message_circular_reference():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
workflow_instance = TrainModel()
|
||||
|
||||
|
||||
# Create circular reference
|
||||
exc1 = ValueError('Error 1')
|
||||
exc2 = ValueError('Error 2')
|
||||
@@ -165,7 +166,7 @@ def test_extract_error_message_duplicate_messages():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
workflow_instance = TrainModel()
|
||||
|
||||
|
||||
# Create chain with duplicate messages
|
||||
exc1 = ValueError('Same error')
|
||||
exc2 = ValueError('Same error')
|
||||
@@ -179,7 +180,9 @@ def test_extract_error_message_duplicate_messages():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_success(mock_workflow_module, sample_input_data, mock_train_params):
|
||||
async def test_validate_training_parameters_success(
|
||||
mock_workflow_module, sample_input_data, mock_train_params
|
||||
):
|
||||
"""Test successful parameter validation."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
@@ -189,9 +192,7 @@ async def test_validate_training_parameters_success(mock_workflow_module, sample
|
||||
workflow_instance = TrainModel()
|
||||
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
|
||||
|
||||
result = await workflow_instance._validate_training_parameters(
|
||||
sample_input_data, 123, metadata
|
||||
)
|
||||
result = await workflow_instance._validate_training_parameters(sample_input_data, 123, metadata)
|
||||
|
||||
assert result == mock_train_params
|
||||
|
||||
@@ -214,9 +215,7 @@ async def test_validate_training_parameters_failure(mock_workflow_module, sample
|
||||
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
|
||||
|
||||
with pytest.raises(ValueError, match='Invalid params'):
|
||||
await workflow_instance._validate_training_parameters(
|
||||
sample_input_data, 123, metadata
|
||||
)
|
||||
await workflow_instance._validate_training_parameters(sample_input_data, 123, metadata)
|
||||
|
||||
# Verify error status update was called
|
||||
assert mock_workflow_module.execute_activity_method.call_count == 2
|
||||
@@ -231,7 +230,7 @@ async def test_train_model_success(mock_workflow_module, mock_train_params):
|
||||
# Setup mocks
|
||||
train_result = {
|
||||
'run_name': 'test-run-123',
|
||||
'run_dir': '/tmp/test-run',
|
||||
'run_dir': '/tmp/test-run', # noqa: S108
|
||||
'mse_val': 0.5,
|
||||
'r2_val': 0.9,
|
||||
}
|
||||
@@ -276,10 +275,10 @@ async def test_train_model_mlflow_error(mock_workflow_module, mock_train_params)
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
# Setup mocks - MLflow save fails
|
||||
mlflow_error = ModelTrainingError(model_trained=True, model_saved=False, message='MLflow save failed')
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[mlflow_error, None]
|
||||
mlflow_error = ModelTrainingError(
|
||||
model_trained=True, model_saved=False, message='MLflow save failed'
|
||||
)
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(side_effect=[mlflow_error, None])
|
||||
|
||||
workflow_instance = TrainModel()
|
||||
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
|
||||
@@ -308,7 +307,7 @@ async def test_cleanup_resources_success(mock_workflow_module):
|
||||
|
||||
await workflow_instance._cleanup_resources(
|
||||
experiment_run_id=123,
|
||||
run_dir='/tmp/test-run',
|
||||
run_dir='/tmp/test-run', # noqa: S108
|
||||
bucket_name='test-bucket',
|
||||
file_name='test-file.csv',
|
||||
metadata=metadata,
|
||||
@@ -334,7 +333,7 @@ async def test_cleanup_resources_failure(mock_workflow_module):
|
||||
with pytest.raises(RuntimeError, match='Cleanup failed'):
|
||||
await workflow_instance._cleanup_resources(
|
||||
experiment_run_id=123,
|
||||
run_dir='/tmp/test-run',
|
||||
run_dir='/tmp/test-run', # noqa: S108
|
||||
bucket_name='test-bucket',
|
||||
file_name='test-file.csv',
|
||||
metadata=metadata,
|
||||
@@ -348,8 +347,8 @@ async def test_cleanup_resources_failure(mock_workflow_module):
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_update_experiment_run_status_only(mock_workflow_module):
|
||||
"""Test update_experiment_run with status only."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
from model_manager.activities.experiment_tracking import UpdateType
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_workflow_module.execute_activity_method = AsyncMock()
|
||||
|
||||
@@ -375,8 +374,8 @@ async def test_update_experiment_run_status_only(mock_workflow_module):
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_update_experiment_run_with_error(mock_workflow_module):
|
||||
"""Test update_experiment_run with error message."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
from model_manager.activities.experiment_tracking import UpdateType
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_workflow_module.execute_activity_method = AsyncMock()
|
||||
|
||||
@@ -400,8 +399,8 @@ async def test_update_experiment_run_with_error(mock_workflow_module):
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_update_experiment_run_with_run_name(mock_workflow_module):
|
||||
"""Test update_experiment_run with run_name."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
from model_manager.activities.experiment_tracking import UpdateType
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_workflow_module.execute_activity_method = AsyncMock()
|
||||
|
||||
@@ -424,24 +423,26 @@ async def test_update_experiment_run_with_run_name(mock_workflow_module):
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
@patch('model_manager.workflows.train_model.POD_ID', 'test-pod-456')
|
||||
async def test_run_complete_workflow_success(mock_workflow_module, sample_input_data, mock_train_params):
|
||||
async def test_run_complete_workflow_success(
|
||||
mock_workflow_module, sample_input_data, mock_train_params
|
||||
):
|
||||
"""Test complete workflow execution success path."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
# Setup mocks for all activities
|
||||
train_result = {
|
||||
'run_name': 'test-run-123',
|
||||
'run_dir': '/tmp/test-run',
|
||||
'run_dir': '/tmp/test-run', # noqa: S108
|
||||
}
|
||||
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
mock_train_params, # validate_train_params
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
train_result, # train_model
|
||||
None, # update status (MLFLOW_SENT)
|
||||
None, # cleanup_resources
|
||||
None, # update status (FILE_DELETED)
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
train_result, # train_model
|
||||
None, # update status (MLFLOW_SENT)
|
||||
None, # cleanup_resources
|
||||
None, # update status (FILE_DELETED)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -464,7 +465,7 @@ async def test_run_workflow_validation_error(mock_workflow_module, sample_input_
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
ValueError('Invalid parameters'), # validate_train_params fails
|
||||
None, # update status (ORCHESTRATOR_VALIDATION_ERROR)
|
||||
None, # update status (ORCHESTRATOR_VALIDATION_ERROR)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -479,17 +480,19 @@ async def test_run_workflow_validation_error(mock_workflow_module, sample_input_
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_workflow_training_error(mock_workflow_module, sample_input_data, mock_train_params):
|
||||
async def test_run_workflow_training_error(
|
||||
mock_workflow_module, sample_input_data, mock_train_params
|
||||
):
|
||||
"""Test workflow handles training errors."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
# Setup mocks - training fails
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
mock_train_params, # validate_train_params
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
RuntimeError('Training failed'), # train_model fails
|
||||
None, # update status (TRAINING_ERROR)
|
||||
mock_train_params, # validate_train_params
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
RuntimeError('Training failed'), # train_model fails
|
||||
None, # update status (TRAINING_ERROR)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -503,24 +506,26 @@ async def test_run_workflow_training_error(mock_workflow_module, sample_input_da
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_workflow_cleanup_error(mock_workflow_module, sample_input_data, mock_train_params):
|
||||
async def test_run_workflow_cleanup_error(
|
||||
mock_workflow_module, sample_input_data, mock_train_params
|
||||
):
|
||||
"""Test workflow handles cleanup errors."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
# Setup mocks - cleanup fails
|
||||
train_result = {
|
||||
'run_name': 'test-run-123',
|
||||
'run_dir': '/tmp/test-run',
|
||||
'run_dir': '/tmp/test-run', # noqa: S108
|
||||
}
|
||||
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
mock_train_params, # validate_train_params
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
train_result, # train_model
|
||||
None, # update status (MLFLOW_SENT)
|
||||
RuntimeError('Cleanup failed'), # cleanup_resources fails
|
||||
None, # update status (FILE_DELETE_ERROR)
|
||||
mock_train_params, # validate_train_params
|
||||
None, # update status (MAGE_WAITING_PROC)
|
||||
train_result, # train_model
|
||||
None, # update status (MLFLOW_SENT)
|
||||
RuntimeError('Cleanup failed'), # cleanup_resources fails
|
||||
None, # update status (FILE_DELETE_ERROR)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -547,13 +552,13 @@ async def test_run_workflow_missing_experiment_run_id():
|
||||
def test_module_constants():
|
||||
"""Test that module-level constants are defined correctly."""
|
||||
from model_manager.workflows.train_model import (
|
||||
TIMEOUT_VALIDATE_PARAMS,
|
||||
TIMEOUT_TRAIN_MODEL,
|
||||
TIMEOUT_DELETE_FILE,
|
||||
TIMEOUT_TRAIN_MODEL,
|
||||
TIMEOUT_UPDATE_DATABASE,
|
||||
TIMEOUT_VALIDATE_PARAMS,
|
||||
database_retry_policy,
|
||||
network_retry_policy,
|
||||
no_retry_policy,
|
||||
database_retry_policy,
|
||||
)
|
||||
|
||||
# Verify timeouts are integers
|
||||
@@ -608,11 +613,11 @@ async def test_train_model_empty_run_dir(mock_workflow_module, mock_train_params
|
||||
mock_workflow_module.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
mock_train_params, # validate
|
||||
None, # update status
|
||||
train_result, # train
|
||||
None, # update status
|
||||
None, # cleanup
|
||||
None, # update status
|
||||
None, # update status
|
||||
train_result, # train
|
||||
None, # update status
|
||||
None, # cleanup
|
||||
None, # update status
|
||||
]
|
||||
)
|
||||
|
||||
@@ -644,4 +649,3 @@ async def test_train_model_empty_run_dir(mock_workflow_module, mock_train_params
|
||||
# Verify cleanup was called with empty string
|
||||
cleanup_call = mock_workflow_module.execute_activity_method.call_args_list[4]
|
||||
assert cleanup_call[0][1]['run_dir'] == ''
|
||||
|
||||
|
||||
Reference in New Issue
Block a user