SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.

This commit is contained in:
Bruno Domingues
2025-10-22 15:37:56 -03:00
parent f2a1c88ff3
commit 5789a13023
31 changed files with 37878 additions and 6097 deletions

View File

@@ -1,494 +0,0 @@
from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest
@pytest.fixture
def mock_env_vars(monkeypatch):
"""Fixture to set up environment variables for tests."""
monkeypatch.setenv('POD_ID', 'test-pod-123')
monkeypatch.setenv('TEMPORAL_HOST', 'test-temporal:7233')
monkeypatch.setenv('TEMPORAL_NAMESPACE', 'test-namespace')
monkeypatch.setenv('HTTP_METRICS_PORT', '9090')
monkeypatch.setenv('HTTP_SDK_METRICS_PORT', '9091')
monkeypatch.setenv('PROJECT_NAME', 'test-project')
monkeypatch.setenv('POSTGRES_HOST', 'localhost')
monkeypatch.setenv('POSTGRES_PORT', '5432')
monkeypatch.setenv('POSTGRES_USER', 'test')
monkeypatch.setenv('POSTGRES_PASSWORD', 'test')
monkeypatch.setenv('POSTGRES_DBNAME', 'test')
monkeypatch.setenv('MLFLOW_HOST', 'http://localhost')
monkeypatch.setenv('MLFLOW_PORT', '5000')
monkeypatch.setenv('MLFLOW_USERNAME', 'test')
monkeypatch.setenv('MLFLOW_PASSWORD', 'test')
monkeypatch.setenv('MINIO_ENDPOINT_URL', 'http://localhost:9000')
monkeypatch.setenv('MINIO_ACCESS_KEY', 'test')
monkeypatch.setenv('MINIO_SECRET_KEY', 'test')
monkeypatch.setenv('MONGODB_USERNAME', 'test')
monkeypatch.setenv('MONGODB_PASSWORD', 'test')
monkeypatch.setenv('MONGODB_URL', 'localhost:27017')
monkeypatch.setenv('MONGODB_DATABASE_NAME', 'test')
@patch('model_manager.worker.worker.start_http_server')
@patch('model_manager.worker.worker.metrics')
def test_start_prometheus_server_success(mock_metrics, mock_start_http_server, mock_env_vars):
"""Test successful Prometheus server startup."""
# Import after patching to ensure mocks are in place
from model_manager.worker.worker import start_prometheus_server
# Act
start_prometheus_server()
# Assert
mock_start_http_server.assert_called_once_with(9090)
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(1)
@patch('model_manager.worker.worker.start_http_server')
@patch('model_manager.worker.worker.os._exit')
def test_start_prometheus_server_failure(mock_exit, mock_start_http_server, mock_env_vars):
"""Test Prometheus server startup failure."""
# Arrange
mock_start_http_server.side_effect = Exception('Port already in use')
# Import after patching
from model_manager.worker.worker import start_prometheus_server
# Act
start_prometheus_server()
# Assert
mock_start_http_server.assert_called_once_with(9090)
mock_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_success(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test successful main function execution."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete successfully
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert
mock_start_prometheus.assert_called_once()
mock_notification_handler.assert_called_once()
mock_activities.assert_called_once()
mock_client_connect.assert_called_once_with(
target_host='test-temporal:7233',
namespace='test-namespace',
runtime=ANY,
)
assert mock_worker.call_count == 1 # Only one worker created
mock_gather.assert_called_once()
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
async def test_main_exception_handling(
mock_metrics,
mock_sys_exit,
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_env_vars,
):
"""Test main function exception handling and cleanup."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to raise an exception
mock_gather.side_effect = Exception('Worker failed')
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - Verify cleanup was performed
mock_logger.custom_error.assert_called_once()
mock_handler.shutdown.assert_called_once()
mock_activities_instance.shutdown.assert_called_once()
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(0)
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_creates_only_one_worker(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main creates only one worker with correct configurations."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify only one worker was created
assert mock_worker.call_count == 1
# Verify worker (train_model-queue)
first_call = mock_worker.call_args_list[0]
assert first_call[1]['task_queue'] == 'train_model-queue'
assert 'TrainModel' in str(first_call[1]['workflows'])
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_initializes_activities_with_configs(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main initializes Activities with correct configurations."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify Activities was initialized with correct parameters
mock_activities.assert_called_once()
call_kwargs = mock_activities.call_args[1]
assert 'postgres_config' in call_kwargs
assert 'mlflow_config' in call_kwargs
assert 'minio_config' in call_kwargs
assert call_kwargs['logger'] == mock_logger
assert call_kwargs['notification_handler'] == mock_handler
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_uses_environment_variables(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test that main uses environment variables correctly."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Import
from model_manager.worker.worker import main
# Mock gather to prevent infinite wait
with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock):
# Act
await main()
# Assert - Verify environment variables were used
mock_client_connect.assert_called_once_with(
target_host='test-temporal:7233',
namespace='test-namespace',
runtime=ANY,
)
mock_notification_handler.assert_called_once()
notification_call_kwargs = mock_notification_handler.call_args[1]
assert notification_call_kwargs['project_name'] == 'test-project'
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_cleanup_with_none_notification_handler(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test cleanup when notification_handler is None (line 194 branch False)."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
# Return None for notification_handler
mock_notification_handler.return_value = None
mock_activities_instance = MagicMock()
mock_activities_instance.shutdown = AsyncMock()
mock_activities.return_value = mock_activities_instance
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - notification_handler.shutdown() should NOT be called (line 194 False)
# Since notification_handler is None, we can't call shutdown on it
mock_activities_instance.shutdown.assert_called_once()
mock_sys_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.sys.exit')
@patch('model_manager.worker.worker.metrics')
@patch('model_manager.worker.worker.asyncio.gather')
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client.connect')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.get_logger')
async def test_main_cleanup_with_falsy_activities(
mock_get_logger,
mock_start_prometheus,
mock_notification_handler,
mock_activities,
mock_runtime,
mock_client_connect,
mock_worker,
mock_gather,
mock_metrics,
mock_sys_exit,
mock_env_vars,
):
"""Test cleanup when activities evaluates to False (line 196 branch False)."""
# Arrange
mock_logger = MagicMock()
mock_get_logger.return_value = mock_logger
mock_handler = MagicMock()
mock_notification_handler.return_value = mock_handler
# Create a falsy activities object (empty list, 0, False, etc.)
# Using an object that evaluates to False but doesn't cause AttributeError
class FalsyActivities:
def __bool__(self):
return False
def __getattr__(self, name):
# Return mock methods to avoid AttributeError during worker creation
return MagicMock()
falsy_activities = FalsyActivities()
mock_activities.return_value = falsy_activities
mock_temporal_client = AsyncMock()
mock_client_connect.return_value = mock_temporal_client
mock_worker_instance = MagicMock()
mock_worker_instance.run = MagicMock(return_value=AsyncMock())
mock_worker.return_value = mock_worker_instance
# Mock gather to complete
mock_gather.return_value = None
# Import and run
from model_manager.worker.worker import main
# Act
await main()
# Assert - notification_handler.shutdown() is called, but activities.shutdown() is NOT
mock_handler.shutdown.assert_called_once()
# activities is falsy, so shutdown should NOT be called
mock_sys_exit.assert_called_once_with(1)