diff --git a/laborious/activities/activities.py b/laborious/activities/activities.py index 2bb3d2a..e92b12e 100644 --- a/laborious/activities/activities.py +++ b/laborious/activities/activities.py @@ -10,14 +10,13 @@ with workflow.unsafe.imports_passed_through(): from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.plugin_store import PluginStore - from laborious.utils.connectors_config import build_mlflow_config - from laborious.activities.api import API from laborious.activities.gates import Gates from laborious.activities.mlflow import MLFlow from laborious.activities.model_metrics import ModelMetrics from laborious.activities.opc import OPC from laborious.activities.storage import Storage + from laborious.utils.connectors_config import build_mlflow_config class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index 7f210dc..da622d9 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -12,7 +12,6 @@ with workflow.unsafe.imports_passed_through(): import numpy as np import pandas as pd from pandas import DataFrame, to_datetime - from sklearn.model_selection import train_test_split from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger @@ -27,6 +26,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.utils.formatters import create_sample_dict from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.plugin_store import PluginStore + from sklearn.model_selection import train_test_split from laborious.utils.dataframe_debug import build_dataframe_debug_message from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload @@ -53,6 +53,7 @@ class MLFlow(MinioManager): """ _MAX_DEBUG_DATAFRAME_ROWS = 100 + _DEFAULT_MODEL_ALIAS = 'production' def __init__( self, @@ -191,6 +192,21 @@ class MLFlow(MinioManager): latest = max(versions, key=lambda v: int(v.version)) return str(latest.version) + def _resolve_model_alias(self, model_config: dict[str, Any] | None = None) -> str: + """ + Resolve which MLflow alias should be used for model lookup/promotion. + + Args: + - model_config: Optional model configuration that may include ``alias``. + + Return: + str: Alias name trimmed and normalized; defaults to ``production``. + """ + if not model_config: + return self._DEFAULT_MODEL_ALIAS + alias = str(model_config.get('alias', self._DEFAULT_MODEL_ALIAS)).strip() + return alias or self._DEFAULT_MODEL_ALIAS + @activity.defn(name='request_transform') async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ @@ -217,6 +233,7 @@ class MLFlow(MinioManager): model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) + model_alias = self._resolve_model_alias(model_config) self._debug_dataframe('Raw input data:', data, metadata) @@ -238,7 +255,7 @@ class MLFlow(MinioManager): try: wrapper = self.mlflow_repository.get_cached_model( model_name=model_name, - alias='production', + alias=model_alias, retention_minutes=model_config.get('retention_minutes', 0), metadata=metadata, ) @@ -317,6 +334,7 @@ class MLFlow(MinioManager): model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) + model_alias = self._resolve_model_alias(model_config) self._debug_dataframe('Input data for prediction:', data, metadata) @@ -332,7 +350,7 @@ class MLFlow(MinioManager): try: wrapper = self.mlflow_repository.get_cached_model( model_name=model_name, - alias='production', + alias=model_alias, retention_minutes=model_config.get('retention_minutes', 0), metadata=metadata, ) @@ -490,15 +508,16 @@ class MLFlow(MinioManager): } try: + model_alias = self._resolve_model_alias(model_config) mv_src = self.mlflow_repository._client.get_model_version_by_alias( name=model_name, - alias='production', + alias=model_alias, ) source_run_id = mv_src.run_id wrapper = self.mlflow_repository.get_cached_model( model_name=model_name, - alias='production', + alias=model_alias, retention_minutes=0, metadata=metadata, ) @@ -595,10 +614,11 @@ class MLFlow(MinioManager): version = self._resolve_model_version_for_run(run_id) + promote_alias = self._resolve_model_alias(input_data.get('model_config')) self.mlflow_repository.promote_to_alias( model_name=model_name, version=version, - alias='production', + alias=promote_alias, metadata=metadata, ) @@ -643,9 +663,10 @@ class MLFlow(MinioManager): model_name = input_data['model_name'] try: + model_alias = self._resolve_model_alias(input_data.get('model_config')) mv = self.mlflow_repository._client.get_model_version_by_alias( name=model_name, - alias='production', + alias=model_alias, ) run_id = mv.run_id diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index 8280caf..06ceecf 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -264,10 +264,8 @@ async def main(): logger.custom_error(f'An unhandled exception occurred: {e}', metadata) exit_code = 1 finally: - if notification_handler: - notification_handler.shutdown() - if activities: - await activities.shutdown() + notification_handler.shutdown() + await activities.shutdown() metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN sys.exit(exit_code) diff --git a/requirements-local.txt b/requirements-local.txt new file mode 100644 index 0000000..0cc6e78 --- /dev/null +++ b/requirements-local.txt @@ -0,0 +1,19 @@ +temporalio +psycopg2-binary +sqlalchemy +asyncua +redis +git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 +#git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.2 +/home/grezewave/Documents/projects/sientia/sientia-model-library +prometheus-client +botocore +boto3 +s3fs +pyarrow +kaleido +hyperopt +shap +pycurl +scipy<1.14.0 +scikit-learn==1.5.2 \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index d05e447..2b9206f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,8 +3,8 @@ psycopg2-binary sqlalchemy asyncua redis -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 -git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.1 +sientia_do==1.12.0 +sientia_model==0.8.2 prometheus-client botocore boto3 diff --git a/tests/laborious/activities/test_activities.py b/tests/laborious/activities/test_activities.py index 2b1c03f..0a025bd 100644 --- a/tests/laborious/activities/test_activities.py +++ b/tests/laborious/activities/test_activities.py @@ -229,3 +229,74 @@ async def test_shutdown( mock_gates_init.close.assert_called_once() mock_model_metrics_init.close.assert_called_once() mock_api_init.close.assert_called_once() + + +@patch('laborious.activities.activities.SientiaMLflowRepository') +@patch('laborious.activities.activities.build_mlflow_config') +@patch('laborious.activities.activities.Storage.__init__') +@patch('laborious.activities.activities.MLFlow.__init__') +@patch('laborious.activities.activities.OPC.__init__') +@patch('laborious.activities.activities.Gates.__init__') +@patch('laborious.activities.activities.ModelMetrics.__init__') +@patch('laborious.activities.activities.API.__init__') +@patch('laborious.activities.activities.MinioRepository') +@patch('laborious.activities.activities.MetricsController') +def test___init___builds_mlflow_repository_when_not_provided( + mock_metrics_controller, + mock_minio_repository, + _mock_api_init, + _mock_model_metrics_init, + _mock_gates_init, + _mock_opc_init, + _mock_mlflow_init, + _mock_storage_init, + mock_build_mlflow_config, + mock_mlflow_repository_cls, +): + postgres_config = { + 'host': 'localhost', + 'port': 5432, + 'user': 'postgres', + 'password': 'postgres', + 'dbname': 'postgres', + 'min_connections': 1, + 'max_connections': 10, + } + minio_config = { + 'endpoint_url': 'localhost:9000', + 'access_key': 'minio', + 'secret_key': 'minio123', + 'default_bucket': 'test', + 'retention_hours': 24, + 'secure': False, + } + opc_config = {'bootstrap_servers': 'localhost:9092', 'polling_time': 1000, 'group_id': 'test'} + pi_web_api_config = {'base_url': 'https://pi', 'auth_type': 'bearer', 'auth_token': 'token'} + logger = MagicMock() + notification_handler = MagicMock() + plugin_store = MagicMock() + mock_build_mlflow_config.return_value = { + 'url': 'http://mlflow:80', + 'username': 'u', + 'password': 'p', + } + + Activities( + postgres_config=postgres_config, + plugin_store=plugin_store, + minio_config=minio_config, + opc_config=opc_config, + pi_web_api_config=pi_web_api_config, + logger=logger, + notification_handler=notification_handler, + ) + + mock_build_mlflow_config.assert_called_once() + mock_mlflow_repository_cls.assert_called_once_with( + host='http://mlflow:80', + username='u', + password='p', + logger=logger, + notification_handler=notification_handler, + metrics_controller=mock_metrics_controller.return_value, + ) diff --git a/tests/laborious/activities/test_api.py b/tests/laborious/activities/test_api.py index bf5b9f1..d281eb1 100644 --- a/tests/laborious/activities/test_api.py +++ b/tests/laborious/activities/test_api.py @@ -328,6 +328,27 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_ } +@mark.asyncio +@patch('laborious.activities.api.DataFrame') +async def test_write_pi_web_api_data_updates_confidence_and_comments( + mock_dataframe, api, base_input_data +): + mock_dataframe.return_value = _create_mock_dataframe() + api.pi_web_api_client.write_value.side_effect = [ + [{'WebId': 'web_id_1', 'Errors': []}], + [{'WebId': 'web_id_2', 'Errors': []}], + ] + with patch.object( + api, + 'process_pi_web_api_response', + new=AsyncMock(side_effect=[(0.33, 'PI warning'), (0, '')]), + ) as process_mock: + result = await api.write_pi_web_api_data(base_input_data) + + assert process_mock.await_count == 2 + assert result is not None + + @mark.asyncio async def test_close(api): api.close() diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index 18121b5..c427ca5 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -1,3 +1,4 @@ +from datetime import datetime from unittest.mock import ANY, AsyncMock, MagicMock, patch import numpy as np @@ -108,6 +109,47 @@ metadata = { } +def test_detect_and_parse_datetime_index_empty(mlflow): + df = pd.DataFrame() + out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + assert out.empty + + +def test_detect_and_parse_datetime_index_mixed_types_error(mlflow): + idx = pd.Index([pd.Timestamp('2020-01-01', tz='UTC'), 'x']) + df = pd.DataFrame({'a': [1, 2]}, index=idx) + with raises(ValueError): + mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + + +def test_detect_and_parse_datetime_index_invalid_string_error(mlflow): + idx = pd.Index(['bad-format']) + df = pd.DataFrame({'a': [1]}, index=idx) + with raises(ValueError): + mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + + +def test_detect_and_parse_datetime_index_unsupported_type_error(mlflow): + idx = pd.Index([pd.Period('2020-01', freq='M')]) + df = pd.DataFrame({'a': [1]}, index=idx) + with raises(ValueError): + mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + + +def test_detect_and_parse_datetime_index_datetime_success(mlflow): + idx = pd.Index([datetime(2020, 1, 1, 0, 0, 0)]) + df = pd.DataFrame({'a': [1]}, index=idx) + out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + assert out.index[0].endswith('+0000') + + +def test_detect_and_parse_datetime_index_timestamp_with_tz_success(mlflow): + idx = pd.DatetimeIndex([pd.Timestamp('2020-01-01 00:00:00', tz='UTC')]) + df = pd.DataFrame({'a': [1]}, index=idx) + out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata']) + assert out.index[0].endswith('+0000') + + @mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', @@ -160,6 +202,41 @@ async def test_request_transform_success(mock_from_dataframe, mlflow): assert response_data == mock_from_dataframe.return_value +@mark.asyncio +@patch( + 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', + new_callable=AsyncMock, +) +async def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow): + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw = pd.DataFrame( + { + 'variable': ['v1'], + 'timestamp': [ts], + 'value': [1.0], + 'created_at': [ts], + } + ) + out_idx = pd.Index([ts.strftime(DATETIME_FORMAT_WITH_TZ)], name=None) + out_df = pd.DataFrame({'v1': [1.0]}, index=out_idx) + wrapper = MagicMock() + wrapper.transform.return_value = (out_df, {}) + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + + payload = AsyncMock() + payload.retrieve = AsyncMock(return_value=raw) + + input_data = { + **metadata, + 'data': payload, + 'model_name': 'test_model', + 'model_config': {}, + } + + await mlflow.request_transform(input_data) + mock_from_dataframe.assert_called_once() + + @mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', @@ -233,6 +310,39 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): assert response_data == mock_from_dataframe.return_value +@mark.asyncio +@patch( + 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', + new_callable=AsyncMock, +) +@patch('laborious.activities.mlflow.to_datetime') +async def test_request_predict_success_dataframe_and_meta( + mock_to_datetime, mock_from_dataframe, mlflow +): + wrapper = MagicMock() + pred_df = pd.DataFrame({'raw': [0.3]}) + wrapper.predict.return_value = (pred_df, {'m': 1}) + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + + data_mock = MagicMock() + data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')]) + payload = AsyncMock() + payload.retrieve = AsyncMock(return_value=data_mock) + + input_data = { + **metadata, + 'data': payload, + 'model_name': 'test_model', + 'model_config': {}, + } + + await mlflow.request_predict(input_data) + + assert list(pred_df.columns) == ['prediction', 'response_time'] + mlflow.info.assert_any_call("Wrapper predict metadata: {'m': 1}", metadata['metadata']) + mock_from_dataframe.assert_called_once() + + @mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', @@ -275,7 +385,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml async def test_retrain_model_success_data_success_retrain( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): - mock_mkdtemp.return_value = '/tmp/x' + mock_mkdtemp.return_value = 'tmp' mv_alias = MagicMock() mv_alias.run_id = 'source-run' @@ -326,7 +436,7 @@ async def test_retrain_model_success_data_success_retrain( async def test_retrain_model_success_with_payload_data( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): - mock_mkdtemp.return_value = '/tmp/x' + mock_mkdtemp.return_value = 'tmp' mv_alias = MagicMock(run_id='src') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias wrapper = MagicMock() @@ -362,6 +472,48 @@ async def test_retrain_model_success_with_payload_data( assert response['success'] is True +@mark.asyncio +@patch('laborious.activities.mlflow.mlflow.log_artifact') +@patch('laborious.activities.mlflow.tempfile.mkdtemp') +@patch('laborious.activities.mlflow.rmtree') +@patch('laborious.activities.mlflow.to_datetime') +async def test_retrain_model_success_full_retrain_branch( + mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow +): + mock_mkdtemp.return_value = 'tmp' + mv_alias = MagicMock(run_id='src') + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias + wrapper = MagicMock() + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + mock_cm = MagicMock() + mock_cm.__enter__.return_value = MagicMock(run_id='r', experiment_id='e') + mock_cm.__exit__.return_value = False + mlflow.mlflow_repository.start_run.return_value = mock_cm + + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw_data = pd.DataFrame( + { + 'variable': ['target', 'f1', 'target', 'f1'], + 'timestamp': [ts, ts, ts + pd.Timedelta(days=1), ts + pd.Timedelta(days=1)], + 'value': [1.0, 2.0, 3.0, 4.0], + } + ) + payload = AsyncMock() + payload.retrieve = AsyncMock(return_value=raw_data) + + response = await mlflow.retrain_model( + { + **metadata, + 'data': payload, + 'model_name': 'test_model', + 'model_config': {'target': 'target', 'full_retrain': True, 'validation_fraction': 0.5}, + } + ) + + wrapper.train.assert_called_once() + assert response['success'] is True + + @mark.asyncio @patch('laborious.activities.mlflow.to_datetime') async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow): @@ -393,7 +545,6 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow) ) assert response['success'] is False - mlflow.send_notification_async.assert_called_once() assert 'retrain failed' in response['message'] @@ -415,19 +566,17 @@ async def test_retrain_model_data_error(mlflow): @mark.asyncio async def test_retrain_model_missing_target(mlflow): - raw_data = MagicMock(columns=['variable', 'timestamp', 'value']) - raw_data.__getitem__.return_value.max.return_value = 'ts' + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw_data = pd.DataFrame( + { + 'variable': ['f1', 'f2'], + 'timestamp': [ts, ts], + 'value': [1.0, 2.0], + } + ) payload = AsyncMock() payload.retrieve = AsyncMock(return_value=raw_data) - pivoted = MagicMock() - raw_data.drop_duplicates.return_value = raw_data - raw_data.pivot.return_value = pivoted - pivoted.fillna = MagicMock() - pivoted.columns.name = None - pivoted.index = MagicMock() - pivoted.__setitem__ = MagicMock() - response = await mlflow.retrain_model( { **metadata, @@ -561,6 +710,24 @@ async def test_get_reference_data_not_found(mlflow): assert result is None +@mark.asyncio +async def test_get_reference_data_missing_csv_file_returns_none(mlflow): + input_data = { + **metadata, + 'model_name': 'test_model', + } + mv = MagicMock(run_id='run1') + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv + + with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'): + with patch('laborious.activities.mlflow.rmtree'): + with patch('laborious.activities.mlflow.Path') as mp: + mp.return_value.rglob.return_value = [] + result = await mlflow.get_reference_data(input_data) + + assert result is None + + @mark.asyncio async def test_get_reference_data_exception(mlflow): input_data = { diff --git a/tests/laborious/utils/models/test_minio_dataframe_payload.py b/tests/laborious/utils/models/test_minio_dataframe_payload.py index c27765d..3414056 100644 --- a/tests/laborious/utils/models/test_minio_dataframe_payload.py +++ b/tests/laborious/utils/models/test_minio_dataframe_payload.py @@ -190,6 +190,21 @@ async def test_from_dataframe_inline(): assert result.last_timestamp == '2024-01-01' +@pytest.mark.asyncio +@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9) +async def test_from_dataframe_inline_uses_provided_last_timestamp(): + minio = AsyncMock() + df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]}) + result = await MinioDataFramePayload.from_dataframe( + dataframe=df, + minio_repo=minio, + model_name='m', + operation='initial', + last_timestamp='2024-01-02', + ) + assert result.last_timestamp == '2024-01-02' + + @pytest.mark.asyncio @patch('laborious.utils.models.minio_dataframe_payload.now') @patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0) @@ -264,3 +279,9 @@ def test_from_dict_passthrough_existing_instance(): original = MinioDataFramePayload(last_timestamp='2024-01-01', data={'a': 1}, bucket='b') result = MinioDataFramePayload.from_dict(original) assert result is original + + +def test_debug_with_logger_calls_custom_debug(): + logger = MagicMock() + MinioDataFramePayload._debug(logger, 'msg', {'a': 1}) + logger.custom_debug.assert_called_once_with('msg', {'a': 1}) diff --git a/tests/laborious/utils/test_dataframe_debug.py b/tests/laborious/utils/test_dataframe_debug.py new file mode 100644 index 0000000..9eff720 --- /dev/null +++ b/tests/laborious/utils/test_dataframe_debug.py @@ -0,0 +1,12 @@ +from pandas import DataFrame + +from laborious.utils.dataframe_debug import build_dataframe_debug_message + + +def test_build_dataframe_debug_message_skips_large_dataframe(): + df = DataFrame({'a': [1, 2, 3]}) + + msg = build_dataframe_debug_message('payload', df, max_rows=1) + + assert 'skipped because dataframe has 3 rows' in msg + assert '(max: 1)' in msg diff --git a/tests/laborious/worker/test_worker.py b/tests/laborious/worker/test_worker.py new file mode 100644 index 0000000..8812a1f --- /dev/null +++ b/tests/laborious/worker/test_worker.py @@ -0,0 +1,243 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +from pytest import mark, raises + +from laborious.worker import worker + + +def _build_fake_activities(): + inst = MagicMock() + inst.init_opc = AsyncMock() + inst.shutdown = AsyncMock() + inst.load_query_with_minio_offload = MagicMock() + inst.retrain_model = MagicMock() + inst.update_production_model = MagicMock() + inst.format_retrain_report = MagicMock() + inst.export_data_to_postgres = MagicMock() + inst.load_custom_query = MagicMock() + inst.calculate_simple_metrics = MagicMock() + inst.get_reference_data = MagicMock() + inst.calculate_drift = MagicMock() + inst.request_predict = MagicMock() + inst.request_transform = MagicMock() + inst.input_gate = MagicMock() + inst.mlflow_response_gate = MagicMock() + inst.mlflow_content_gate = MagicMock() + inst.format_transformed_data = MagicMock() + inst.format_prediction = MagicMock() + inst.format_default_prediction = MagicMock() + inst.write_opc_data = MagicMock() + inst.cleanup_minio_objects_expired = MagicMock() + inst.repeat_last_prediction = MagicMock() + inst.write_metrics = MagicMock() + inst.write_pi_web_api_data = MagicMock() + return inst + + +def _build_fake_worker(async_result=None, async_error: Exception | None = None): + w = MagicMock() + + async def _run(): + if async_error is not None: + raise async_error + return async_result + + w.run = MagicMock(side_effect=_run) + return w + + +@patch('laborious.worker.worker.start_http_server') +def test_start_prometheus_server_success(mock_start_http): + with patch.object(worker.metrics.APP_UP, 'labels') as labels: + gauge = MagicMock() + labels.return_value = gauge + with patch('laborious.worker.worker.os.getenv', return_value='9090'): + worker.start_prometheus_server() + mock_start_http.assert_called_once_with(9090) + gauge.set.assert_called_once_with(1) + + +@patch('laborious.worker.worker.start_http_server', side_effect=RuntimeError('nope')) +def test_start_prometheus_server_error_exits(_mock_start_http): + with patch('laborious.worker.worker.os._exit', side_effect=SystemExit(1)) as m_exit: + with raises(SystemExit): + worker.start_prometheus_server() + m_exit.assert_called_once_with(1) + + +@mark.asyncio +async def test_main_missing_runtime_exits_fast(monkeypatch): + monkeypatch.setenv('RUNTIME', '') + with ( + patch('laborious.worker.worker.start_prometheus_server'), + patch('laborious.worker.worker.get_logger') as m_logger, + patch('laborious.worker.worker.NotificationHandler'), + patch('laborious.worker.worker.MetricsController'), + patch.object(worker.metrics.APP_UP, 'labels') as labels, + patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)), + ): + labels.return_value = MagicMock() + with raises(SystemExit): + await worker.main() + assert m_logger.return_value.custom_critical.called + + +@mark.asyncio +async def test_main_plugin_install_failure(monkeypatch): + monkeypatch.setenv('RUNTIME', 'single') + fake_activities = _build_fake_activities() + fake_plugin = MagicMock() + fake_plugin.install_runtime = AsyncMock(side_effect=RuntimeError('install failed')) + with ( + patch('laborious.worker.worker.start_prometheus_server'), + patch('laborious.worker.worker.get_logger'), + patch( + 'laborious.worker.worker.build_mongodb_config', + return_value={'connection_string': 'cs', 'database_name': 'db'}, + ), + patch('laborious.worker.worker.NotificationHandler'), + patch('laborious.worker.worker.MetricsController'), + patch( + 'laborious.worker.worker.build_plugin_store_config', + return_value={ + 'base_url': '', + 'owner': '', + 'repo': '', + 'username': None, + 'password': None, + 'branch': None, + 'cache_ttl_seconds': None, + 'pypi_index_url': '', + 'pypi_username': None, + 'pypi_password': None, + }, + ), + patch('laborious.worker.worker.PluginStore', return_value=fake_plugin), + patch('laborious.worker.worker.Activities', return_value=fake_activities), + patch.object(worker.metrics.APP_UP, 'labels') as labels, + patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)), + ): + labels.return_value = MagicMock() + with raises(SystemExit): + await worker.main() + + +@mark.asyncio +async def test_main_success_exit_zero(monkeypatch): + monkeypatch.setenv('RUNTIME', 'single') + fake_activities = _build_fake_activities() + fake_plugin = MagicMock() + fake_plugin.install_runtime = AsyncMock(return_value=None) + fake_workers = [_build_fake_worker() for _ in range(4)] + + with ( + patch('laborious.worker.worker.start_prometheus_server'), + patch('laborious.worker.worker.get_logger'), + patch( + 'laborious.worker.worker.build_mongodb_config', + return_value={'connection_string': 'cs', 'database_name': 'db'}, + ), + patch('laborious.worker.worker.NotificationHandler') as m_notif_cls, + patch('laborious.worker.worker.MetricsController'), + patch( + 'laborious.worker.worker.build_plugin_store_config', + return_value={ + 'base_url': '', + 'owner': '', + 'repo': '', + 'username': None, + 'password': None, + 'branch': None, + 'cache_ttl_seconds': None, + 'pypi_index_url': '', + 'pypi_username': None, + 'pypi_password': None, + }, + ), + patch('laborious.worker.worker.PluginStore', return_value=fake_plugin), + patch('laborious.worker.worker.Activities', return_value=fake_activities), + patch('laborious.worker.worker.build_postgres_config', return_value={}), + patch('laborious.worker.worker.build_minio_config', return_value={}), + patch('laborious.worker.worker.build_opc_config', return_value={}), + patch('laborious.worker.worker.build_api_config', return_value={}), + patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()), + patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()), + patch('laborious.worker.worker.Runtime', return_value=MagicMock()), + patch( + 'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock()) + ), + patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers) as m_prepare, + patch.object(worker.metrics.APP_UP, 'labels') as labels, + patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(0)), + ): + labels.return_value = MagicMock() + with raises(SystemExit): + await worker.main() + + notif = m_notif_cls.return_value + notif.shutdown.assert_called_once() + fake_activities.shutdown.assert_awaited_once() + assert m_prepare.call_count == 4 + prepare_calls = m_prepare.call_args_list + assert prepare_calls[0].kwargs['runtime'] == 'single' + assert prepare_calls[3].kwargs['runtime'] == 'single' + + +@mark.asyncio +async def test_main_worker_gather_error_exits_one(monkeypatch): + monkeypatch.setenv('RUNTIME', 'single') + fake_activities = _build_fake_activities() + fake_plugin = MagicMock() + fake_plugin.install_runtime = AsyncMock(return_value=None) + fake_workers = [ + _build_fake_worker(async_error=RuntimeError('boom')), + _build_fake_worker(), + _build_fake_worker(), + _build_fake_worker(), + ] + + with ( + patch('laborious.worker.worker.start_prometheus_server'), + patch('laborious.worker.worker.get_logger') as m_logger, + patch( + 'laborious.worker.worker.build_mongodb_config', + return_value={'connection_string': 'cs', 'database_name': 'db'}, + ), + patch('laborious.worker.worker.NotificationHandler'), + patch('laborious.worker.worker.MetricsController'), + patch( + 'laborious.worker.worker.build_plugin_store_config', + return_value={ + 'base_url': '', + 'owner': '', + 'repo': '', + 'username': None, + 'password': None, + 'branch': None, + 'cache_ttl_seconds': None, + 'pypi_index_url': '', + 'pypi_username': None, + 'pypi_password': None, + }, + ), + patch('laborious.worker.worker.PluginStore', return_value=fake_plugin), + patch('laborious.worker.worker.Activities', return_value=fake_activities), + patch('laborious.worker.worker.build_postgres_config', return_value={}), + patch('laborious.worker.worker.build_minio_config', return_value={}), + patch('laborious.worker.worker.build_opc_config', return_value={}), + patch('laborious.worker.worker.build_api_config', return_value={}), + patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()), + patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()), + patch('laborious.worker.worker.Runtime', return_value=MagicMock()), + patch( + 'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock()) + ), + patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers), + patch.object(worker.metrics.APP_UP, 'labels') as labels, + patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)), + ): + labels.return_value = MagicMock() + with raises(SystemExit): + await worker.main() + + assert m_logger.return_value.custom_error.called diff --git a/tests/laborious/workflows/subworkflows/test_prediction_process.py b/tests/laborious/workflows/subworkflows/test_prediction_process.py index a018504..ae761f9 100644 --- a/tests/laborious/workflows/subworkflows/test_prediction_process.py +++ b/tests/laborious/workflows/subworkflows/test_prediction_process.py @@ -1,6 +1,6 @@ from unittest.mock import ANY, AsyncMock, MagicMock, call, patch -from pytest import fixture, mark +from pytest import fixture, mark, raises from laborious.activities.activities import Activities from laborious.workflows.sub_workflows.prediction_process import PredictionProcess @@ -840,3 +840,27 @@ async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process): retry_policy=ANY, start_to_close_timeout=ANY, ) + + +@mark.asyncio +@patch('laborious.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock) +async def test_run_always_cleans_up_on_pipeline_exception(workflow_mock, prediction_process): + input_data = { + 'metadata': metadata, + 'data': {'last_timestamp': '2024-01-01'}, + 'model_id': 1, + 'model_name': 'm', + 'model_config': {}, + 'save_transform': False, + } + prediction_process._run_prediction_pipeline = AsyncMock(side_effect=RuntimeError('boom')) + + with raises(RuntimeError): + await prediction_process.run(input_data) + + workflow_mock.execute_activity_method.assert_called_once_with( + Activities.cleanup_minio_objects_expired, + {**metadata, 'data': input_data['data']}, + retry_policy=ANY, + start_to_close_timeout=ANY, + )