From 503d9aa485315e23e442f573f815cf15a3bd9d4c Mon Sep 17 00:00:00 2001 From: vitor-aignosi Date: Mon, 23 Mar 2026 14:45:21 -0300 Subject: [PATCH] SIENTIAPDE-1712 Update dependencies and refactor input filter handling for consistency - Updated sientia-dataops-library dependency version from 1.10.3 to 1.10.4 in requirements.txt. - Refactored input filter handling in the Gates class to read policy and config keys in a case-insensitive manner. - Updated test cases to ensure consistency in filter key naming conventions across various scenarios. --- e2e/conftest.py | 32 ++++--- e2e/test_predictions_batch_format_export.py | 58 +++++------- e2e/test_predictions_batch_main_workflow.py | 18 ++-- ...st_predictions_batch_prediction_process.py | 33 +++---- laborious/activities/gates.py | 44 ++++++--- laborious/activities/mlflow.py | 5 +- laborious/metrics.py | 4 +- .../utils/models/minio_dataframe_payload.py | 3 +- laborious/worker/worker.py | 5 +- laborious/workflows/drift.py | 4 +- laborious/workflows/minimal_retrain.py | 4 +- laborious/workflows/predictions_batch.py | 8 +- laborious/workflows/simple_metrics.py | 2 +- .../sub_workflows/prediction_process.py | 2 - requirements.txt | 2 +- tests/laborious/activities/test_gates.py | 73 +++++++++++++-- tests/laborious/activities/test_mlflow.py | 2 + tests/laborious/activities/test_storage.py | 12 +-- .../models/test_minio_dataframe_payload.py | 14 ++- .../subworkflows/test_prediction_process.py | 47 +++++----- tests/laborious/workflows/test_drift.py | 26 +++--- .../workflows/test_minimal_retrain.py | 35 +++++-- .../workflows/test_predictions_batch.py | 40 +++++--- .../workflows/test_simple_metrics.py | 92 +++++++++---------- 24 files changed, 331 insertions(+), 234 deletions(-) diff --git a/e2e/conftest.py b/e2e/conftest.py index 61f86f2..5f3aa96 100644 --- a/e2e/conftest.py +++ b/e2e/conftest.py @@ -247,15 +247,17 @@ def mock_minio_repository(): def mock_pi_web_api_repository(): """Mock PI Web API repository for PI Web API operations.""" mock_repo = MagicMock() - mock_repo.write_value = AsyncMock( - return_value={ - 'Items': [ - { - 'WebId': 'web_id_1' - } - ] - } - ) + + async def _write_value(web_ids, value, metadata=None, **kwargs): + """ + Mirror successful PI writes: one response item per requested web_id. + + write_pi_web_api_data passes the list into process_pi_web_api_response (not a + wrapped {'Items': ...} envelope). + """ + return [{'WebId': wid, 'Errors': []} for wid in web_ids] + + mock_repo.write_value = AsyncMock(side_effect=_write_value) mock_repo.close = MagicMock() return mock_repo @@ -266,7 +268,7 @@ def mock_opc_repository(): mock_repo.write_data = AsyncMock( return_value=(True, {'response_time': 0.1}) ) - mock_repo.disconnect = MagicMock() + mock_repo.disconnect = AsyncMock() return mock_repo @pytest_asyncio.fixture @@ -279,7 +281,8 @@ def patch_create_engine(postgres_engine): @pytest_asyncio.fixture def patch_minio_repository(mock_minio_repository): """Patch MinioRepository to return mock.""" - with patch('sientia_do.repository.minio_repository.MinioRepository', return_value=mock_minio_repository): + # Patch where Activities resolves the symbol (import binds the original class). + with patch('laborious.activities.activities.MinioRepository', return_value=mock_minio_repository): yield @pytest_asyncio.fixture @@ -407,10 +410,13 @@ async def test_activities( 'password': 'test', }, minio_config={ - 'endpoint_url': 'http://localhost:9000', + # Host:port only; Minio() prepends http(s):// from the secure flag. + 'endpoint_url': 'localhost:9000', 'access_key': 'test', 'secret_key': 'test', 'default_bucket': 'test-bucket', + 'retention_hours': 24, + 'secure': False, }, opc_config={}, pi_web_api_config={ @@ -452,7 +458,6 @@ async def temporal_worker(temporal_test_env, test_activities): test_activities.load_custom_query, test_activities.load_query_with_minio_offload, test_activities.cleanup_minio_objects_expired, - test_activities.get_last_timestamp, test_activities.input_gate, test_activities.request_transform, test_activities.mlflow_response_gate, @@ -465,6 +470,7 @@ async def temporal_worker(temporal_test_env, test_activities): test_activities.write_pi_web_api_data, test_activities.write_opc_data, test_activities.export_data_to_postgres, + test_activities.export_payload_to_postgres, test_activities.write_metrics, ], ) as worker: diff --git a/e2e/test_predictions_batch_format_export.py b/e2e/test_predictions_batch_format_export.py index e7bdd91..d7e7621 100644 --- a/e2e/test_predictions_batch_format_export.py +++ b/e2e/test_predictions_batch_format_export.py @@ -4,6 +4,7 @@ End-to-end tests for PredictionsBatch workflow - Format and Export scenarios. import asyncio from datetime import datetime +from typing import Any, cast from unittest.mock import ANY, AsyncMock, patch, call import pandas as pd @@ -25,13 +26,13 @@ base_input_data = { 'table_name': 'predictions', 'transform_table_name': 'transformed_data', 'input_filters': { - 'EMPTY_DATA': {'policy': 'STOP', 'config': {}}, + 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_transform_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_predict_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'opc_output_config': {}, @@ -191,7 +192,6 @@ async def test_scenario_3_1_1_default_prediction_export( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0.5, }, - endpoint='test_endpoint', metadata={ 'model_id': 311, 'model_name': 'test_model', @@ -205,7 +205,6 @@ async def test_scenario_3_1_1_default_prediction_export( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0, }, - endpoint='test_endpoint', metadata={ 'model_id': 311, 'model_name': 'test_model', @@ -217,7 +216,8 @@ async def test_scenario_3_1_1_default_prediction_export( any_order=True, ) - test_activities.opc_repository['1'].write_data.assert_has_calls( + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.assert_has_calls( [ call('addr_1', 0.5, 'float', ANY, { @@ -300,7 +300,8 @@ async def test_scenario_3_1_2_export_with_opc_only( await start_and_await_workflow(client, input_data, workflow_id) - test_activities.opc_repository['1'].write_data.assert_has_calls( + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.assert_has_calls( [ call('addr_1', 0.5, 'float', ANY, { @@ -382,7 +383,6 @@ async def test_scenario_3_1_3_export_with_pi_web_api_only( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0.5, }, - endpoint='test_endpoint', metadata={ 'model_id': 313, 'model_name': 'test_model', @@ -396,7 +396,6 @@ async def test_scenario_3_1_3_export_with_pi_web_api_only( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0, }, - endpoint='test_endpoint', metadata={ 'model_id': 313, 'model_name': 'test_model', @@ -408,7 +407,8 @@ async def test_scenario_3_1_3_export_with_pi_web_api_only( any_order=True, ) - test_activities.opc_repository['1'].write_data.assert_not_called() + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.assert_not_called() assert_prediction(postgres_engine, model_id) @@ -459,7 +459,8 @@ async def test_scenario_3_1_4_export_without_optional_outputs( await start_and_await_workflow(client, input_data, workflow_id) test_activities.pi_web_api_client.write_value.assert_not_called() - test_activities.opc_repository['1'].write_data.assert_not_called() + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.assert_not_called() assert_prediction(postgres_engine, model_id) @@ -535,7 +536,6 @@ async def test_scenario_3_1_5_export_without_transformed_data( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0.5, }, - endpoint='test_endpoint', metadata={ 'model_id': 315, 'model_name': 'test_model', @@ -549,7 +549,6 @@ async def test_scenario_3_1_5_export_without_transformed_data( 'Timestamp': '2024-01-01 12:00:00+0000', 'Value': 0, }, - endpoint='test_endpoint', metadata={ 'model_id': 315, 'model_name': 'test_model', @@ -561,7 +560,8 @@ async def test_scenario_3_1_5_export_without_transformed_data( any_order=True, ) - test_activities.opc_repository['1'].write_data.assert_has_calls( + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.assert_has_calls( [ call('addr_1', 0.5, 'float', ANY, { @@ -695,7 +695,8 @@ async def test_scenario_3_2_2_opc_write_error( insert_sample_data(postgres_engine, model_id, [23.5, 78.2]) print("[TEST] ✓ Data inserted successfully") - test_activities.opc_repository['1'].write_data.return_value = (False, { + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + opc_write_data.return_value = (False, { 'notification_id': 'OPC_WRITE_DATA_ERROR_1', 'message': 'OPC server unavailable', 'block': 'opc_repository', @@ -774,25 +775,14 @@ async def test_scenario_3_2_3_pi_web_api_partial_write_error( insert_sample_data(postgres_engine, model_id, [23.5, 78.2]) print("[TEST] ✓ Data inserted successfully") - test_activities.pi_web_api_client.write_value = AsyncMock(side_effect=[ - { - 'Items': [ - { - 'WebId': 'web_id_1', - 'Errors': [], - }, - ] - }, - Exception('Tag write failed'), - { - 'Items': [ - { - 'WebId': 'web_id_2', - 'Errors': [], - }, - ] - }, - ]) + test_activities.pi_web_api_client.write_value = AsyncMock( + side_effect=[ + # Prediction batch: two web_ids requested, only one acknowledged. + [{'WebId': 'web_id_1', 'Errors': []}], + # Confidence write succeeds. + [{'WebId': 'web_id_2', 'Errors': []}], + ] + ) input_data = get_base_input_data(model_id) input_data['pi_web_api_output_config'] = { diff --git a/e2e/test_predictions_batch_main_workflow.py b/e2e/test_predictions_batch_main_workflow.py index f8469ef..ddac7ec 100644 --- a/e2e/test_predictions_batch_main_workflow.py +++ b/e2e/test_predictions_batch_main_workflow.py @@ -100,13 +100,13 @@ async def test_scenario_1_1_1_happy_path_complete_success( 'table_name': 'predictions', 'transform_table_name': 'transformed_data', 'input_filters': { - 'EMPTY_DATA': {'policy': 'STOP', 'config': {}}, + 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_transform_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_predict_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'opc_output_config': {}, @@ -238,13 +238,13 @@ async def test_scenario_1_2_1_sql_query_execution_error( 'table_name': 'predictions', 'transform_table_name': 'transformed_data', 'input_filters': { - 'EMPTY_DATA': {'policy': 'STOP', 'config': {}}, + 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_transform_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_predict_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'opc_output_config': {}, @@ -407,13 +407,13 @@ async def test_scenario_1_2_3_invalid_datetime_column_specification( 'table_name': 'predictions', 'transform_table_name': 'transformed_data', 'input_filters': { - 'EMPTY_DATA': {'policy': 'STOP', 'config': {}}, + 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_transform_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_predict_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'opc_output_config': {}, diff --git a/e2e/test_predictions_batch_prediction_process.py b/e2e/test_predictions_batch_prediction_process.py index d32980e..98f9429 100644 --- a/e2e/test_predictions_batch_prediction_process.py +++ b/e2e/test_predictions_batch_prediction_process.py @@ -5,6 +5,7 @@ End-to-end tests for PredictionsBatch workflow - Prediction Process scenarios. import asyncio from datetime import datetime from decimal import Decimal +from typing import Any from unittest.mock import MagicMock, patch import pandas as pd @@ -27,15 +28,15 @@ base_input_data = { 'transform_table_name': 'transformed_data', 'input_filters': { 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'policy': 'CONTINUE', # Continue despite issues, not STOP - 'config': {'variables': ['sensor_1']}, + 'POLICY': 'CONTINUE', # Continue despite issues, not STOP + 'CONFIG': {'variables': ['sensor_1']}, }, }, 'mlflow_transform_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'mlflow_predict_filters': { - 'API_ERROR': {'policy': 'STOP', 'config': {}}, + 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, }, 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'opc_output_config': {}, @@ -59,7 +60,7 @@ def get_base_input_data(model_id): 'query': base_query.format(model_id=model_id), } -def insert_sample_data(postgres_engine, model_id, values: list[tuple]): +def insert_sample_data(postgres_engine, model_id, values: list[Any]): with postgres_engine.begin() as conn: conn.execute(text(f"DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}")) @@ -111,7 +112,7 @@ async def start_and_await_workflow(client, input_data, workflow_id): pytest.fail("Workflow execution timed out after 60 seconds") def assert_continue( - postgres_engine, model_id, prediction_confidence: Decimal = 2, + postgres_engine, model_id, prediction_confidence: Decimal = Decimal(2), comments: str = 'Input data with bad quality', ): print("\n[TEST] 4. Verifying prediction was created despite warnings...") @@ -236,7 +237,7 @@ async def test_scenario_2_1_2_input_gate_triggers_stop( print("[TEST] ✓ Data inserted successfully") input_data = get_base_input_data(model_id) - input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['policy'] = 'STOP' + input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'STOP' print("\n[TEST] 2. Starting workflow that should stop at input gate...") workflow_id = f'test-input-stop-{datetime.now().timestamp()}' @@ -285,7 +286,7 @@ async def test_scenario_2_1_3_input_gate_triggers_repeat( print("[TEST] ✓ Data and previous prediction inserted") input_data = get_base_input_data(model_id) - input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['policy'] = 'REPEAT' + input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'REPEAT' print("\n[TEST] 2. Starting workflow that should trigger REPEAT...") workflow_id = f'test-input-repeat-{datetime.now().timestamp()}' @@ -345,7 +346,7 @@ async def test_scenario_2_2_1_transform_gate_triggers_continue( insert_sample_data(postgres_engine, model_id, [60.0, 78.2]) input_data = get_base_input_data(model_id) - input_data['mlflow_transform_filters']['API_ERROR']['policy'] = 'CONTINUE' + input_data['mlflow_transform_filters']['API_ERROR']['POLICY'] = 'CONTINUE' print("\n[TEST] 2. Starting workflow that should trigger CONTINUE at transform gate...") workflow_id = f'test-transform-continue-{datetime.now().timestamp()}' @@ -355,7 +356,7 @@ async def test_scenario_2_2_1_transform_gate_triggers_continue( postgres_engine=postgres_engine, model_id=model_id, prediction_confidence=Decimal(10), - comments='Bad data model', + comments='Unknown MLFlow API error', ) print("\n[TEST] ✓ All assertions passed!") @@ -397,7 +398,7 @@ async def test_scenario_2_2_2_transform_gate_triggers_stop( print("[TEST] ✓ Data inserted successfully") input_data = get_base_input_data(model_id) - input_data['mlflow_transform_filters']['API_ERROR']['policy'] = 'STOP' + input_data['mlflow_transform_filters']['API_ERROR']['POLICY'] = 'STOP' print("\n[TEST] 2. Starting workflow that should trigger STOP at transform gate...") workflow_id = f'test-transform-stop-{datetime.now().timestamp()}' @@ -449,7 +450,7 @@ async def test_scenario_2_2_3_transform_gate_triggers_repeat( print("[TEST] ✓ Data inserted successfully") input_data = get_base_input_data(model_id) - input_data['mlflow_transform_filters']['API_ERROR']['policy'] = 'REPEAT' + input_data['mlflow_transform_filters']['API_ERROR']['POLICY'] = 'REPEAT' print("\n[TEST] 2. Starting workflow that should trigger REPEAT at transform gate...") workflow_id = f'test-transform-repeat-{datetime.now().timestamp()}' @@ -516,7 +517,7 @@ async def test_scenario_2_3_1_predict_gate_triggers_continue( insert_sample_data(postgres_engine, model_id, [60.0, 78.2]) input_data = get_base_input_data(model_id) - input_data['mlflow_predict_filters']['API_ERROR']['policy'] = 'CONTINUE' + input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'CONTINUE' print("\n[TEST] 2. Starting workflow that should trigger CONTINUE at predict gate...") workflow_id = f'test-predict-continue-{datetime.now().timestamp()}' @@ -526,7 +527,7 @@ async def test_scenario_2_3_1_predict_gate_triggers_continue( postgres_engine=postgres_engine, model_id=model_id, prediction_confidence=Decimal(10), - comments='Bad predict model', + comments='Unknown MLFlow API error', ) print("\n[TEST] ✓ All assertions passed!") @@ -567,7 +568,7 @@ async def test_scenario_2_3_2_predict_gate_triggers_stop( print("[TEST] ✓ Data inserted successfully") input_data = get_base_input_data(model_id) - input_data['mlflow_predict_filters']['API_ERROR']['policy'] = 'STOP' + input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'STOP' print("\n[TEST] 2. Starting workflow that should stop at predict gate...") workflow_id = f'test-predict-stop-{datetime.now().timestamp()}' @@ -617,7 +618,7 @@ async def test_scenario_2_3_3_predict_gate_triggers_repeat( print("[TEST] ✓ Data and previous prediction inserted") input_data = get_base_input_data(model_id) - input_data['mlflow_predict_filters']['API_ERROR']['policy'] = 'REPEAT' + input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'REPEAT' input_data['path_priority'] = ['REPEAT', 'STOP', 'CONTINUE'] # REPEAT first print("\n[TEST] 2. Starting workflow that should trigger REPEAT at predict gate...") diff --git a/laborious/activities/gates.py b/laborious/activities/gates.py index cac6c38..1d561c5 100644 --- a/laborious/activities/gates.py +++ b/laborious/activities/gates.py @@ -118,6 +118,22 @@ class Gates(MinioManager): def __del__(self): self.close() + @staticmethod + def _read_filter_entry(config: dict[str, Any]) -> tuple[str, dict[str, Any]]: + """ + Read filter policy/config keys in a case-insensitive way. + + Args: + config (dict[str, Any]): Filter configuration dictionary. + + Return: + tuple[str, dict[str, Any]]: Parsed policy and config payload. + """ + normalized = {str(key).upper(): value for key, value in config.items()} + policy = normalized['POLICY'] + filter_config = normalized.get('CONFIG', {}) + return policy, filter_config + @activity.defn(name='input_gate') async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ @@ -171,10 +187,11 @@ class Gates(MinioManager): if fil not in input_filter_functions: self.error(f'Filter {fil} not found', metadata) continue + policy, filter_config = self._read_filter_entry(config) try: - if input_filter_functions[fil](data, config['config']): + if input_filter_functions[fil](data, filter_config): self.debug(f'Data not passed the input filter {fil}:{config}', metadata) - filter_output.append(config['policy']) + filter_output.append(policy) except Exception as e: trace = traceback.format_exc() await self.send_notification_async( @@ -235,8 +252,10 @@ class Gates(MinioManager): self.info('Performing mlflow response gate...', metadata) raw_data = input_data['data'] filters = input_data['filters'] - - self.debug(f'Input data: \n {create_sample_dict(raw_data, max_items=5, max_depth=5)}', metadata) + + self.debug( + f'Input data: \n {create_sample_dict(raw_data, max_items=5, max_depth=5)}', metadata + ) self.debug(f'Filters: {filters}', metadata) payload = MinioDataFramePayload.from_dict(raw_data) @@ -254,17 +273,18 @@ class Gates(MinioManager): for fil, config in filters.items(): if fil not in mlflow_response_filter_functions: continue + policy, filter_config = self._read_filter_entry(config) try: - if mlflow_response_filter_functions[fil](status, config): - filter_output.append(config['policy']) - comments.append(status['message']) + if mlflow_response_filter_functions[fil](status, filter_config): + filter_output.append(policy) + comments.append(status.get('message', 'Unknown MLFlow API error')) await self.send_notification_async( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}', - message=data['content']['message'], + message=status.get('message', 'Unknown MLFlow API error'), block='mlflow_gate', level=NotificationLevel.ERROR, - attachment_content=data['content']['traceback'], + attachment_content=status.get('traceback'), ) except Exception as e: trace = traceback.format_exc() @@ -341,9 +361,10 @@ class Gates(MinioManager): for fil, config in filters.items(): if fil not in mlflow_content_filter_functions: continue + policy, filter_config = self._read_filter_entry(config) try: - if mlflow_content_filter_functions[fil](data, config): - filter_output.append(config['policy']) + if mlflow_content_filter_functions[fil](data, filter_config): + filter_output.append(policy) await self.send_notification_async( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}', @@ -703,7 +724,6 @@ class Gates(MinioManager): self.info(f'Writing metrics for model {metadata["model_name"]}', metadata) - core_tags = { 'pod_id': self.pod_id, 'runtime': self.runtime, diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index ea97f3b..ad9b6ae 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -1,7 +1,5 @@ from temporalio import activity, workflow - - with workflow.unsafe.imports_passed_through(): import traceback from typing import Any @@ -20,8 +18,9 @@ with workflow.unsafe.imports_passed_through(): now, ) from sientia_do.utils.formatters import create_sample_dict - from laborious.utils.repository.minio_manager import MinioManager + from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload + from laborious.utils.repository.minio_manager import MinioManager from laborious.utils.repository.model_repository import MLFlowRepository diff --git a/laborious/metrics.py b/laborious/metrics.py index a67e574..776b826 100644 --- a/laborious/metrics.py +++ b/laborious/metrics.py @@ -25,9 +25,7 @@ Metric Labels: """ from prometheus_client import Counter, Gauge, Histogram -from sientia_do.observability.metrics import ( - CORE_LABELS -) +from sientia_do.observability.metrics import CORE_LABELS # Application health metric APP_UP = Gauge( diff --git a/laborious/utils/models/minio_dataframe_payload.py b/laborious/utils/models/minio_dataframe_payload.py index 42b60f4..bb3f2fd 100644 --- a/laborious/utils/models/minio_dataframe_payload.py +++ b/laborious/utils/models/minio_dataframe_payload.py @@ -162,7 +162,7 @@ class MinioDataFramePayload: """ Return True if the payload has some data internally or in MinIO. """ - return (self.data is not None and not self.data != {}) or self.object_key is not None + return (self.data is not None and self.data != {}) or self.object_key is not None @classmethod async def from_dataframe( @@ -200,7 +200,6 @@ class MinioDataFramePayload: data=None, last_timestamp=now().strftime(DATETIME_FORMAT_WITH_TZ), status=status ) - if last_timestamp is None: last_timestamp = max(dataframe['timestamp'].values.tolist()) diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index 5202f2d..fa09259 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -219,20 +219,21 @@ async def main(): logger.custom_info('Workers started successfully', metadata) + exit_code = 0 try: # This will run the workers and wait for them to complete. # If an exception occurs in any of the worker handlers, it will be propagated here. await asyncio.gather(*handlers) except BaseException as e: # NOSONAR 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() - # Exit with a non-zero status code to indicate failure to Kubernetes metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN - sys.exit(1) + sys.exit(exit_code) def start_prometheus_server(): diff --git a/laborious/workflows/drift.py b/laborious/workflows/drift.py index b98fc56..2faff32 100644 --- a/laborious/workflows/drift.py +++ b/laborious/workflows/drift.py @@ -46,7 +46,7 @@ class Drift: ORDER BY timestamp ASC """ # nosec B608 - values come from internal Temporal workflow config, not user input - target_data_handler = workflow.start_local_activity_method( + target_data_handler = workflow.start_activity_method( Activities.load_custom_query, { **metadata, @@ -58,7 +58,7 @@ class Drift: start_to_close_timeout=timedelta(seconds=300), ) - reference_data_handler = workflow.start_local_activity_method( + reference_data_handler = workflow.start_activity_method( Activities.get_reference_data, {**metadata, 'model_name': input_data['model_name']}, retry_policy=retry_policy, diff --git a/laborious/workflows/minimal_retrain.py b/laborious/workflows/minimal_retrain.py index f0cf0e8..f8f5151 100644 --- a/laborious/workflows/minimal_retrain.py +++ b/laborious/workflows/minimal_retrain.py @@ -7,6 +7,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.temporal.policies import retry_policy from laborious.activities.activities import Activities + from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload @workflow.defn(name='minimal_retrain') @@ -84,7 +85,8 @@ class MinimalRetrain: start_to_close_timeout=timedelta(seconds=600), ) - if not storage_result.has_data(): + storage_payload = MinioDataFramePayload.from_dict(storage_result) + if not storage_payload.has_data(): raise ValueError('No data returned from query') experiment_response = await workflow.execute_activity_method( diff --git a/laborious/workflows/predictions_batch.py b/laborious/workflows/predictions_batch.py index 427a470..36e2e97 100644 --- a/laborious/workflows/predictions_batch.py +++ b/laborious/workflows/predictions_batch.py @@ -105,12 +105,14 @@ class PredictionsBatch: 'transform_table_name': input_data['transform_table_name'], 'model_id': input_data['model_id'], 'model_name': input_data['model_name'], - 'input_filters': input_data.get('input_filters', {'EMPTY_DATA': {'POLICY': 'STOP'}}), + 'input_filters': input_data.get( + 'input_filters', {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}} + ), 'mlflow_transform_filters': input_data.get( - 'mlflow_transform_filters', {'API_ERROR': {'POLICY': 'STOP'}} + 'mlflow_transform_filters', {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}} ), 'mlflow_predict_filters': input_data.get( - 'mlflow_predict_filters', {'API_ERROR': {'POLICY': 'STOP'}} + 'mlflow_predict_filters', {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}} ), 'model_config': input_data.get('model_config', {}), 'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']), diff --git a/laborious/workflows/simple_metrics.py b/laborious/workflows/simple_metrics.py index 20b60ab..38b47c4 100644 --- a/laborious/workflows/simple_metrics.py +++ b/laborious/workflows/simple_metrics.py @@ -47,7 +47,7 @@ class SimpleMetrics: p."timestamp" desc; """ # nosec B608 - values come from internal Temporal workflow config, not user input - target_data = await workflow.execute_local_activity_method( + target_data = await workflow.execute_activity_method( Activities.load_custom_query, { **metadata, diff --git a/laborious/workflows/sub_workflows/prediction_process.py b/laborious/workflows/sub_workflows/prediction_process.py index fbac978..3493b59 100644 --- a/laborious/workflows/sub_workflows/prediction_process.py +++ b/laborious/workflows/sub_workflows/prediction_process.py @@ -7,7 +7,6 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.temporal.policies import retry_policy from laborious.activities.activities import Activities - from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload @workflow.defn(name='subworkflow.prediction_process') @@ -112,7 +111,6 @@ class PredictionProcess: start_to_close_timeout=timedelta(minutes=5), ) raise e - async def _run_prediction_pipeline( self, diff --git a/requirements.txt b/requirements.txt index dd4fa1a..f0eedf8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,7 +3,7 @@ psycopg2-binary sqlalchemy asyncua redis -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.10.3 +git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.10.4 git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.40.7 prometheus-client botocore diff --git a/tests/laborious/activities/test_gates.py b/tests/laborious/activities/test_gates.py index 367f5b4..3d072b9 100644 --- a/tests/laborious/activities/test_gates.py +++ b/tests/laborious/activities/test_gates.py @@ -90,7 +90,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac ) input_data = { **metadata, - 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, + 'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}}, 'data': _minio_payload(DataFrame({'value': []})), 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], } @@ -103,7 +103,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac gates_activity.send_notification_async.assert_called_once_with( metadata=metadata['metadata'], notification_id='INTPUT_GATE_ERROR__EMPTY_DATA', - message="Error in filter EMPTY_DATA:{'policy': 'STOP', 'config': {}}: \n Test error", + message="Error in filter EMPTY_DATA:{'POLICY': 'STOP', 'CONFIG': {}}: \n Test error", block='input_gate', level=NotificationLevel.ERROR, attachment_content=ANY, @@ -133,7 +133,7 @@ async def test_input_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, - 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, + 'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}}, 'data': _minio_payload(DataFrame({'value': []})), 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], } @@ -147,11 +147,45 @@ async def test_input_gate_with_filter(gates_activity): @mark.asyncio -async def test_input_gate_with_filter_not_caught(gates_activity): +async def test_input_gate_with_filter_lowercase_keys(gates_activity): # Arrange input_data = { **metadata, 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, + 'data': _minio_payload(DataFrame({'value': []})), + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], + } + + # Act + result = await gates_activity.input_gate(input_data) + + # Assert + assert result == ('STOP', -1, 'Input data with bad quality') + + +@mark.asyncio +async def test_input_gate_with_filter_capitalized_keys(gates_activity): + # Arrange + input_data = { + **metadata, + 'filters': {'EMPTY_DATA': {'Policy': 'STOP', 'Config': {}}}, + 'data': _minio_payload(DataFrame({'value': []})), + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], + } + + # Act + result = await gates_activity.input_gate(input_data) + + # Assert + assert result == ('STOP', -1, 'Input data with bad quality') + + +@mark.asyncio +async def test_input_gate_with_filter_not_caught(gates_activity): + # Arrange + input_data = { + **metadata, + 'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})), 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], } @@ -248,7 +282,7 @@ async def test_mlflow_response_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, - 'filters': {'API_ERROR': {'policy': 'STOP'}}, + 'filters': {'API_ERROR': {'POLICY': 'STOP'}}, 'data': _minio_payload( {'content': {'message': 'API error occurred', 'traceback': 'error trace'}}, status={'success': False, 'message': 'API error occurred'}, @@ -266,12 +300,33 @@ async def test_mlflow_response_gate_with_filter(gates_activity): gates_activity.send_notification_async.assert_called() +@mark.asyncio +async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity): + # Arrange + input_data = { + **metadata, + 'filters': {'API_ERROR': {'Policy': 'STOP'}}, + 'data': _minio_payload( + {'content': {'message': 'API error occurred', 'traceback': 'error trace'}}, + status={'success': False, 'message': 'API error occurred'}, + ), + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], + } + + # Act + result = await gates_activity.mlflow_response_gate(input_data) + + # Assert + assert result == ('STOP', -1, 'API error occurred') + + @mark.asyncio async def test_mlflow_response_gate_with_filter_not_caught(gates_activity): # Arrange input_data = { **metadata, - 'filters': {'API_ERROR': {'policy': 'STOP'}}, + 'filters': {'API_ERROR': {'POLICY': 'STOP'}}, 'data': _minio_payload( {'content': {'message': 'success'}}, status={'success': True}, @@ -364,7 +419,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, - 'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}}, + 'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}}, 'data': _minio_payload(DataFrame({'value': [None, None, None]})), 'type': 'test', 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], @@ -402,7 +457,7 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity): async def test_mlflow_content_gate_filter_returns_false(gates_activity): input_data = { **metadata, - 'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}}, + 'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})), 'type': 'test', 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], @@ -823,6 +878,7 @@ async def test_write_metrics(mock_metrics, gates_activity): core_tags = { 'pod_id': gates_activity.pod_id, 'runtime': gates_activity.runtime, + 'operation_type': 'predict', 'model_name': metadata['metadata']['model_name'], 'workflow_name': metadata['metadata']['workflow_name'], } @@ -930,6 +986,7 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act tags={ 'pod_id': gates_activity.pod_id, 'runtime': gates_activity.runtime, + 'operation_type': 'predict', 'model_name': metadata['metadata']['model_name'], 'workflow_name': metadata['metadata']['workflow_name'], 'opc_server_id': 'server1', diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index e495be8..daaebab 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -180,6 +180,7 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow): operation='transform', status=transform_response, workflow_metadata=metadata['metadata'], + last_timestamp=payload.last_timestamp, ) assert response_data == mock_from_dataframe.return_value @@ -250,6 +251,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml operation='predict', status=predict_response, workflow_metadata=metadata['metadata'], + last_timestamp=payload.last_timestamp, ) assert response_data == mock_from_dataframe.return_value diff --git a/tests/laborious/activities/test_storage.py b/tests/laborious/activities/test_storage.py index 9e76c5f..ce99ff7 100644 --- a/tests/laborious/activities/test_storage.py +++ b/tests/laborious/activities/test_storage.py @@ -204,9 +204,7 @@ async def test_cleanup_minio_objects_expired(mock_now, storage): data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired( - {**metadata, 'data': data_mock} - ) + result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 1 assert result['failed_count'] == 0 @@ -291,9 +289,7 @@ async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage): data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired( - {**metadata, 'data': data_mock} - ) + result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 0 assert result['failed_count'] == 1 @@ -312,9 +308,7 @@ async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storag data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired( - {**metadata, 'data': data_mock} - ) + result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 0 assert result['failed_count'] == 0 diff --git a/tests/laborious/utils/models/test_minio_dataframe_payload.py b/tests/laborious/utils/models/test_minio_dataframe_payload.py index de1a334..c27765d 100644 --- a/tests/laborious/utils/models/test_minio_dataframe_payload.py +++ b/tests/laborious/utils/models/test_minio_dataframe_payload.py @@ -90,22 +90,20 @@ async def test_retrieve_downloads_parquet_when_offloaded(): def test_build_object_key(): key, prefix = _build_object_key('my-model', 'initial', '2024-01-01_00-00-00') - assert key == 'training_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet' - assert prefix == 'training_datasets/my-model' + assert key == 'prediction_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet' + assert prefix == 'prediction_datasets/my-model' def test_build_object_key_strips_slashes(): key, prefix = _build_object_key(' /my-model/ ', 'transform', '2024-06-15_10-30-45') - assert prefix == 'training_datasets/my-model' - assert key.startswith('training_datasets/my-model/') + assert prefix == 'prediction_datasets/my-model' + assert key.startswith('prediction_datasets/my-model/') def test_estimate_size_bytes_fallback(): df = DataFrame({'a': [1, 2]}) - original_to_dict = df.to_dict - df.to_dict = lambda *a, **kw: (_ for _ in ()).throw(RuntimeError('to_dict failed')) - size = MinioDataFramePayload.estimate_size_bytes(df) - df.to_dict = original_to_dict + with patch.object(df, 'to_dict', side_effect=RuntimeError('to_dict failed')): + size = MinioDataFramePayload.estimate_size_bytes(df) assert isinstance(size, int) assert size > 0 diff --git a/tests/laborious/workflows/subworkflows/test_prediction_process.py b/tests/laborious/workflows/subworkflows/test_prediction_process.py index db75b8c..60f3744 100644 --- a/tests/laborious/workflows/subworkflows/test_prediction_process.py +++ b/tests/laborious/workflows/subworkflows/test_prediction_process.py @@ -6,15 +6,6 @@ from laborious.activities.activities import Activities from laborious.workflows.sub_workflows.prediction_process import PredictionProcess -@fixture(autouse=True) -def _passthrough_from_dict(): - with patch( - 'laborious.workflows.sub_workflows.prediction_process.MinioDataFramePayload.from_dict', - side_effect=lambda x: x, - ): - yield - - @fixture def prediction_process(): return PredictionProcess() @@ -38,7 +29,9 @@ async def test_run(workflow_mock, prediction_process): # Arrange data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, @@ -128,7 +121,7 @@ async def test_run(workflow_mock, prediction_process): { **metadata, 'filters': input_data['mlflow_transform_filters'], - 'data': 'transformed_data', + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'type': 'transform', 'path_priority': input_data['path_priority'], }, @@ -143,7 +136,7 @@ async def test_run(workflow_mock, prediction_process): Activities.request_predict, { **metadata, - 'data': 'transformed_data', + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'model_name': input_data['model_name'], 'model_config': input_data['model_config'], }, @@ -174,8 +167,8 @@ async def test_run(workflow_mock, prediction_process): { 'metadata': metadata, 'path_flag': 'continue', - 'data': 'predicted_data', - 'transformed_data': 'transformed_data', + 'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'}, + 'transformed_data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'prediction_confidence': 0.95, 'timestamp': '2024-01-01', 'model_id': 1, @@ -199,7 +192,9 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process): # Arrange data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, @@ -251,7 +246,9 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_ # Arrange data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, @@ -336,7 +333,9 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process # Arrange data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, @@ -421,7 +420,7 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process Activities.mlflow_content_gate, { 'filters': input_data['mlflow_transform_filters'], - 'data': 'transformed_data', + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'type': 'transform', 'path_priority': input_data['path_priority'], **metadata, @@ -441,7 +440,9 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p # Arrange data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, @@ -527,7 +528,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p Activities.mlflow_content_gate, { 'filters': input_data['mlflow_transform_filters'], - 'data': 'transformed_data', + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'type': 'transform', 'path_priority': input_data['path_priority'], **metadata, @@ -542,7 +543,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p call( Activities.request_predict, { - 'data': 'transformed_data', + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'model_name': input_data['model_name'], 'model_config': input_data['model_config'], **metadata, @@ -784,7 +785,9 @@ async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process): data_payload = MagicMock() data_payload.cleanup_prefix.return_value = 'training_datasets/test' - data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + data_payload.__getitem__ = ( + lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock() + ) input_data = { 'metadata': metadata, 'data': data_payload, diff --git a/tests/laborious/workflows/test_drift.py b/tests/laborious/workflows/test_drift.py index 8f013d3..c31dafc 100644 --- a/tests/laborious/workflows/test_drift.py +++ b/tests/laborious/workflows/test_drift.py @@ -45,10 +45,9 @@ async def test_run(workflow_mock: AsyncMock, drift: Drift): reference_data = {'data': 'test_reference_data'} drift_data = {'drift': 'test_drift_data'} - workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data] - - workflow_mock.execute_local_activity_method.return_value = drift_data - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.start_activity_method.side_effect = [target_data, reference_data] + workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data) + workflow_mock.execute_activity_method = AsyncMock(return_value=None) # Act await drift.run(input_data) @@ -64,7 +63,7 @@ async def test_run(workflow_mock: AsyncMock, drift: Drift): ORDER BY timestamp ASC """ - workflow_mock.start_local_activity_method.assert_has_calls( + workflow_mock.start_activity_method.assert_has_calls( [ call( Activities.load_custom_query, @@ -140,16 +139,15 @@ async def test_run_empty_target_data(workflow_mock: AsyncMock, drift: Drift): target_data = None reference_data = {'data': 'test_reference_data'} - workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data] + workflow_mock.start_activity_method.side_effect = [target_data, reference_data] - workflow_mock.execute_local_activity_method = AsyncMock() + workflow_mock.execute_activity_method = AsyncMock() workflow_mock.execute_activity_method = AsyncMock() # Act await drift.run(input_data) # Assert - Should not call calculate_drift or export - workflow_mock.execute_local_activity_method.assert_not_called() workflow_mock.execute_activity_method.assert_not_called() @@ -175,9 +173,8 @@ async def test_run_empty_drift_data(workflow_mock: AsyncMock, drift: Drift): reference_data = {'data': 'test_reference_data'} drift_data = None - workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data] - - workflow_mock.execute_local_activity_method.return_value = drift_data + workflow_mock.start_activity_method.side_effect = [target_data, reference_data] + workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data) workflow_mock.execute_activity_method = AsyncMock() # Act @@ -226,10 +223,9 @@ async def test_run_default_chunk_period(workflow_mock: AsyncMock, drift: Drift): reference_data = {'data': 'test_reference_data'} drift_data = {'drift': 'test_drift_data'} - workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data] - - workflow_mock.execute_local_activity_method.return_value = drift_data - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.start_activity_method.side_effect = [target_data, reference_data] + workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data) + workflow_mock.execute_activity_method = AsyncMock(return_value=None) # Act await drift.run(input_data) diff --git a/tests/laborious/workflows/test_minimal_retrain.py b/tests/laborious/workflows/test_minimal_retrain.py index be49c89..94d86da 100644 --- a/tests/laborious/workflows/test_minimal_retrain.py +++ b/tests/laborious/workflows/test_minimal_retrain.py @@ -1,4 +1,4 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, AsyncMock, call, patch from pytest import fixture, mark @@ -39,8 +39,15 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain): }, } - storage_result = MagicMock() - storage_result.has_data.return_value = True + storage_result = { + 'last_timestamp': '2024-01-01 00:00:00+0000', + 'status': {'success': True}, + 'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}}, + 'bucket': None, + 'object_key': None, + 'object_prefix': None, + 'uri': None, + } workflow_mock.execute_activity_method = AsyncMock( side_effect=[ @@ -163,8 +170,15 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim }, } - storage_result = MagicMock() - storage_result.has_data.return_value = False + storage_result = { + 'last_timestamp': '2024-01-01 00:00:00+0000', + 'status': {'success': True}, + 'data': {}, + 'bucket': None, + 'object_key': None, + 'object_prefix': None, + 'uri': None, + } workflow_mock.execute_activity_method = AsyncMock( side_effect=[ @@ -218,8 +232,15 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim }, } - storage_result = MagicMock() - storage_result.has_data.return_value = True + storage_result = { + 'last_timestamp': '2024-01-01 00:00:00+0000', + 'status': {'success': True}, + 'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}}, + 'bucket': None, + 'object_key': None, + 'object_prefix': None, + 'uri': None, + } workflow_mock.execute_activity_method = AsyncMock( side_effect=[ diff --git a/tests/laborious/workflows/test_predictions_batch.py b/tests/laborious/workflows/test_predictions_batch.py index 279499f..de52ad8 100644 --- a/tests/laborious/workflows/test_predictions_batch.py +++ b/tests/laborious/workflows/test_predictions_batch.py @@ -12,12 +12,10 @@ def predictions_batch() -> PredictionsBatch: metadata = { - 'metadata': { - 'model_id': 'test_model_id', - 'model_name': 'test_model', - 'workflow_name': 'predictions_batch', - 'schedule_name': 'test_schedule', - }, + 'model_id': 'test_model_id', + 'model_name': 'test_model', + 'workflow_name': 'predictions_batch', + 'schedule_name': 'test_schedule', } @@ -49,7 +47,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch call( Activities.load_query_with_minio_offload, { - **metadata, + 'metadata': metadata, 'query': input_data['query'], 'datetime_columns': input_data.get('datetime_columns', []), 'model_name': input_data['model_name'], @@ -60,19 +58,39 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch ] ) prediction_input = { - 'metadata': metadata, + 'metadata': {'metadata': metadata}, 'data': activity_return, 'schema': input_data['schema'], 'table_name': input_data['table_name'], 'transform_table_name': input_data['transform_table_name'], 'model_id': input_data['model_id'], 'model_name': input_data['model_name'], - 'input_filters': input_data.get('input_filters', {'EMPTY_DATA': {'POLICY': 'STOP'}}), + 'input_filters': input_data.get( + 'input_filters', + { + 'EMPTY_DATA': { + 'POLICY': 'STOP', + 'CONFIG': {}, + } + }, + ), 'mlflow_transform_filters': input_data.get( - 'mlflow_transform_filters', {'API_ERROR': {'POLICY': 'STOP'}} + 'mlflow_transform_filters', + { + 'API_ERROR': { + 'POLICY': 'STOP', + 'CONFIG': {}, + } + }, ), 'mlflow_predict_filters': input_data.get( - 'mlflow_predict_filters', {'API_ERROR': {'POLICY': 'STOP'}} + 'mlflow_predict_filters', + { + 'API_ERROR': { + 'POLICY': 'STOP', + 'CONFIG': {}, + } + }, ), 'model_config': input_data.get('model_config', {}), 'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']), diff --git a/tests/laborious/workflows/test_simple_metrics.py b/tests/laborious/workflows/test_simple_metrics.py index 0d96bc5..dbda735 100644 --- a/tests/laborious/workflows/test_simple_metrics.py +++ b/tests/laborious/workflows/test_simple_metrics.py @@ -42,9 +42,8 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics): target_data = {'data': 'test_target_data'} simple_metrics_data = {'metrics': 'test_simple_metrics_data'} - workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data] - - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None]) + workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data) # Act await simple_metrics.run(input_data) @@ -66,7 +65,7 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics): p."timestamp" desc; """ - workflow_mock.execute_local_activity_method.assert_has_calls( + workflow_mock.execute_activity_method.assert_has_calls( [ call( Activities.load_custom_query, @@ -80,29 +79,30 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics): start_to_close_timeout=ANY, ), call( - Activities.calculate_simple_metrics, + Activities.export_data_to_postgres, { **metadata, - 'model_id': input_data['model_id'], - 'target_data': target_data, - 'metrics': input_data['metrics'], - 'interval_minutes': input_data['interval_minutes'], + 'data': simple_metrics_data, + 'schema': input_data['schema'], + 'table_name': input_data['target_table_name'], + 'timestamp_conversion': { + 'column': 'timestamp', + 'format': DATETIME_FORMAT_WITH_TZ, + }, }, retry_policy=ANY, start_to_close_timeout=ANY, ), ] ) - - # Assert - Check export_data_to_postgres call - workflow_mock.execute_activity_method.assert_called_once_with( - Activities.export_data_to_postgres, + workflow_mock.execute_local_activity_method.assert_called_once_with( + Activities.calculate_simple_metrics, { **metadata, - 'data': simple_metrics_data, - 'schema': input_data['schema'], - 'table_name': input_data['target_table_name'], - 'timestamp_conversion': {'column': 'timestamp', 'format': DATETIME_FORMAT_WITH_TZ}, + 'model_id': input_data['model_id'], + 'target_data': target_data, + 'metrics': input_data['metrics'], + 'interval_minutes': input_data['interval_minutes'], }, retry_policy=ANY, start_to_close_timeout=ANY, @@ -128,15 +128,13 @@ async def test_run_empty_target_data(workflow_mock: AsyncMock, simple_metrics: S target_data = None - workflow_mock.execute_local_activity_method.return_value = target_data - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.execute_activity_method = AsyncMock(return_value=target_data) # Act await simple_metrics.run(input_data) # Assert - Should not call calculate_simple_metrics or export - assert workflow_mock.execute_local_activity_method.call_count == 1 - workflow_mock.execute_activity_method.assert_not_called() + assert workflow_mock.execute_activity_method.call_count == 1 @mark.asyncio @@ -159,16 +157,15 @@ async def test_run_empty_simple_metrics(workflow_mock: AsyncMock, simple_metrics target_data = {'data': 'test_target_data'} simple_metrics_data = None - workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data] - - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.execute_activity_method = AsyncMock(return_value=target_data) + workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data) # Act await simple_metrics.run(input_data) # Assert - Should call calculate_simple_metrics but not export - assert workflow_mock.execute_local_activity_method.call_count == 2 - workflow_mock.execute_activity_method.assert_not_called() + workflow_mock.execute_activity_method.assert_called_once() + workflow_mock.execute_local_activity_method.assert_called_once() @mark.asyncio @@ -191,33 +188,28 @@ async def test_run_default_metrics(workflow_mock: AsyncMock, simple_metrics: Sim target_data = {'data': 'test_target_data'} simple_metrics_data = {'metrics': 'test_simple_metrics_data'} - workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data] - - workflow_mock.execute_activity_method = AsyncMock() + workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None]) + workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data) # Act await simple_metrics.run(input_data) # Assert - Check calculate_simple_metrics call with default metrics - workflow_mock.execute_local_activity_method.assert_has_calls( - [ - call( - Activities.load_custom_query, - ANY, - retry_policy=ANY, - start_to_close_timeout=ANY, - ), - call( - Activities.calculate_simple_metrics, - { - **metadata, - 'model_id': input_data['model_id'], - 'target_data': target_data, - 'metrics': ['rmse', 'mse', 'mae', 'r2'], # Default value - 'interval_minutes': input_data['interval_minutes'], - }, - retry_policy=ANY, - start_to_close_timeout=ANY, - ), - ] + workflow_mock.execute_activity_method.assert_any_call( + Activities.load_custom_query, + ANY, + retry_policy=ANY, + start_to_close_timeout=ANY, + ) + workflow_mock.execute_local_activity_method.assert_called_once_with( + Activities.calculate_simple_metrics, + { + **metadata, + 'model_id': input_data['model_id'], + 'target_data': target_data, + 'metrics': ['rmse', 'mse', 'mae', 'r2'], # Default value + 'interval_minutes': input_data['interval_minutes'], + }, + retry_policy=ANY, + start_to_close_timeout=ANY, )