Files
sientia-dataops-model-manager/tests/worker/test_worker.py

556 lines
19 KiB
Python
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Unit tests for worker module."""
import asyncio
import os
import sys
from unittest.mock import AsyncMock, Mock, patch
import pytest
@pytest.fixture
def mock_env_vars():
"""Set up test environment variables."""
env_vars = {
'POD_ID': 'test-pod-123',
'HTTP_METRICS_PORT': '9090',
'HTTP_SDK_METRICS_PORT': '9091',
'TEMPORAL_HOST': 'localhost:7233',
'TEMPORAL_NAMESPACE': 'test-namespace',
'PROJECT_NAME': 'test-project',
}
with patch.dict(os.environ, env_vars, clear=False):
yield env_vars
@pytest.fixture
def mock_logger():
"""Create a mock logger."""
logger = Mock()
logger.custom_info = Mock()
logger.custom_error = Mock()
return logger
@pytest.fixture
def mock_temporal_client():
"""Create a mock Temporal client."""
client_mock = AsyncMock()
client_mock.connect = AsyncMock()
return client_mock
@pytest.fixture
def mock_worker():
"""Create a mock Temporal worker."""
worker_mock = Mock()
worker_mock.run = AsyncMock(return_value=None)
return worker_mock
@pytest.fixture
def mock_notification_handler():
"""Create a mock notification handler."""
handler = Mock()
handler.shutdown = Mock()
return handler
@pytest.fixture
def mock_activities():
"""Create a mock Activities instance."""
activities = AsyncMock()
activities.update_experiment_run = Mock()
activities.validate_train_params = Mock()
activities.train_model = Mock()
activities.cleanup_resources = Mock()
activities.shutdown = AsyncMock()
return activities
def test_pod_id_from_env():
"""Test that POD_ID is correctly read from environment."""
with patch.dict(os.environ, {'POD_ID': 'pod-test-123'}):
# Re-import to get new env value
import importlib
import model_manager.worker.worker as worker_module
importlib.reload(worker_module)
assert worker_module.POD_ID == 'pod-test-123'
def test_sdk_metrics_port_default():
"""Test that SDK_METRICS_PORT uses default value."""
with patch.dict(os.environ, {}, clear=True):
import importlib
import model_manager.worker.worker as worker_module
importlib.reload(worker_module)
assert worker_module.SDK_METRICS_PORT == 9091
def test_sdk_metrics_port_from_env():
"""Test that SDK_METRICS_PORT is read from environment."""
with patch.dict(os.environ, {'HTTP_SDK_METRICS_PORT': '8888'}):
import importlib
import model_manager.worker.worker as worker_module
importlib.reload(worker_module)
assert worker_module.SDK_METRICS_PORT == 8888
@patch('model_manager.worker.worker.POD_ID', 'test-pod-123')
@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."""
from model_manager.worker.worker import start_prometheus_server
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
start_prometheus_server()
# Verify HTTP server started
mock_start_http_server.assert_called_once_with(9090)
# Verify APP_UP metric was set to 1
mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123')
mock_app_up.set.assert_called_once_with(1)
@patch('model_manager.worker.worker.start_http_server')
@patch('model_manager.worker.worker.metrics')
def test_start_prometheus_server_custom_port(mock_metrics, mock_start_http_server):
"""Test Prometheus server startup with custom port."""
from model_manager.worker.worker import start_prometheus_server
with patch.dict(os.environ, {'HTTP_METRICS_PORT': '8080', 'POD_ID': 'custom-pod'}):
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
start_prometheus_server()
mock_start_http_server.assert_called_once_with(8080)
@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
):
"""Test Prometheus server startup failure."""
from model_manager.worker.worker import start_prometheus_server
mock_start_http_server.side_effect = OSError('Port already in use')
start_prometheus_server()
# Verify exit was called with code 1
mock_exit.assert_called_once_with(1)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.build_mongodb_config')
@patch('model_manager.worker.worker.build_postgres_config')
@patch('model_manager.worker.worker.build_mlflow_config')
@patch('model_manager.worker.worker.build_minio_config')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.metrics')
async def test_main_successful_startup(
mock_metrics,
mock_start_prometheus,
mock_get_logger,
mock_build_minio,
mock_build_mlflow,
mock_build_postgres,
mock_build_mongodb,
mock_notification_handler_class,
mock_activities_class,
mock_runtime_class,
mock_client_class,
mock_worker_class,
mock_env_vars,
mock_logger,
mock_temporal_client,
mock_worker,
mock_notification_handler,
mock_activities,
):
"""Test successful main() execution until workers start."""
from model_manager.worker.worker import main
# Setup mocks
mock_get_logger.return_value = mock_logger
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 = {}
mock_notification_handler_class.return_value = mock_notification_handler
mock_activities_class.return_value = mock_activities
mock_runtime = Mock()
mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock()
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_class.return_value = mock_worker_instance
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
# Run main() and expect it to exit due to CancelledError
with pytest.raises(SystemExit) as exc_info:
await main()
assert exc_info.value.code == 1
# Verify all initialization steps were called
mock_get_logger.assert_called_once()
mock_start_prometheus.assert_called_once()
mock_notification_handler_class.assert_called_once()
mock_activities_class.assert_called_once()
mock_client_class.connect.assert_called_once()
# Agora sao criados dois Workers: um para train_model-queue e outro para cleanup-queue
assert mock_worker_class.call_count == 2
# Verify cleanup was performed
mock_notification_handler.shutdown.assert_called_once()
mock_activities.shutdown.assert_called_once()
mock_app_up.set.assert_called_with(0)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.build_mongodb_config')
@patch('model_manager.worker.worker.build_postgres_config')
@patch('model_manager.worker.worker.build_mlflow_config')
@patch('model_manager.worker.worker.build_minio_config')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.metrics')
async def test_main_handles_exception(
mock_metrics,
mock_start_prometheus,
mock_get_logger,
mock_build_minio,
mock_build_mlflow,
mock_build_postgres,
mock_build_mongodb,
mock_notification_handler_class,
mock_activities_class,
mock_runtime_class,
mock_client_class,
mock_worker_class,
mock_env_vars,
mock_logger,
):
"""Test main() handles exceptions and performs cleanup."""
from model_manager.worker.worker import main
# Setup mocks
mock_get_logger.return_value = mock_logger
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 = {}
mock_notification_handler = Mock()
mock_notification_handler.shutdown = Mock()
mock_notification_handler_class.return_value = mock_notification_handler
mock_activities = AsyncMock()
mock_activities.shutdown = AsyncMock()
mock_activities_class.return_value = mock_activities
mock_runtime = Mock()
mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock()
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock()
mock_worker_instance.run = AsyncMock(side_effect=RuntimeError('Worker failed'))
mock_worker_class.return_value = mock_worker_instance
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
# Run main() and expect SystemExit
with pytest.raises(SystemExit) as exc_info:
await main()
assert exc_info.value.code == 1
# Verify error was logged
mock_logger.custom_error.assert_called_once()
assert 'Worker failed' in str(mock_logger.custom_error.call_args)
# Verify cleanup was performed
mock_notification_handler.shutdown.assert_called_once()
mock_activities.shutdown.assert_called_once()
mock_app_up.set.assert_called_with(0)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.build_mongodb_config')
@patch('model_manager.worker.worker.build_postgres_config')
@patch('model_manager.worker.worker.build_mlflow_config')
@patch('model_manager.worker.worker.build_minio_config')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.metrics')
async def test_main_temporal_client_configuration(
mock_metrics,
mock_start_prometheus,
mock_get_logger,
mock_build_minio,
mock_build_mlflow,
mock_build_postgres,
mock_build_mongodb,
mock_notification_handler_class,
mock_activities_class,
mock_runtime_class,
mock_client_class,
mock_worker_class,
mock_logger,
):
"""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'},
):
# Setup mocks
mock_get_logger.return_value = mock_logger
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 = {}
mock_notification_handler = Mock()
mock_notification_handler.shutdown = Mock()
mock_notification_handler_class.return_value = mock_notification_handler
mock_activities = AsyncMock()
mock_activities.shutdown = AsyncMock()
mock_activities_class.return_value = mock_activities
mock_runtime = Mock()
mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock()
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock()
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
mock_worker_class.return_value = mock_worker_instance
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
# Run main()
with pytest.raises(SystemExit):
await main()
# Verify Temporal client was configured with correct parameters
mock_client_class.connect.assert_called_once_with(
target_host='temporal.example.com:7233', namespace='production', runtime=mock_runtime
)
@pytest.mark.asyncio
@patch('model_manager.worker.worker.Worker')
@patch('model_manager.worker.worker.client.Client')
@patch('model_manager.worker.worker.Runtime')
@patch('model_manager.worker.worker.Activities')
@patch('model_manager.worker.worker.NotificationHandler')
@patch('model_manager.worker.worker.build_mongodb_config')
@patch('model_manager.worker.worker.build_postgres_config')
@patch('model_manager.worker.worker.build_mlflow_config')
@patch('model_manager.worker.worker.build_minio_config')
@patch('model_manager.worker.worker.get_logger')
@patch('model_manager.worker.worker.start_prometheus_server')
@patch('model_manager.worker.worker.metrics')
async def test_main_worker_configuration(
mock_metrics,
mock_start_prometheus,
mock_get_logger,
mock_build_minio,
mock_build_mlflow,
mock_build_postgres,
mock_build_mongodb,
mock_notification_handler_class,
mock_activities_class,
mock_runtime_class,
mock_client_class,
mock_worker_class,
mock_env_vars,
mock_logger,
):
"""Test that Temporal worker is configured with correct parameters."""
from model_manager.worker.worker import main
# Setup mocks
mock_get_logger.return_value = mock_logger
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 = {}
mock_notification_handler = Mock()
mock_notification_handler_class.return_value = mock_notification_handler
mock_activities = AsyncMock()
mock_activities.update_experiment_run = Mock()
mock_activities.validate_train_params = Mock()
mock_activities.train_model = Mock()
mock_activities.cleanup_resources = Mock()
mock_activities.shutdown = AsyncMock()
mock_activities_class.return_value = mock_activities
mock_runtime = Mock()
mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock()
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock()
mock_worker_instance.run = AsyncMock(side_effect=asyncio.CancelledError())
mock_worker_class.return_value = mock_worker_instance
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
# Run main()
with pytest.raises(SystemExit):
await main()
# Verify Worker was created with correct configuration
assert mock_worker_class.call_count == 2
# Primeira chamada: worker de treinamento (train_model-queue)
train_call_args = mock_worker_class.call_args_list[0]
assert train_call_args[0][0] == mock_client_instance # temporal_client
assert train_call_args[1]['task_queue'] == 'train_model-queue'
assert train_call_args[1]['max_concurrent_workflow_tasks'] == 10
assert train_call_args[1]['max_concurrent_activities'] == 10
assert train_call_args[1]['max_concurrent_local_activities'] == 10
assert train_call_args[1]['max_cached_workflows'] == 100
train_activities_list = train_call_args[1]['activities']
assert mock_activities.update_experiment_run in train_activities_list
assert mock_activities.validate_train_params in train_activities_list
assert mock_activities.train_model in train_activities_list
assert mock_activities.cleanup_resources in train_activities_list
# Segunda chamada: worker de cleanup (cleanup-queue)
cleanup_call_args = mock_worker_class.call_args_list[1]
assert cleanup_call_args[0][0] == mock_client_instance # temporal_client
assert cleanup_call_args[1]['task_queue'] == 'cleanup-queue'
assert cleanup_call_args[1]['max_concurrent_workflow_tasks'] == 20
assert cleanup_call_args[1]['max_concurrent_activities'] == 20
assert cleanup_call_args[1]['max_concurrent_local_activities'] == 20
assert cleanup_call_args[1]['max_cached_workflows'] == 100
@patch('model_manager.worker.worker.asyncio.run')
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 model_manager.worker.worker as worker_module
# Simulate running the module
worker_module.main = AsyncMock()
# This would normally be called by asyncio.run(main())
# We just verify the pattern is correct
assert callable(worker_module.main)
def test_worker_module_docstring():
"""Test that worker module has comprehensive documentation."""
import model_manager.worker.worker as worker_module
assert worker_module.__doc__ is not None
assert 'Temporal' in worker_module.__doc__
assert 'worker' in worker_module.__doc__
@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
):
"""Test that start_prometheus_server prints success message."""
from model_manager.worker.worker import start_prometheus_server
mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up
start_prometheus_server()
captured = capsys.readouterr()
assert 'Prometheus server started on port 9090' in captured.out
@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
):
"""Test that start_prometheus_server prints failure message."""
from model_manager.worker.worker import start_prometheus_server
mock_start_http_server.side_effect = Exception('Test error')
start_prometheus_server()
captured = capsys.readouterr()
assert 'Failed to start Prometheus server' in captured.out
assert 'Test error' in captured.out