From e6018af23f2240e5bbdd2312787d637e6a7ac115 Mon Sep 17 00:00:00 2001 From: vitor-aignosi Date: Thu, 7 May 2026 17:02:25 -0300 Subject: [PATCH] SIENTIAPDE-1646 Update README, requirements, and E2E tests for improved configuration and functionality - Enhanced the README with updated model configuration examples, including the addition of an alias for production. - Removed the `requirements-light.txt` file and updated `requirements-local.txt` and `requirements.txt` to replace `asyncua` with `opcua`. - Refactored E2E test scenarios to utilize scenario input files for better maintainability and clarity. - Improved test coverage for MinIO offload functionality and added new helper functions for loading scenario inputs. - Updated `values.yaml` to reflect new global configurations and environment variables for the laborious worker. --- README.md | 97 +- e2e/conftest.py | 901 ++++++++++-------- e2e/helpers.py | 177 ++++ e2e/scenario_inputs/drift_base.json | 14 + e2e/scenario_inputs/format_export_base.json | 37 + e2e/scenario_inputs/main_happy_path.json | 45 + .../main_invalid_datetime.json | 45 + .../main_missing_required.json | 16 + e2e/scenario_inputs/main_sql_error.json | 44 + e2e/scenario_inputs/minimal_retrain_base.json | 12 + .../minio_offload_load_query.json | 11 + .../minio_offload_workflow.json | 37 + .../prediction_process_base.json | 39 + e2e/scenario_inputs/simple_metrics_base.json | 14 + e2e/scenarios.md | 788 +++++++-------- e2e/test_drift.py | 786 +++++++++++++++ e2e/test_minimal_retrain.py | 335 +++++++ e2e/test_minio_offload.py | 77 +- e2e/test_predictions_batch_format_export.py | 89 +- e2e/test_predictions_batch_main_workflow.py | 129 +-- ...st_predictions_batch_prediction_process.py | 71 +- e2e/test_simple_metrics.py | 322 +++++++ input_sample.json | 72 ++ laborious/activities/activities.py | 6 +- laborious/activities/api.py | 24 +- laborious/activities/gates.py | 70 +- laborious/activities/mlflow.py | 49 +- laborious/activities/model_metrics.py | 148 ++- laborious/activities/opc.py | 210 ++-- laborious/activities/storage.py | 72 +- .../utils/models/minio_dataframe_payload.py | 30 +- laborious/utils/repository/minio_manager.py | 32 - laborious/utils/repository/opc_repository.py | 281 +++--- laborious/worker/worker.py | 4 +- .../format_and_export_prediction.py | 2 +- .../sub_workflows/prediction_process.py | 2 +- requirements-light.txt | 12 - requirements-local.txt | 9 +- requirements.txt | 2 +- tests/conftest.py | 13 + tests/laborious/activities/test_activities.py | 11 +- tests/laborious/activities/test_api.py | 96 +- tests/laborious/activities/test_gates.py | 238 +++-- tests/laborious/activities/test_mlflow.py | 168 ++-- .../activities/test_model_metrics.py | 197 ++-- tests/laborious/activities/test_opc.py | 206 ++-- tests/laborious/activities/test_storage.py | 141 +-- .../models/test_minio_dataframe_payload.py | 67 +- .../utils/repository/test_opc_repository.py | 279 +++--- tests/laborious/worker/test_worker.py | 6 +- values.yaml | 535 +++++------ 51 files changed, 4408 insertions(+), 2660 deletions(-) create mode 100644 e2e/scenario_inputs/drift_base.json create mode 100644 e2e/scenario_inputs/format_export_base.json create mode 100644 e2e/scenario_inputs/main_happy_path.json create mode 100644 e2e/scenario_inputs/main_invalid_datetime.json create mode 100644 e2e/scenario_inputs/main_missing_required.json create mode 100644 e2e/scenario_inputs/main_sql_error.json create mode 100644 e2e/scenario_inputs/minimal_retrain_base.json create mode 100644 e2e/scenario_inputs/minio_offload_load_query.json create mode 100644 e2e/scenario_inputs/minio_offload_workflow.json create mode 100644 e2e/scenario_inputs/prediction_process_base.json create mode 100644 e2e/scenario_inputs/simple_metrics_base.json create mode 100644 e2e/test_drift.py create mode 100644 e2e/test_minimal_retrain.py create mode 100644 e2e/test_simple_metrics.py create mode 100644 input_sample.json delete mode 100644 laborious/utils/repository/minio_manager.py delete mode 100644 requirements-light.txt diff --git a/README.md b/README.md index bf037dc..c864f7c 100644 --- a/README.md +++ b/README.md @@ -463,7 +463,7 @@ flowchart LR "source_table_name": "laborious_data", "target_table_name": "drift_metrics", "interval": 60, - "model_config": { "target": "temperature" }, + "model_config": { "target": "temperature", "alias": "production" }, "drift_metrics": ["kolmogorov_smirnov", "jensen_shannon", "wasserstein"], "chunk_period": "min" } @@ -498,7 +498,7 @@ flowchart LR "data_table_name": "laborious_data", "target_table_name": "simple_metrics", "interval_minutes": 60, - "model_config": { "target": "temperature" }, + "model_config": { "target": "temperature", "alias": "production" }, "metrics": ["rmse", "mse", "mae", "r2"] } ``` @@ -953,7 +953,7 @@ For single OPC server, use individual environment variables: ### PI Web API Configuration -PI Web API configuration is built from environment variables using the `build_api_config` function from `sientia_do.connectors_config`. The configuration includes: +PI Web API configuration is built from environment variables using the `build_api_config` function from `sientia_do.utils.connectors_config`. The configuration includes: - `PI_WEB_API_BASE_URL`: Base URL of the PI Web API server - `PI_WEB_API_AUTH_TYPE`: Authentication type ('basic' or 'bearer') @@ -986,9 +986,30 @@ Where: MongoDB pipeline configuration: -#### Predictions Batch Workflow configuration sample +#### MongoDB input samples (updated) -This is the configuration for the Predictions Batch Workflow, to be inserted into the MongoDB pipeline collection. +Updated examples are available in `input_sample.json` at the repository root. +The sample already reflects the runtime-aware and alias-based flow: + +- `model_config` uses `target`, `retention_minutes`, and `alias`. +- `transform_flavor` / `predict_flavor` are not used anymore. + +Example model document: + +```json +{ + "id": "4", + "name": "vcm-nox", + "active": false, + "model_config": { + "alias": "production", + "retention_minutes": 60, + "target": "CI-W3W01A3" + } +} +``` + +Example predictions_batch schedule document: ```json { @@ -998,41 +1019,67 @@ This is the configuration for the Predictions Batch Workflow, to be inserted int "frequency": "30s", "max_retry_policy": 1, "query": "select * from sientia_data.laborious_data where model_id = 1 and \"timestamp\" > NOW() - INTERVAL '5 minutes' order by \"timestamp\" desc limit 30;", + "retention_time": 60, "write_tags": [ { - "server_id": "server1", + "server_id": "1", "type": "prediction", "addr": "ns=2;i=5", "data_type": "double" }, { - "server_id": "server1", + "server_id": "1", "type": "confidence", - "addr": "ns=2;i=6", + "addr": "ns=2;i=5", "data_type": "double" } ], - "input_filters": { - "EMPTY_DATA": {"POLICY": "STOP"}, - "SPECIFIC_VARIABLES_NULL_VALUES": { - "POLICY": "CONTINUE", - "config": {"variables": ["Counter"]} + "input_filters": [ + { + "filter_name": "EMPTY_DATA", + "policy": "STOP" + }, + { + "filter_name": "SPECIFIC_VARIABLES_NULL_VALUES", + "policy": "CONTINUE", + "config": { + "variables": ["Counter"] + } } - }, - "mlflow_transform_filters": { - "API_ERROR": {"POLICY": "REPEAT"}, - "NAN_VALUES": {"POLICY": "STOP"} - }, - "mlflow_predict_filters": { - "API_ERROR": {"POLICY": "CONTINUE"} - }, + ], + "mlflow_transform_filters": [ + { + "filter_name": "API_ERROR", + "policy": "REPEAT" + }, + { + "filter_name": "NAN_VALUES", + "policy": "STOP" + } + ], + "mlflow_predict_filters": [ + { + "filter_name": "API_ERROR", + "policy": "CONTINUE" + } + ], "path_priority": ["STOP", "CONTINUE", "REPEAT"], "active": true, "updated_at": { - "$date": "2025-09-16T10:00:00.000Z" + "$date": "2026-01-27T17:35:01.600Z" }, - "datetime_columns": ["timestamp", "created_at"], - "predictions_storage_policy": "lts:1" + "save_transform": false, + "pi_web_api_output_config": { + "endpoint": "/streamsets/value", + "prediction_tags": {}, + "confidence_tags": {} + }, + "model_config": { + "alias": "production", + "retention_minutes": 60, + "target": "CI-W3W01A3" + }, + "datetime_columns": ["timestamp", "created_at"] } ``` @@ -1125,7 +1172,7 @@ laborious/ 2. **MLFlow Connection Issues** - Verify MLFlow server is running and accessible - Check authentication credentials and permissions - - Ensure model names and versions exist + - Ensure model names exist and the expected alias (for example `production`) is registered 3. **Database Connection Issues** - Verify PostgreSQL service is running diff --git a/e2e/conftest.py b/e2e/conftest.py index 0a55f00..e66febc 100644 --- a/e2e/conftest.py +++ b/e2e/conftest.py @@ -1,37 +1,134 @@ -""" -Pytest configuration and fixtures for E2E tests. -""" +"""Pytest configuration and fixtures for E2E tests.""" -from unittest.mock import AsyncMock, MagicMock, patch -from io import BytesIO +from concurrent.futures import ThreadPoolExecutor +from typing import Any +from unittest.mock import MagicMock import pandas as pd import pytest import pytest_asyncio from sqlalchemy import create_engine, text +from testcontainers.core.container import DockerContainer from testcontainers.minio import MinioContainer from testcontainers.postgres import PostgresContainer from temporalio.testing import WorkflowEnvironment from temporalio.worker import Worker -from laborious.activities.activities import Activities + +class _FakeModelAnalysis: + """ + Configurable double for ``sientia_model.analytics.model_analysis.ModelAnalysis`` + used by E2E drift tests. + + Replaces the real analyzer in ``sientia_model.analytics.model_analysis`` + BEFORE the production import chain runs so + ``laborious.activities.model_metrics`` resolves ``ModelAnalysis`` to this + class at import time. Keeps drift assertions stable across runs (the real + analyzer is data-dependent). + + Tests configure responses through class-level attributes which are reset + between tests by ``reset_model_analysis_stub``: + + - ``_drift_response_factory``: callable ``(univariate, multivariate) -> DataFrame`` + controlling the consolidated drift dataframe seen by ``calculate_drift``. + - ``_univariate_side_effect`` / ``_multivariate_side_effect``: optional + side effects (Exception or callable) for the raw detection methods. + - ``_drift_metrics_dataframe_exception``: when set, raised by + ``get_drift_metrics_dataframe`` to simulate analyzer failures. + """ + + _drift_response_factory: Any = None + _univariate_side_effect: Any = None + _multivariate_side_effect: Any = None + _drift_metrics_dataframe_exception: Exception | None = None + last_instance: Any = None + + def __init__(self, config): + self.config = config + type(self).last_instance = self + self.detect_univariate_drift_calls = [] + self.detect_multivariate_drift_calls = [] + self.get_drift_metrics_dataframe_calls = [] + + def detect_univariate_drift(self, **kwargs): + """Record arguments and return ``{}`` unless ``_univariate_side_effect`` overrides it.""" + self.detect_univariate_drift_calls.append(kwargs) + side_effect = type(self)._univariate_side_effect + if isinstance(side_effect, Exception): + raise side_effect + if callable(side_effect): + return side_effect(**kwargs) + return {} + + def detect_multivariate_drift(self, **kwargs): + """Record arguments and return ``{}`` unless ``_multivariate_side_effect`` overrides it.""" + self.detect_multivariate_drift_calls.append(kwargs) + side_effect = type(self)._multivariate_side_effect + if isinstance(side_effect, Exception): + raise side_effect + if callable(side_effect): + return side_effect(**kwargs) + return {} + + def get_drift_metrics_dataframe(self, univariate_drift, multivariate_drift): + """Return the configured drift dataframe (copy) or raise the configured exception.""" + self.get_drift_metrics_dataframe_calls.append( + {'univariate_drift': univariate_drift, 'multivariate_drift': multivariate_drift} + ) + if type(self)._drift_metrics_dataframe_exception is not None: + raise type(self)._drift_metrics_dataframe_exception + factory = type(self)._drift_response_factory + if factory is None: + return pd.DataFrame() + result = factory(univariate_drift, multivariate_drift) + return result.copy() if isinstance(result, pd.DataFrame) else result + + @classmethod + def reset(cls): + """Clear all configured side effects and the last constructed instance.""" + cls._drift_response_factory = None + cls._univariate_side_effect = None + cls._multivariate_side_effect = None + cls._drift_metrics_dataframe_exception = None + cls.last_instance = None + + +# Patch ``sientia_model.analytics.model_analysis.ModelAnalysis`` BEFORE the +# production import chain runs so drift E2E tests can drive deterministic +# analyzer outputs (the real implementation is data-dependent and would yield +# values that drift across runs). The patch is applied to the real installed +# module so that ``from laborious.activities.activities import Activities`` +# resolves ``ModelAnalysis`` to ``_FakeModelAnalysis`` at import time. +import sientia_model.analytics.model_analysis as _sientia_model_analysis_module # noqa: E402 + +_sientia_model_analysis_module.ModelAnalysis = _FakeModelAnalysis + +from laborious.activities.activities import Activities # noqa: E402 +from laborious.workflows.drift import Drift +from laborious.workflows.minimal_retrain import MinimalRetrain from laborious.workflows.predictions_batch import PredictionsBatch +from laborious.workflows.simple_metrics import SimpleMetrics +from laborious.workflows.sub_workflows.format_and_export_prediction import ( + FormatAndExportPrediction, +) from laborious.workflows.sub_workflows.prediction_process import PredictionProcess -from laborious.workflows.sub_workflows.format_and_export_prediction import FormatAndExportPrediction from sientia_do.notifications.handlers import CoreNotificationHandler from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController -# Test constants -TEST_MONGODB_CONNECTION_STRING = 'mongodb://localhost:27017' -TEST_DATABASE_NAME = 'test_db' + +@pytest_asyncio.fixture(scope='session') +def postgres_container(): + """PostgreSQL testcontainer used by all E2E tests.""" + postgres = PostgresContainer('postgres:15') + postgres.start() + yield postgres + postgres.stop() @pytest_asyncio.fixture(scope='session') def minio_container(): - """ - MinIO S3-compatible storage for E2E tests that exercise real offload uploads. - """ + """MinIO testcontainer used by E2E offload and payload retrieval paths.""" minio = MinioContainer() minio.start() yield minio @@ -39,307 +136,206 @@ def minio_container(): @pytest_asyncio.fixture(scope='session') -def postgres_container(): - """ - Create a PostgreSQL container using testcontainers. - - This fixture creates a real PostgreSQL database in a Docker container - that will be used for all tests in the session. - """ - postgres = PostgresContainer('postgres:15') - postgres.start() - yield postgres - postgres.stop() +def mongo_container(): + """MongoDB testcontainer used by real CoreNotificationHandler.""" + mongo = DockerContainer('mongo:7').with_exposed_ports(27017) + mongo.start() + yield mongo + mongo.stop() @pytest_asyncio.fixture def postgres_engine(postgres_container): - """ - Create SQLAlchemy engine for PostgreSQL test database. - - This fixture creates a connection to the PostgreSQL container - created by the postgres_container fixture. - """ + """SQLAlchemy engine bound to the PostgreSQL testcontainer.""" engine = create_engine(postgres_container.get_connection_url()) - yield engine - engine.dispose() def _create_schema_and_tables(engine): - """ - Helper function to create schema and tables in the given engine. - - Creates predictions_schema with: - - laborious_data: Input data table for queries - - predictions: Output predictions table - - transformed_data: Output transformed data table - """ - # Use begin() to ensure transaction is properly committed + """Create all schemas/tables required by workflow and activity paths.""" with engine.begin() as conn: - # Create predictions_schema - conn.execute(text("CREATE SCHEMA IF NOT EXISTS predictions_schema")) - - # Create laborious_data table (input data from sensors) - create_laborious_data_sql = """ - CREATE TABLE IF NOT EXISTS predictions_schema.laborious_data ( - id SERIAL NOT NULL, - model_id int4 NOT NULL, - variable text NOT NULL, - value numeric NULL, - "timestamp" timestamptz NOT NULL, - created_at timestamptz NOT NULL, - PRIMARY KEY (id) - ); - """ - conn.execute(text(create_laborious_data_sql)) - - # Create predictions table - create_predictions_sql = """ - CREATE TABLE if not exists predictions_schema.predictions ( - id SERIAL NOT NULL , - model_id int4 NOT NULL, - prediction numeric NULL, - prediction_confidence numeric NOT NULL, - response_time numeric NOT NULL, - prediction_status text NOT NULL, - "timestamp" timestamptz NOT NULL, - created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, - "comments" text NULL, - PRIMARY KEY (id, created_at) - ); - """ - conn.execute(text(create_predictions_sql)) - - # Create transformed_data table - create_transformed_sql = """ - CREATE TABLE IF NOT EXISTS predictions_schema.transformed_data ( - id SERIAL NOT NULL, - model_id int4 NOT NULL, - variable text NOT NULL, - value numeric NULL, - "timestamp" timestamptz NOT NULL, - created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, - PRIMARY KEY (id) - ); - """ - conn.execute(text(create_transformed_sql)) + conn.execute(text('CREATE SCHEMA IF NOT EXISTS predictions_schema')) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.laborious_data ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + variable text NOT NULL, + value numeric NULL, + "timestamp" timestamptz NOT NULL, + created_at timestamptz NOT NULL, + PRIMARY KEY (id) + ); + """ + ) + ) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.predictions ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + prediction numeric NULL, + prediction_confidence numeric NOT NULL, + response_time numeric NOT NULL, + prediction_status text NOT NULL, + "timestamp" timestamptz NOT NULL, + created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + comments text NULL, + PRIMARY KEY (id, created_at) + ); + """ + ) + ) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.transformed_data ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + variable text NOT NULL, + value numeric NULL, + "timestamp" timestamptz NOT NULL, + created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + PRIMARY KEY (id) + ); + """ + ) + ) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.drift ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + feature text NOT NULL, + method text NOT NULL, + value numeric NULL, + drift bool NOT NULL, + chunk int4 NOT NULL, + "timestamp" timestamptz NOT NULL, + timestamp_end text NULL, + accurate bool NOT NULL, + created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + updated_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + PRIMARY KEY (id) + ); + """ + ) + ) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.simple_metrics_data ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + metric text NOT NULL, + value numeric NOT NULL, + "timestamp" timestamptz NOT NULL, + data_size int4 NOT NULL, + interval_minutes int4 NOT NULL, + created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + PRIMARY KEY (id) + ); + """ + ) + ) + conn.execute( + text( + """ + CREATE TABLE IF NOT EXISTS predictions_schema.retrain_reports ( + id SERIAL NOT NULL, + model_id int4 NOT NULL, + model_name text NOT NULL, + "timestamp" text NOT NULL, + status text NOT NULL, + version text NULL, + mlflow_run_id text NULL, + mlflow_experiment_id text NULL, + created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL, + PRIMARY KEY (id) + ); + """ + ) + ) @pytest_asyncio.fixture(autouse=True) def setup_postgres_schema_and_tables(postgres_engine): - """ - Automatically create necessary schema and tables before each test. - - This fixture runs automatically (autouse=True) and ensures - that the predictions_schema and tables exist with the correct structure. - """ + """Ensure required schema and tables exist before each E2E test.""" _create_schema_and_tables(postgres_engine) yield +@pytest.fixture(autouse=True) +def reset_model_analysis_stub(): + """Reset class-level state on the ModelAnalysis stub between tests.""" + _FakeModelAnalysis.reset() + yield + _FakeModelAnalysis.reset() + + @pytest_asyncio.fixture def mock_logger(): - """Mock logger for testing.""" - def message(message): - print(f"[LOG] {message}") - def custom_message(message, _metadata={}): - print(f"[LOG] {message}") - logger = MagicMock() - logger.info = MagicMock( - side_effect=message - ) - logger.debug = MagicMock( - side_effect=message - ) - logger.error = MagicMock( - side_effect=message - ) - logger.warning = MagicMock( - side_effect=message - ) - logger.custom_info = MagicMock( - side_effect=custom_message - ) - logger.custom_debug = MagicMock( - side_effect=custom_message - ) - logger.custom_error = MagicMock( - side_effect=custom_message - ) - logger.custom_warning = MagicMock( - side_effect=custom_message - ) + """Logger double with readable console output for E2E runs.""" + logger = MagicMock(spec=Logger) + logger.info = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}')) + logger.debug = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}')) + logger.error = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}')) + logger.warning = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}')) + logger.custom_info = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}')) + logger.custom_debug = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}')) + logger.custom_error = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}')) + logger.custom_warning = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}')) return logger -@pytest_asyncio.fixture -def mock_mongo_client(): - """ - Mock MongoDB client to avoid real connections. - - This fixture mocks the pymongo.MongoClient used by CoreNotificationHandler, - allowing us to use a real NotificationHandler instance without connecting to MongoDB. - """ - mock_client = MagicMock() - mock_db = MagicMock() - mock_collection = MagicMock() - - # Configure the mock chain: client[database] -> db[collection] -> collection - mock_client.__getitem__.return_value = mock_db - mock_db.__getitem__.return_value = mock_collection - - # Mock server_info() to avoid connection attempts - mock_client.server_info = MagicMock() - - # Mock insert_one for notifications - mock_collection.insert_one = MagicMock() - - return mock_client - - -@pytest.fixture -def notification_inserts(mock_mongo_client): - """ - Mongo insert_one mock used by CoreNotificationHandler for notification persistence. - - Yields: - MagicMock for insert_one, reset before each test. - """ - mock_db = mock_mongo_client.__getitem__.return_value - mock_collection = mock_db.__getitem__.return_value - mock_collection.insert_one.reset_mock() - yield mock_collection.insert_one - - -@pytest_asyncio.fixture -def notification_handler(mock_logger, mock_mongo_client): - """ - Create a real NotificationHandler instance with mocked MongoDB client. - - This fixture creates a real CoreNotificationHandler instance but mocks - the underlying MongoDB connection to avoid real database connections. - """ - # Patch MongoClient where it's imported in the handlers module - with patch('sientia_do.notifications.handlers.MongoClient', return_value=mock_mongo_client): - handler = CoreNotificationHandler( - connection_string=TEST_MONGODB_CONNECTION_STRING, - database=TEST_DATABASE_NAME, - logger=mock_logger, - project_name='laborious', - ) - yield handler - handler.shutdown() - - @pytest_asyncio.fixture def metrics_controller(mock_logger): - """Create a real MetricsController instance.""" + """Real metrics controller for E2E observability paths.""" return MetricsController(logger=mock_logger) @pytest_asyncio.fixture -def mock_minio_repository(): - """Mock MinIO repository for object storage operations.""" - mock_repo = MagicMock() - - # Provide at least valid parquet bytes so that MinioDataFramePayload.retrieve() - # can decode the payload if offloading is exercised in an integration scenario. - parquet_df = pd.DataFrame({'a': [1]}) - parquet_buffer = BytesIO() - parquet_df.to_parquet(parquet_buffer, engine='pyarrow', index=True) - parquet_bytes = parquet_buffer.getvalue() - - # sientia_do MinioRepository API - mock_repo.bucket = 'test-bucket' - mock_repo.upload_file = AsyncMock( - side_effect=lambda file_bytes, relative_key, content_type='application/octet-stream', bucket=None, metadata=None: { - 'minio_object_name': f'sientia/streamlit-connectors/{relative_key}', - 'original_filename': relative_key.rsplit('/', 1)[-1], - 'uploaded_at': '2024-01-01T00:00:00Z', - 'sha256_hash': 'deadbeef', - } +def notification_handler(mock_logger, mongo_container): + """Real notification handler using MongoDB testcontainer.""" + mongo_port = mongo_container.get_exposed_port(27017) + handler = CoreNotificationHandler( + connection_string=f'mongodb://localhost:{mongo_port}', + database='test_db', + logger=mock_logger, + project_name='laborious', ) - mock_repo.download_file = AsyncMock(return_value=parquet_bytes) - mock_repo.list_objects = AsyncMock(return_value=[]) - mock_repo.delete_file = AsyncMock() - mock_repo.close = MagicMock() - - return mock_repo + try: + yield handler + finally: + handler.shutdown() -@pytest_asyncio.fixture -def mock_pi_web_api_repository(): - """Mock PI Web API repository for PI Web API operations.""" - mock_repo = MagicMock() - - 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 - -@pytest_asyncio.fixture -def mock_opc_repository(): - """Mock OPC repository for OPC operations.""" - mock_repo = MagicMock() - mock_repo.write_data = AsyncMock( - return_value=(True, {'response_time': 0.1}) - ) - mock_repo.disconnect = AsyncMock() - return mock_repo - -@pytest_asyncio.fixture -def patch_create_engine(postgres_engine): - """Patch create_engine to return test postgres_engine.""" - with patch('sientia_do.temporal.activities.postgres.create_engine', return_value=postgres_engine): - yield +@pytest.fixture +def notification_inserts(notification_handler): + """Spy on real Mongo insert calls issued by notification handler.""" + collection = notification_handler.mongo_collection + original_insert_one = collection.insert_one + spy = MagicMock(wraps=original_insert_one) + collection.insert_one = spy + try: + yield spy + finally: + collection.insert_one = original_insert_one -@pytest_asyncio.fixture -def patch_minio_repository(mock_minio_repository): - """Patch MinioRepository to return mock.""" - # Patch where Activities resolves the symbol (import binds the original class). - with patch('laborious.activities.activities.MinioRepository', return_value=mock_minio_repository): - yield +class _FakeModelWrapper: + """External MLflow wrapper double used by repository stub.""" -@pytest_asyncio.fixture -def patch_pi_web_api_repository(mock_pi_web_api_repository): - """Patch PI Web API client to return mock.""" - with patch('laborious.activities.api.PIWebAPIClient', return_value=mock_pi_web_api_repository): - yield + def __init__(self): + self.transform = MagicMock(side_effect=self._default_transform) + self.predict = MagicMock(side_effect=self._default_predict) - -@pytest_asyncio.fixture -def plugin_store_stub(): - """ - PluginStore stub for Activities construction. - - Runtime installation happens in the worker process; activities only hold a reference. - """ - - return MagicMock() - - -@pytest_asyncio.fixture -def mlflow_repository_stub(): - """ - SientiaMLflowRepository stub that returns a SientiaModel-like wrapper for E2E tests. - - Transform/predict mirror the legacy sklearn/pyfunc mock behavior using pandas outputs. - """ - - repo = MagicMock() - - def _transform_side_effect(data: pd.DataFrame): + @staticmethod + def _default_transform(data: pd.DataFrame): result = pd.DataFrame( { 'feature_1': [0.234] * len(data), @@ -349,115 +345,94 @@ def mlflow_repository_stub(): result.index = data.index return result, {} - def _predict_side_effect(_params: dict, data: pd.DataFrame): + @staticmethod + def _default_predict(_params: dict, data: pd.DataFrame): pred = pd.DataFrame([0.5] * len(data), columns=['placeholder']) pred.index = data.index return pred, {} - wrapper = MagicMock() - wrapper.transform.side_effect = _transform_side_effect - wrapper.predict.side_effect = _predict_side_effect - repo.get_cached_model = MagicMock(return_value=wrapper) +@pytest_asyncio.fixture +def mlflow_repository_stub(): + """External MLflow repository stub.""" + repo = MagicMock() + wrapper = _FakeModelWrapper() repo.stub_wrapper = wrapper + repo.get_cached_model = MagicMock(return_value=wrapper) repo._client = MagicMock() return repo +class _FakePIWebAPIClient: + """External PI Web API client stub with deterministic responses.""" + + def __init__(self): + self._responses = None + self.write_value = MagicMock(side_effect=self._write_value) + self.close = MagicMock() + + def set_side_effect(self, side_effect): + self._responses = side_effect + + def _write_value(self, web_ids, value, metadata=None, **kwargs): + if isinstance(self._responses, Exception): + raise self._responses + if isinstance(self._responses, list): + item = self._responses.pop(0) + if isinstance(item, Exception): + raise item + return item + if callable(self._responses): + return self._responses(web_ids=web_ids, value=value, metadata=metadata, **kwargs) + return [{'WebId': wid, 'Errors': []} for wid in web_ids] + + +@pytest_asyncio.fixture +def pi_web_api_client_stub(): + """PI Web API stub fixture.""" + return _FakePIWebAPIClient() + + +@pytest_asyncio.fixture +def opc_repository_stub(): + """OPC external dependency stub.""" + repo = MagicMock() + repo.write_data = MagicMock(return_value=(True, {'response_time': 0.1})) + repo.disconnect = MagicMock() + return repo + + +@pytest_asyncio.fixture +def plugin_store_stub(): + """Plugin store external dependency stub.""" + return MagicMock() + + @pytest_asyncio.fixture(scope='function') async def test_activities( - postgres_engine, - postgres_container, - mock_logger, - notification_handler, - metrics_controller, - mock_minio_repository, - patch_create_engine, - patch_minio_repository, - mlflow_repository_stub, - plugin_store_stub, - patch_pi_web_api_repository, - mock_opc_repository -): - """ - Create Activities instance with test dependencies. - - This fixture creates a real Activities instance with: - - PostgreSQL database (via testcontainers) - - Mocked MinIO client - - Real NotificationHandler and MetricsController (with mocked underlying services) - """ - activities = Activities( - postgres_config={ - 'host': 'localhost', - 'port': postgres_container.get_exposed_port(5432), - 'user': 'test', - 'password': 'test', - 'dbname': 'test', - 'min_connections': 1, - 'max_connections': 5, - }, - plugin_store=plugin_store_stub, - minio_config={ - # 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={ - 'base_url': 'http://localhost:8080', - 'auth_type': 'bearer', - 'auth_token': 'test_token', - }, - logger=mock_logger, - notification_handler=notification_handler, - metrics_controller=metrics_controller, - mlflow_repository=mlflow_repository_stub, - ) - - activities.opc_repository = { - '1': mock_opc_repository, - } - - try: - yield activities - finally: - # Cleanup - ALWAYS runs, even if test fails - await activities.shutdown() - - -@pytest_asyncio.fixture(scope='function') -async def test_activities_real_minio( - postgres_engine, postgres_container, minio_container, mock_logger, notification_handler, metrics_controller, - patch_create_engine, mlflow_repository_stub, plugin_store_stub, - patch_pi_web_api_repository, - mock_opc_repository, + pi_web_api_client_stub, + opc_repository_stub, ): - """ - Activities with a real MinIO testcontainer (no MinioRepository patch) for offload tests. - """ + """Activities with real infra and external-system stubs only.""" minio_client = minio_container.get_client() if not minio_client.bucket_exists('test-bucket'): minio_client.make_bucket('test-bucket') minio_port = minio_container.get_exposed_port(9000) + activities = Activities( postgres_config={ 'host': 'localhost', - 'port': postgres_container.get_exposed_port(5432), - 'user': 'test', - 'password': 'test', - 'dbname': 'test', + 'port': int(postgres_container.get_exposed_port(5432)), + 'user': postgres_container.username, + 'password': postgres_container.password, + 'dbname': postgres_container.dbname, 'min_connections': 1, 'max_connections': 5, }, @@ -481,16 +456,73 @@ async def test_activities_real_minio( metrics_controller=metrics_controller, mlflow_repository=mlflow_repository_stub, ) - activities.opc_repository = {'1': mock_opc_repository} + activities.pi_web_api_client = pi_web_api_client_stub + activities.opc_repository = {'1': opc_repository_stub} try: yield activities finally: - await activities.shutdown() + activities.shutdown() + + +@pytest_asyncio.fixture(scope='function') +async def test_activities_real_minio( + postgres_container, + minio_container, + mock_logger, + notification_handler, + metrics_controller, + mlflow_repository_stub, + plugin_store_stub, + pi_web_api_client_stub, + opc_repository_stub, +): + """Compatibility alias for offload tests.""" + minio_client = minio_container.get_client() + if not minio_client.bucket_exists('test-bucket'): + minio_client.make_bucket('test-bucket') + minio_port = minio_container.get_exposed_port(9000) + + activities = Activities( + postgres_config={ + 'host': 'localhost', + 'port': int(postgres_container.get_exposed_port(5432)), + 'user': postgres_container.username, + 'password': postgres_container.password, + 'dbname': postgres_container.dbname, + 'min_connections': 1, + 'max_connections': 5, + }, + plugin_store=plugin_store_stub, + minio_config={ + 'endpoint_url': f'localhost:{minio_port}', + 'access_key': 'minioadmin', + 'secret_key': 'minioadmin', + 'default_bucket': 'test-bucket', + 'retention_hours': 24, + 'secure': False, + }, + opc_config={}, + pi_web_api_config={ + 'base_url': 'http://localhost:8080', + 'auth_type': 'bearer', + 'auth_token': 'test_token', + }, + logger=mock_logger, + notification_handler=notification_handler, + metrics_controller=metrics_controller, + mlflow_repository=mlflow_repository_stub, + ) + activities.pi_web_api_client = pi_web_api_client_stub + activities.opc_repository = {'1': opc_repository_stub} + try: + yield activities + finally: + activities.shutdown() def _worker_activity_list(test_activities: Activities): + """List of registered activity callables used by Temporal worker in E2E.""" return [ - test_activities.load_custom_query, test_activities.load_query_with_minio_offload, test_activities.cleanup_minio_objects_expired, test_activities.input_gate, @@ -512,7 +544,7 @@ def _worker_activity_list(test_activities: Activities): @pytest_asyncio.fixture(scope='function') async def temporal_test_env(): - """Create Temporal test environment.""" + """Temporal test environment with time-skipping.""" env = await WorkflowEnvironment.start_time_skipping() async with env: yield env @@ -520,23 +552,134 @@ async def temporal_test_env(): @pytest_asyncio.fixture(scope='function') async def temporal_worker(temporal_test_env, test_activities): - """Create Temporal worker with test activities.""" - async with Worker( - temporal_test_env.client, - task_queue='test-queue', - workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction], - activities=_worker_activity_list(test_activities), - ) as worker: - yield worker + """Temporal worker for full predictions-batch and child workflows.""" + with ThreadPoolExecutor(max_workers=32) as activity_executor: + async with Worker( + temporal_test_env.client, + task_queue='test-queue', + workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction], + activities=_worker_activity_list(test_activities), + activity_executor=activity_executor, + ) as worker: + yield worker @pytest_asyncio.fixture(scope='function') async def temporal_worker_real_minio(temporal_test_env, test_activities_real_minio): - """Temporal worker backed by Activities using real MinIO testcontainer.""" - async with Worker( - temporal_test_env.client, - task_queue='test-queue', - workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction], - activities=_worker_activity_list(test_activities_real_minio), - ) as worker: - yield worker + """Temporal worker alias for tests that emphasize MinIO behavior.""" + with ThreadPoolExecutor(max_workers=32) as activity_executor: + async with Worker( + temporal_test_env.client, + task_queue='test-queue', + workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction], + activities=_worker_activity_list(test_activities_real_minio), + activity_executor=activity_executor, + ) as worker: + yield worker + + +def _drift_worker_activity_list(test_activities: Activities): + """Activity callables registered on the drift Temporal worker.""" + return [ + test_activities.load_custom_query, + test_activities.get_reference_data, + test_activities.calculate_drift, + test_activities.export_data_to_postgres, + ] + + +@pytest_asyncio.fixture(scope='function') +async def temporal_worker_drift(temporal_test_env, test_activities): + """Temporal worker registered with the Drift workflow and its activities.""" + with ThreadPoolExecutor(max_workers=32) as activity_executor: + async with Worker( + temporal_test_env.client, + task_queue='test-queue', + workflows=[Drift], + activities=_drift_worker_activity_list(test_activities), + activity_executor=activity_executor, + ) as worker: + yield worker + + +def _simple_metrics_worker_activity_list(test_activities: Activities): + """Activity callables registered on the simple-metrics Temporal worker.""" + return [ + test_activities.load_custom_query, + test_activities.calculate_simple_metrics, + test_activities.export_data_to_postgres, + ] + + +@pytest_asyncio.fixture(scope='function') +async def temporal_worker_simple_metrics(temporal_test_env, test_activities): + """Temporal worker registered with the SimpleMetrics workflow and its activities.""" + with ThreadPoolExecutor(max_workers=32) as activity_executor: + async with Worker( + temporal_test_env.client, + task_queue='test-queue', + workflows=[SimpleMetrics], + activities=_simple_metrics_worker_activity_list(test_activities), + activity_executor=activity_executor, + ) as worker: + yield worker + + +def _minimal_retrain_worker_activity_list(test_activities: Activities): + """Activity callables registered on the minimal-retrain Temporal worker.""" + return [ + test_activities.load_query_with_minio_offload, + test_activities.retrain_model, + test_activities.update_production_model, + test_activities.format_retrain_report, + test_activities.export_data_to_postgres, + ] + + +@pytest_asyncio.fixture(scope='function') +async def temporal_worker_minimal_retrain(temporal_test_env, test_activities): + """Temporal worker registered with the MinimalRetrain workflow and its activities.""" + with ThreadPoolExecutor(max_workers=32) as activity_executor: + async with Worker( + temporal_test_env.client, + task_queue='test-queue', + workflows=[MinimalRetrain], + activities=_minimal_retrain_worker_activity_list(test_activities), + activity_executor=activity_executor, + ) as worker: + yield worker + + +@pytest.fixture +def model_analysis_stub(): + """ + Expose the ``_FakeModelAnalysis`` class so drift tests can configure responses. + + Use class-level attributes to control the analyzer outputs: + + - ``stub.set_drift_dataframe(factory)`` to provide rows for ``calculate_drift``. + - ``stub.raise_on_get_drift_metrics_dataframe(exception)`` to simulate failures. + """ + + class _Helper: + """Thin convenience wrapper around _FakeModelAnalysis class state.""" + + cls = _FakeModelAnalysis + + def set_drift_dataframe(self, factory): + self.cls._drift_response_factory = factory + + def raise_on_get_drift_metrics_dataframe(self, exc: Exception): + self.cls._drift_metrics_dataframe_exception = exc + + def set_univariate_side_effect(self, side_effect): + self.cls._univariate_side_effect = side_effect + + def set_multivariate_side_effect(self, side_effect): + self.cls._multivariate_side_effect = side_effect + + @property + def last_instance(self): + return self.cls.last_instance + + return _Helper() diff --git a/e2e/helpers.py b/e2e/helpers.py index 1aaa4f2..0ada308 100644 --- a/e2e/helpers.py +++ b/e2e/helpers.py @@ -3,13 +3,60 @@ Shared helpers for E2E tests (Temporal workflows + PostgreSQL). """ import asyncio +import json from datetime import datetime from decimal import Decimal +from pathlib import Path from typing import Any +import pandas as pd from sqlalchemy import text from sqlalchemy.engine import Engine +SCENARIO_INPUTS_DIR = Path(__file__).parent / 'scenario_inputs' + + +def _replace_template_values(payload: Any, model_id: int) -> Any: + """ + Replace string placeholders in scenario payloads with the concrete model id. + + Args: + payload: JSON-like structure loaded from scenario input file. + model_id: Model id used to render template placeholders. + + Return: + Any: Payload with ``{{MODEL_ID}}`` replaced where applicable. + """ + if isinstance(payload, dict): + return {key: _replace_template_values(value, model_id) for key, value in payload.items()} + if isinstance(payload, list): + return [_replace_template_values(item, model_id) for item in payload] + if isinstance(payload, str): + if payload == '{{MODEL_ID}}': + return model_id + return payload.replace('{{MODEL_ID}}', str(model_id)) + return payload + + +def load_scenario_input(file_name: str, model_id: int | None = None) -> dict[str, Any]: + """ + Load a scenario input JSON from ``e2e/scenario_inputs``. + + Args: + file_name: JSON file name inside ``e2e/scenario_inputs``. + model_id: Optional model id used to render ``{{MODEL_ID}}`` placeholders. + + Return: + dict[str, Any]: Input payload ready to be passed to workflow/activity calls. + """ + file_path = SCENARIO_INPUTS_DIR / file_name + with file_path.open('r', encoding='utf-8') as f: + payload = json.load(f) + + if model_id is not None: + return _replace_template_values(payload, model_id) + return payload + async def start_and_await_workflow(client, workflow_run, input_data: dict, workflow_id: str, timeout: float = 60.0): """ @@ -172,3 +219,133 @@ def assert_repeat(postgres_engine: Engine, model_id: int, last_prediction: tuple def make_workflow_id(prefix: str) -> str: """Build a unique workflow id using a prefix and current timestamp.""" return f'{prefix}-{datetime.now().timestamp()}' + + +def insert_target_data_for_drift( + postgres_engine: Engine, + model_id: int, + timestamps: list[str], + variables_values: dict[str, list[float]], +) -> None: + """ + Insert one row per (timestamp, variable) pair into ``laborious_data``. + + Used by drift scenarios that need wide-format input where the pivot keeps a + full row for every timestamp. + + Args: + - postgres_engine: SQLAlchemy engine bound to the test container. + - model_id: Model id stamped on every row. + - timestamps: ISO-8601 strings used both as ``timestamp`` and ``created_at``. + - variables_values: Mapping of variable name to a list of values; each list + must be the same length as ``timestamps``. + """ + for var_name, values in variables_values.items(): + if len(values) != len(timestamps): + raise ValueError( + f"Variable '{var_name}' has {len(values)} values but {len(timestamps)} timestamps" + ) + + rows_sql = [] + for index, ts in enumerate(timestamps): + for var_name, values in variables_values.items(): + rows_sql.append( + f"({model_id}, '{var_name}', {values[index]}, '{ts}', '{ts}')" + ) + + with postgres_engine.begin() as conn: + conn.execute( + text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}') + ) + if rows_sql: + conn.execute( + text( + 'INSERT INTO predictions_schema.laborious_data ' + '(model_id, variable, value, "timestamp", created_at) VALUES ' + + ', '.join(rows_sql) + ) + ) + + +def build_drift_dataframe( + timestamps: list[str], + features: list[str], + methods: list[str], + statistic: float = 1.0, + drift_flags: dict[tuple[str, str], bool] | None = None, + include_multivariate: bool = True, + multivariate_value: float = 16.0, + multivariate_drift: bool = True, + extra_rows: list[dict[str, Any]] | None = None, +) -> Any: + """ + Build a deterministic dataframe matching the schema returned by + ``sientia_model.analytics.model_analysis.ModelAnalysis.get_drift_metrics_dataframe`` + so drift E2E tests can pin the exact rows persisted to PostgreSQL. + + The output mirrors the analyzer's canonical schema: + ``timestamp, feature, metric, statistic, p_value, alert, chunk_index, + chunk_start_date, chunk_end_date``. ``calculate_drift`` then renames + ``alert -> drift``, ``chunk_index -> chunk``, ``chunk_end_date -> + timestamp_end`` and drops ``p_value`` / ``chunk_start_date``. + + Args: + - timestamps: Truncated chunk start timestamps (``'2024-01-01 12:00'`` for ``min``). + - features: Univariate feature names (one row per feature/method/timestamp). + - methods: Univariate methods such as ``kolmogorov_smirnov``. + - statistic: Default univariate statistic value. + - drift_flags: Optional override of the ``alert`` flag per ``(feature, method)`` pair. + - include_multivariate: Whether to add a final multivariate row block. + - multivariate_value: Value placed on multivariate rows. + - multivariate_drift: Drift flag placed on multivariate rows. + - extra_rows: Additional pre-built rows to append (used for dedup/p_value tests). + + Return: + pandas.DataFrame with columns: timestamp, feature, metric, statistic, + p_value, alert, chunk_index, chunk_start_date, chunk_end_date. + """ + rows: list[dict[str, Any]] = [] + drift_flags = drift_flags or {} + # Synthetic chunk-end offset that mirrors the high-precision boundary + # (``...:59.999999999``) emitted by ``ModelAnalysis`` for minute chunks. + # Computing via ``Timedelta`` instead of string concatenation keeps the + # helper safe for both minute- and second-precision timestamps. + chunk_span = pd.Timedelta(seconds=59, nanoseconds=999999999) + + for chunk_index, ts in enumerate(timestamps): + chunk_start = pd.Timestamp(ts) + chunk_end = chunk_start + chunk_span + for feature in features: + for method in methods: + rows.append( + { + 'timestamp': chunk_start, + 'feature': feature, + 'metric': method, + 'statistic': statistic, + 'p_value': 0.5, + 'alert': drift_flags.get((feature, method), False), + 'chunk_index': chunk_index, + 'chunk_start_date': chunk_start, + 'chunk_end_date': chunk_end, + } + ) + if include_multivariate: + rows.append( + { + 'timestamp': chunk_start, + 'feature': 'multivariate', + 'metric': 'multivariate', + 'statistic': multivariate_value, + 'p_value': 0.0, + 'alert': multivariate_drift, + 'chunk_index': chunk_index, + 'chunk_start_date': chunk_start, + 'chunk_end_date': chunk_end, + } + ) + + if extra_rows: + rows.extend(extra_rows) + + return pd.DataFrame(rows) diff --git a/e2e/scenario_inputs/drift_base.json b/e2e/scenario_inputs/drift_base.json new file mode 100644 index 0000000..1d02c5f --- /dev/null +++ b/e2e/scenario_inputs/drift_base.json @@ -0,0 +1,14 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "schema": "predictions_schema", + "source_table_name": "laborious_data", + "target_table_name": "drift", + "interval": 60, + "drift_metrics": ["kolmogorov_smirnov", "jensen_shannon", "wasserstein"], + "chunk_period": "min", + "model_config": { + "target": "sensor_1" + } +} diff --git a/e2e/scenario_inputs/format_export_base.json b/e2e/scenario_inputs/format_export_base.json new file mode 100644 index 0000000..151342e --- /dev/null +++ b/e2e/scenario_inputs/format_export_base.json @@ -0,0 +1,37 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "EMPTY_DATA": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": true, + "prediction_store_policy": "lts:1", + "model_config": { + "retention_minutes": 0, + "target": "sensor_1" + }, + "datetime_columns": ["timestamp", "created_at"] +} diff --git a/e2e/scenario_inputs/main_happy_path.json b/e2e/scenario_inputs/main_happy_path.json new file mode 100644 index 0000000..e951311 --- /dev/null +++ b/e2e/scenario_inputs/main_happy_path.json @@ -0,0 +1,45 @@ +{ + "metadata": { + "metadata": { + "model_id": "{{MODEL_ID}}", + "model_name": "test_model", + "schedule_name": "test-schedule", + "workflow_name": "predictions_batch" + } + }, + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "EMPTY_DATA": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": true, + "prediction_store_policy": "lts:1", + "model_config": { + "target": "sensor_1", + "retention_minutes": 0 + }, + "datetime_columns": ["timestamp", "created_at"] +} diff --git a/e2e/scenario_inputs/main_invalid_datetime.json b/e2e/scenario_inputs/main_invalid_datetime.json new file mode 100644 index 0000000..1a3c40b --- /dev/null +++ b/e2e/scenario_inputs/main_invalid_datetime.json @@ -0,0 +1,45 @@ +{ + "metadata": { + "metadata": { + "model_id": "{{MODEL_ID}}", + "model_name": "test_model", + "schedule_name": "test-schedule", + "workflow_name": "predictions_batch" + } + }, + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "EMPTY_DATA": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": true, + "prediction_store_policy": "lts:1", + "model_config": { + "target": "sensor_1", + "retention_minutes": 0 + }, + "datetime_columns": ["nonexistent_column"] +} diff --git a/e2e/scenario_inputs/main_missing_required.json b/e2e/scenario_inputs/main_missing_required.json new file mode 100644 index 0000000..e491755 --- /dev/null +++ b/e2e/scenario_inputs/main_missing_required.json @@ -0,0 +1,16 @@ +{ + "metadata": { + "metadata": { + "model_id": "{{MODEL_ID}}", + "model_name": "test_model", + "schedule_name": "test-schedule", + "workflow_name": "predictions_batch" + } + }, + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data" +} diff --git a/e2e/scenario_inputs/main_sql_error.json b/e2e/scenario_inputs/main_sql_error.json new file mode 100644 index 0000000..ffd6436 --- /dev/null +++ b/e2e/scenario_inputs/main_sql_error.json @@ -0,0 +1,44 @@ +{ + "metadata": { + "metadata": { + "model_id": "{{MODEL_ID}}", + "model_name": "test_model", + "schedule_name": "test-schedule", + "workflow_name": "predictions_batch" + } + }, + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT * FROM nonexistent_table WHERE invalid_syntax =", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "EMPTY_DATA": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": true, + "prediction_store_policy": "lts:1", + "model_config": { + "target": "sensor_1", + "retention_minutes": 0 + } +} diff --git a/e2e/scenario_inputs/minimal_retrain_base.json b/e2e/scenario_inputs/minimal_retrain_base.json new file mode 100644 index 0000000..98cd94c --- /dev/null +++ b/e2e/scenario_inputs/minimal_retrain_base.json @@ -0,0 +1,12 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "retrain_reports", + "datetime_columns": ["timestamp", "created_at"], + "model_config": { + "target": "sensor_1" + } +} diff --git a/e2e/scenario_inputs/minio_offload_load_query.json b/e2e/scenario_inputs/minio_offload_load_query.json new file mode 100644 index 0000000..3d4c09c --- /dev/null +++ b/e2e/scenario_inputs/minio_offload_load_query.json @@ -0,0 +1,11 @@ +{ + "metadata": { + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "workflow_name": "predictions_batch" + }, + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "model_name": "test_model", + "datetime_columns": ["timestamp", "created_at"] +} diff --git a/e2e/scenario_inputs/minio_offload_workflow.json b/e2e/scenario_inputs/minio_offload_workflow.json new file mode 100644 index 0000000..4e7b82a --- /dev/null +++ b/e2e/scenario_inputs/minio_offload_workflow.json @@ -0,0 +1,37 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "EMPTY_DATA": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": false, + "prediction_store_policy": "lts:1", + "model_config": { + "retention_minutes": 0, + "target": "sensor_1" + }, + "datetime_columns": ["timestamp", "created_at"] +} diff --git a/e2e/scenario_inputs/prediction_process_base.json b/e2e/scenario_inputs/prediction_process_base.json new file mode 100644 index 0000000..ff46bb9 --- /dev/null +++ b/e2e/scenario_inputs/prediction_process_base.json @@ -0,0 +1,39 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "query": "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {{MODEL_ID}}", + "schema": "predictions_schema", + "table_name": "predictions", + "transform_table_name": "transformed_data", + "input_filters": { + "SPECIFIC_VARIABLES_NULL_VALUES": { + "POLICY": "CONTINUE", + "CONFIG": { + "variables": ["sensor_1"] + } + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "STOP", + "CONFIG": {} + } + }, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "opc_output_config": {}, + "pi_web_api_output_config": {}, + "save_transform": true, + "prediction_store_policy": "lts:1", + "model_config": { + "retention_minutes": 0, + "target": "sensor_1" + }, + "datetime_columns": ["timestamp", "created_at"] +} diff --git a/e2e/scenario_inputs/simple_metrics_base.json b/e2e/scenario_inputs/simple_metrics_base.json new file mode 100644 index 0000000..7ba84ed --- /dev/null +++ b/e2e/scenario_inputs/simple_metrics_base.json @@ -0,0 +1,14 @@ +{ + "schedule_name": "test-schedule", + "model_name": "test_model", + "model_id": "{{MODEL_ID}}", + "schema": "predictions_schema", + "predictions_table_name": "predictions", + "data_table_name": "laborious_data", + "target_table_name": "simple_metrics_data", + "interval_minutes": 60, + "metrics": ["rmse", "mse", "mae", "r2"], + "model_config": { + "target": "sensor_target" + } +} diff --git a/e2e/scenarios.md b/e2e/scenarios.md index 5732cd3..da6aa4f 100644 --- a/e2e/scenarios.md +++ b/e2e/scenarios.md @@ -1,520 +1,402 @@ -# Test Scenarios for Predictions Batch Workflow +# E2E Scenario Documentation - Predictions Batch -This document describes all possible test scenarios for the `predictions_batch` workflow and its child workflows `prediction_process` and `format_and_export_prediction`. +This document describes the end-to-end scenarios for `predictions_batch` and its child workflows: +`prediction_process` and `format_and_export_prediction`. -## Running automated E2E tests (`e2e/`) +It is a functional reference of scenario behavior, inputs, and expected outcomes. -- **Runtime**: Docker (or a Docker-compatible daemon) must be available so [testcontainers](https://testcontainers.com/) can start **PostgreSQL** and **MinIO** containers. -- **Dependencies**: install dev requirements (includes `testcontainers[postgres,minio]`). -- **Invocation**: run only integration-marked tests, for example: `pytest e2e/ -m integration`. -- **MinIO tests**: `e2e/test_minio_offload.py` exercises real S3 uploads; other E2E modules continue to mock MinIO on the worker used by most scenarios. +## Execution Context -## Workflow Overview - -The `predictions_batch` workflow: -1. Loads data using a custom SQL query -2. Prepares prediction configuration -3. Delegates to `prediction_process` child workflow which: - - Retrieves last timestamp for incremental processing - - Applies input data quality gates - - Executes MLFlow transform operation - - Validates transform response - - Executes MLFlow predict operation - - Validates predict response - - Delegates to `format_and_export_prediction` child workflow -4. The `format_and_export_prediction` workflow: - - Formats prediction data (normal or default) - - Exports to PI Web API (optional) - - Exports to OPC server (optional) - - Exports to PostgreSQL - - Writes metrics +- Tests run under `e2e/` and are marked with `@pytest.mark.integration`. +- PostgreSQL and MinIO are provisioned with testcontainers. +- `test_minio_offload.py` uses real MinIO I/O; other scenario suites may use stubs/mocks for optional outputs. --- -## 1. Predictions Batch - Main Workflow Scenarios +## 1. Main Workflow Scenarios +Source: `e2e/test_predictions_batch_main_workflow.py` -### 1.1 Success Scenarios +### 1.1.1 Happy Path - Complete Success +**Summary**: Full workflow succeeds with valid query and default gate behavior. -#### Scenario 1.1.1: Happy Path - Complete Success -**Description**: Workflow completes successfully with valid SQL query and all activities succeed +**Description**: +- Query returns rows for a model. +- `prediction_process` runs transform and predict paths. +- Final prediction and transformed data are persisted. -**Input**: -- Valid `schedule_name`, `model_name`, `model_id` -- Valid `query` returning non-empty DataFrame -- Valid `schema`, `table_name`, `transform_table_name` -- Optional `datetime_columns` for timestamp parsing -- Optional `input_filters`, `mlflow_transform_filters`, `mlflow_predict_filters` -- Optional `path_priority`, `opc_output_config`, `pi_web_api_output_config` +**Expected Outcome**: +- Exactly one prediction row is created. +- Transform rows are created. +- Confidence/status/comments are success values. -**Expected Behavior**: -- `load_custom_query` returns DataFrame with data -- Workflow prepares prediction input with all configurations -- `prediction_process` child workflow executes successfully -- All gates pass with no issues -- Transform and predict operations succeed -- Data exported to PostgreSQL -- Metrics written +### 1.2.1 SQL Query Execution Error +**Summary**: Invalid SQL leads to no persisted prediction. -**Assertions**: -- SQL query executed once -- `prediction_process` workflow called with correct parameters -- Data exists in PostgreSQL (predictions table) -- Metrics recorded -- No errors raised +**Description**: +- Input query is invalid. +- Load step fails and workflow follows error/short-circuit path. + +**Expected Outcome**: +- No prediction rows for the model. +- Workflow does not require retry-loop assumptions in assertions. + +### 1.2.2 Missing Required Parameters +**Summary**: Missing required fields prevent workflow completion path. + +**Description**: +- Required input key (e.g. `query`) is omitted. +- Workflow fails to produce actionable input for child flow. + +**Expected Outcome**: +- No prediction rows are persisted. +- Workflow handle may require explicit terminate in E2E harness. + +### 1.2.3 Invalid Datetime Column Specification (de-prioritized) +**Summary**: Legacy invalid datetime-column case is retained only as low-priority legacy coverage. + +**Description**: +- `datetime_columns` references non-existing columns. +- Behavior may vary by query shape and parser fallback. + +**Expected Outcome**: +- No predictions persisted in the covered legacy assertion path. +- Scenario is not considered primary behavior coverage. --- -### 1.2 Error Scenarios +## 2. Prediction Process Scenarios +Source: `e2e/test_predictions_batch_prediction_process.py` -#### Scenario 1.2.1: SQL Query Execution Error -**Description**: SQL query fails due to syntax error or connection issue +### 2.1 Input Gate Path Decisions -**Input**: -- Invalid SQL query (syntax error) -- Or database connection unavailable +#### 2.1.1 CONTINUE +**Summary**: Input filter flags quality issue but allows continuation via default path. -**Expected Behavior**: -- `load_custom_query` raises exception (caught by Temporal retry policy) -- Notification sent with SQL error details -- After retries, activity may return empty data or workflow may fail -- If empty data returned, workflow completes with early exit via input gate +**Description**: +- Input gate returns `CONTINUE`. +- MLFlow transform/predict are skipped. +- Export path persists default-style prediction with warning context. -**Assertions**: -- Error notification sent -- Workflow completes (either fails or exits early) -- No data in predictions table +#### 2.1.2 STOP +**Summary**: Input filter blocks processing. + +**Description**: +- Input gate returns `STOP`. +- Workflow exits without export. + +#### 2.1.3 REPEAT with history +**Summary**: Prior prediction is reused. + +**Description**: +- Input gate returns `REPEAT`. +- `repeat_last_prediction` path is executed using existing historical row. + +#### 2.1.4 REPEAT without history +**Summary**: Repeat requested but no previous prediction exists. + +**Description**: +- Input gate returns `REPEAT`. +- No prior row is available to duplicate. + +**Expected Outcome**: +- No new prediction rows are created for the model. + +### 2.2 Transform Gate Decisions + +#### 2.2.1 CONTINUE on transform response error +**Summary**: Transform response is degraded, but workflow continues. + +#### 2.2.2 STOP on transform response error +**Summary**: Transform response error blocks downstream processing. + +#### 2.2.3 REPEAT on transform response error +**Summary**: Transform response error triggers repeat-last-prediction path. + +#### 2.2.4 STOP on transform content NaN +**Summary**: Content gate (`NAN_VALUES`) blocks on all-NaN transform payload. + +### 2.3 Predict Gate Decisions + +#### 2.3.1 CONTINUE on predict response error +**Summary**: Predict response degraded; workflow exports with degraded metadata. + +#### 2.3.2 STOP on predict response error +**Summary**: Predict response error blocks export. + +#### 2.3.3 REPEAT on predict response error +**Summary**: Predict response error routes to repeat-last-prediction. + +### 2.4.1 Priority Conflict Resolution +**Summary**: Deterministic selection when multiple filters produce different flags. + +**Description**: +- Multiple filters may produce `STOP`, `CONTINUE`, and/or `REPEAT`. +- `path_priority` defines precedence. + +**Expected Outcome**: +- Highest-priority flag is applied consistently. +- Executed branch matches configured priority ordering. --- -#### Scenario 1.2.2: Missing Required Parameters -**Description**: Essential parameters missing from input +## 3. Format and Export Scenarios +Source: `e2e/test_predictions_batch_format_export.py` -**Input**: -- Missing `query` or `model_id` or `schema` or `table_name` +### 3.1 Output Combination Scenarios -**Expected Behavior**: -- Workflow or activity raises KeyError or validation error -- Workflow fails immediately +#### 3.1.1 Default prediction export +**Summary**: Non-`None` path flag uses `format_default_prediction`. -**Assertions**: -- Workflow fails with parameter error -- Error notification sent -- No child workflow called +**Description**: +- Default prediction is generated. +- Transform export is skipped. +- Optional outputs (PI/OPC) still execute when configured. + +#### 3.1.2 OPC only +**Summary**: Postgres + OPC writes, PI Web API disabled. + +#### 3.1.3 PI Web API only +**Summary**: Postgres + PI writes, OPC disabled. + +#### 3.1.4 Postgres only +**Summary**: Both optional outputs disabled; only Postgres persistence and metrics. + +#### 3.1.5 No transformed data export +**Summary**: Prediction is persisted; transformed table is not written. + +### 3.2 Degraded-but-successful Completion + +#### 3.2.1 PI Web API write error +**Summary**: PI write failure does not fail workflow. + +**Expected Outcome**: +- Workflow completes. +- Prediction persisted with degraded confidence/comments (PI error semantics). + +#### 3.2.2 OPC write error +**Summary**: OPC write failure does not fail workflow. + +**Expected Outcome**: +- Workflow completes. +- Prediction persisted with OPC degraded confidence/comments. + +#### 3.2.3 PI Web API partial write error +**Summary**: Partial PI acknowledgement is treated as degraded success. + +**Expected Outcome**: +- Workflow completes. +- Prediction persisted with PI error confidence and descriptive comment. + +### 3.3.1 Combined Optional Outputs (PI + OPC) +**Summary**: Both external output channels are enabled together. + +**Description**: +- PI Web API and OPC configs are both present. +- Output mutation order matters for final persisted payload. + +**Expected Outcome**: +- PI write executes before OPC write in workflow sequence. +- Final Postgres payload reflects any confidence/comment updates. +- OPC metrics are emitted when tag writes return response times. --- -#### Scenario 1.2.3: Invalid Datetime Column Specification -**Description**: Datetime column specified doesn't exist in query results +## 4. MinIO Offload Scenarios +Source: `e2e/test_minio_offload.py` -**Input**: -- `datetime_columns: ['nonexistent_column']` -- Query results don't have this column +### 4.1.1 Forced offload to MinIO +**Summary**: Very low threshold forces parquet upload. -**Expected Behavior**: -- `load_custom_query` may raise KeyError or warning -- Depending on implementation, workflow may fail or continue -- Error notification sent +**Description**: +- Payload is offloaded (`object_key` present, inline data absent/empty). +- Object is present in MinIO under `prediction_datasets/...`. +- Retrieval reconstructs the dataframe. -**Assertions**: -- Error raised or warning logged -- Workflow behavior depends on error handling policy +### 4.1.2 Full workflow with offloaded load payload +**Summary**: Offload path works during full `predictions_batch` execution. + +**Expected Outcome**: +- Workflow completes. +- Prediction row is persisted. + +### 4.2.1 Inline payload below threshold +**Summary**: Data remains inline when threshold is not exceeded. + +**Expected Outcome**: +- Payload stores inline `data`. +- `object_key` is `None`. +- Downstream persistence behavior matches offload scenario semantics. --- -## 2. Prediction Process - Child Workflow Scenarios +## 5. Drift Workflow Scenarios +Source: `e2e/test_drift.py` -### 2.1 Input gate Early Exit Scenarios +The drift suite mocks `sientia.ModelAnalysis.ModelAnalysis` (not installed; see +`CODE_ISSUES.md` issue #1) through the controllable `_FakeModelAnalysis` stub +exposed by the `model_analysis_stub` fixture. The `mlflow_repository_stub` +provides the reference-data CSV via `download_artifacts`. Every scenario asserts +postgres rows in `predictions_schema.drift` against this canonical schema: -#### Scenario 2.1.1: Input Gate Triggers CONTINUE -**Description**: Input gate determines data should use previous prediction +`id, model_id, feature, method, value, drift, chunk, timestamp, timestamp_end, accurate, created_at, updated_at`. -**Input**: -- Data that should continue with input data as prediction -- `input_filters` configured with `POLICY: 'CONTINUE'` -- `path_priority` includes CONTINUE +### 5.1 Happy paths -**Expected Behavior**: -- `input_gate` returns `path_flag='CONTINUE'` -- `path_flag_handler` calls export workflow with input data directly -- MLFlow transform and predict skipped -- Data exported as-is +#### D.1.1 Full pipeline persists all columns with reference data +**Summary**: ModelAnalysis returns a deterministic drift dataframe; the +reference CSV is downloaded from the MLflow stub. -**Assertions**: -- `input_gate` called -- MLFlow operations NOT called -- Export workflow called with original data -- Workflow completes +**Expected Outcome**: +- One row per `(chunk, feature, method)` plus a `multivariate` block per chunk. +- Every drift column is populated and `accurate=True`. +- `timestamp_end` preserves the high-precision string (`HH:MM:59.999999999`). +- `p_value` is dropped before persistence. +- `drift` flags propagate per `(feature, method)` configuration. +#### D.1.2 30% fallback when reference data is unavailable +**Summary**: `get_reference_data` fails alias resolution and returns `None`; +`calculate_drift` uses the first 30% of target rows as reference. -#### Scenario 2.1.2: Input Gate Triggers STOP -**Description**: Input data quality gate fails with STOP policy +**Expected Outcome**: +- Persisted rows carry `accurate=False`. +- A `MODEL_METRICS_REFERENCE_DATA_WARNING` notification is emitted to MongoDB. -**Input**: -- Data with EMPTY_DATA or other critical issues -- `input_filters` configured with `POLICY: 'STOP'` +### 5.2 Filtering / dedup invariants -**Expected Behavior**: -- `input_gate` returns `path_flag='STOP'` -- `path_flag_handler` detects STOP -- Workflow returns early without calling MLFlow -- No prediction exported +#### D.2.1 Deduplication and `p_value` removal +**Summary**: ModelAnalysis returns duplicate `(timestamp, method, feature)` rows +plus a `p_value` column. -**Assertions**: -- `input_gate` called -- `path_flag_handler` returns True (early exit) -- MLFlow transform NOT called -- Export workflow NOT called -- Workflow completes without error +**Expected Outcome**: +- Duplicates are collapsed keeping the first occurrence. +- `p_value` is absent from the persisted rows. +#### D.2.2 Out-of-range timestamps filtered +**Summary**: Drift rows whose timestamps are not present in the target window +must be discarded before persistence. -#### Scenario 2.1.3: Input Gate Triggers REPEAT -**Description**: Input gate determines data should repeat last prediction +### 5.3 Failure paths -**Input**: -- Data with quality issues that require using previous prediction -- `input_filters` configured with `POLICY: 'REPEAT'` -- `path_priority` includes REPEAT +#### D.3.1 Empty target data short-circuits the workflow +**Summary**: `load_custom_query` returns no rows; ModelAnalysis is never +instantiated and no drift rows are written. -**Expected Behavior**: -- `input_gate` returns `path_flag='REPEAT'` -- `path_flag_handler` calls `repeat_last_prediction` activity -- MLFlow transform and predict skipped -- Last prediction repeated and exported +#### D.3.2 ModelAnalysis raises during dataframe assembly +**Summary**: `get_drift_metrics_dataframe` raises. The activity catches the +error, sends a `MODEL_METRICS_GET_DRIFT_METRICS_ERROR` notification, and the +workflow completes without persisting drift rows. -**Assertions**: -- `input_gate` called -- MLFlow operations NOT called -- `repeat_last_prediction` activity called -- Workflow completes +### 5.4 Configuration paths + +#### D.4.1 Default drift metrics propagated to analyzer +**Summary**: Omitting `drift_metrics` defaults to +`['kolmogorov_smirnov', 'jensen_shannon', 'wasserstein']` and forwards the +exact list to `detect_univariate_drift`. + +#### D.4.2 Invalid `chunk_period` raises ValueError +**Summary**: Anything other than `min` / `s` is rejected by `calculate_drift`. + +#### D.4.3 `chunk_period='s'` keeps seconds in timestamp filtering +**Summary**: Truncated `YYYY-MM-DD HH:MM` rows are filtered out when chunking +runs at second granularity. --- -### 2.2 Transform gate Early Exit Scenarios +## 6. Simple Metrics Workflow Scenarios +Source: `e2e/test_simple_metrics.py` -#### Scenario 2.2.1: Transform Gate Triggers CONTINUE -**Description**: Transform response gate determines data should continue despite issues +Validates `predictions_schema.simple_metrics_data` columns: +`id, model_id, metric, value, timestamp, data_size, interval_minutes, created_at`. -**Input**: -- Valid input data -- Transform response has quality issues but policy is CONTINUE -- `mlflow_transform_filters` configured with `POLICY: 'CONTINUE'` -- `path_priority` includes CONTINUE +### 6.1 Happy paths -**Expected Behavior**: -- `request_transform` succeeds -- `mlflow_response_gate` for transform returns `path_flag='CONTINUE'` -- `path_flag_handler` calls export workflow with transform data -- MLFlow predict skipped -- Transform data exported as-is +#### S.1.1 rmse/mse/mae/r2 happy path +**Summary**: Prediction/target pairs are inserted; the activity computes all +four metrics with closed-form expected values. -**Assertions**: -- Transform completed -- `mlflow_response_gate` called for transform -- MLFlow predict NOT called -- Export workflow called with transform data -- Workflow completes +**Expected Outcome**: +- One row per metric is persisted; all columns populated. +- `data_size` matches the joined row count and `interval_minutes=60`. + +#### S.1.2 Subset metrics +**Summary**: Requesting `metrics=['rmse']` writes only the rmse row. + +### 6.2 Edge cases + +#### S.2.1 Zero-variance target returns r2=0 +**Summary**: When all targets are equal, `ss_tot=0`; the activity must guard +against division by zero and return `r2=0`. + +### 6.3 Failure paths + +#### S.3.1 No overlapping data short-circuits persistence +**Summary**: With no `laborious_data` rows for the configured target variable +the workflow exits before `calculate_simple_metrics` and writes nothing. --- -#### Scenario 2.2.2: Transform Gate Triggers STOP -**Description**: Transform response validation fails with STOP policy +## 7. Minimal Retrain Workflow Scenarios +Source: `e2e/test_minimal_retrain.py` -**Input**: -- Valid input data -- Transform response has critical errors -- `mlflow_transform_filters` configured with `POLICY: 'STOP'` +The MLflow registry is fully mocked (no real artifacts in test container). +Validates `predictions_schema.retrain_reports` columns: +`id, model_id, model_name, timestamp, status, version, mlflow_run_id, mlflow_experiment_id, created_at`. -**Expected Behavior**: -- `request_transform` succeeds but response invalid -- `mlflow_response_gate` for transform returns `path_flag='STOP'` -- Workflow exits without calling predict or export +### 7.1 Happy path -**Assertions**: -- Transform completed but validation failed -- `mlflow_response_gate` called for transform -- MLFlow predict NOT called -- Export workflow NOT called -- Workflow completes without error +#### MR.1.1 Successful retrain + promotion +**Summary**: Training data loads via MinIO offload, `wrapper.retrain` succeeds, +the new version is promoted to the `production` alias. + +**Expected Outcome**: +- Report row has success status, `version='7'`, `mlflow_run_id='retrain-run-id'`, + `mlflow_experiment_id='experiment-id'`. +- `mlflow.log_artifact` is called with the input CSV. +- `promote_to_alias` is called once with the resolved version and alias. + +### 7.2 Failure paths + +#### MR.2.1 Wrapper retrain raises +**Summary**: `wrapper.retrain` raises `RuntimeError`. The activity returns +`success=False`, `update_production_model` is NOT invoked. + +**Expected Outcome**: +- Report row carries the error message and `version`/`mlflow_*` columns are NULL. + +#### MR.2.2 Missing `model_config.target` +**Summary**: Empty model config short-circuits before any MLflow call. + +**Expected Outcome**: +- Report row carries the explicit guard message. +- `get_cached_model` is never invoked. + +#### MR.3.1 No training data +**Summary**: The training query returns no rows; the workflow does not +persist any report row. The current code raises plain `ValueError` from the +workflow function, which Temporal treats as a workflow-task failure (see +`CODE_ISSUES.md` issue MR-1). --- -#### Scenario 2.2.3: Transform Gate Triggers REPEAT -**Description**: Transform response gate determines data should repeat last prediction +## Input Contract Reference -**Input**: -- Valid input data -- Transform response has quality issues that require using previous prediction -- `mlflow_transform_filters` configured with `POLICY: 'REPEAT'` -- `path_priority` includes REPEAT +Common scenario input fields: +- `schedule_name` +- `model_name` +- `model_id` +- `query` +- `schema` +- `table_name` +- `transform_table_name` +- `input_filters` +- `mlflow_transform_filters` +- `mlflow_predict_filters` +- `path_priority` (default order: `STOP`, `CONTINUE`, `REPEAT`) +- `save_transform` +- `prediction_store_policy` +- `model_config.target` +- `datetime_columns` (when query returns temporal fields) -**Expected Behavior**: -- `request_transform` succeeds but response has issues -- `mlflow_response_gate` for transform returns `path_flag='REPEAT'` -- `path_flag_handler` calls `repeat_last_prediction` activity -- MLFlow predict skipped -- Last prediction repeated and exported - -**Assertions**: -- Transform completed but validation triggered REPEAT -- `mlflow_response_gate` called for transform -- MLFlow predict NOT called -- `repeat_last_prediction` activity called -- Workflow completes - ---- - -### 2.3 Predict gate Early Exit Scenarios - -#### Scenario 2.3.1: Predict Gate Triggers CONTINUE -**Description**: Predict response gate determines data should continue despite issues - -**Input**: -- Valid input and transform data -- Predict response has quality issues but policy is CONTINUE -- `mlflow_predict_filters` configured with `POLICY: 'CONTINUE'` -- `path_priority` includes CONTINUE - -**Expected Behavior**: -- `request_predict` succeeds -- `mlflow_response_gate` for predict returns `path_flag='CONTINUE'` -- `path_flag_handler` calls export workflow with predict data -- Prediction exported despite quality issues - -**Assertions**: -- Transform and predict completed -- `mlflow_response_gate` called for predict -- Export workflow called with predict data -- Workflow completes - ---- - -#### Scenario 2.3.2: Predict Gate Triggers STOP -**Description**: Prediction validation fails with STOP policy - -**Input**: -- Valid input and transform -- Predict response has critical errors -- `mlflow_predict_filters` configured with `POLICY: 'STOP'` - -**Expected Behavior**: -- `request_predict` succeeds but response invalid -- `mlflow_response_gate` for predict returns `path_flag='STOP'` -- Workflow exits without export - -**Assertions**: -- Transform completed -- Predict completed but validation failed -- Export workflow NOT called -- Workflow completes without error - ---- - -#### Scenario 2.3.3: Predict Gate Triggers REPEAT -**Description**: Predict response gate determines data should repeat last prediction - -**Input**: -- Valid input and transform data -- Predict response has quality issues that require using previous prediction -- `mlflow_predict_filters` configured with `POLICY: 'REPEAT'` -- `path_priority` includes REPEAT - -**Expected Behavior**: -- `request_predict` succeeds but response has issues -- `mlflow_response_gate` for predict returns `path_flag='REPEAT'` -- `path_flag_handler` calls `repeat_last_prediction` activity -- Last prediction repeated and exported - -**Assertions**: -- Transform and predict completed but validation triggered REPEAT -- `mlflow_response_gate` called for predict -- `repeat_last_prediction` activity called -- Export workflow NOT called with current prediction -- Workflow completes - ---- - -## 3. Format and Export Prediction - Child Workflow Scenarios - -### 3.1 Success Scenarios - -#### Scenario 3.1.1: Default Prediction Export -**Description**: Error prediction path creates default prediction - -**Input**: -- `path_flag: 'ERROR'` or other non-None value (not STOP/CONTINUE/REPEAT) -- `comment` provided with error details - -**Expected Behavior**: -- `format_default_prediction` called instead of `format_prediction` -- Default prediction created with error metadata -- Exported to PostgreSQL only -- Transformed data NOT processed -- Metrics written - -**Assertions**: -- `format_default_prediction` called -- `format_prediction` NOT called -- `format_transformed_data` NOT called -- One PostgreSQL export only -- Default values in prediction data -- Comment included - ---- - -#### Scenario 3.1.2: Export with OPC only -**Description**: Export to PostgreSQL and OPC server only (no PI Web API) - -**Input**: -- `path_flag: None` -- `opc_output_config` configured with valid OPC settings -- `pi_web_api_output_config: None` or `{}` - -**Expected Behavior**: -- Normal formatting -- PostgreSQL export executed -- OPC export executed -- PI Web API activity skipped -- Metrics written with OPC metrics - -**Assertions**: -- PI Web API activity NOT called -- OPC activity called -- PostgreSQL export called -- Metrics written with `opc_metrics` populated - ---- - -#### Scenario 3.1.3: Export with PI Web API only -**Description**: Export to PostgreSQL and PI Web API only (no OPC) - -**Input**: -- `path_flag: None` -- `pi_web_api_output_config` configured with valid PI Web API settings -- `opc_output_config: None` or `{}` - -**Expected Behavior**: -- Normal formatting -- PostgreSQL export executed -- PI Web API export executed -- OPC activity skipped -- Metrics written without OPC metrics - -**Assertions**: -- OPC activity NOT called -- PI Web API activity called -- PostgreSQL export called -- Metrics written with empty `opc_metrics` - ---- - -#### Scenario 3.1.4: Export Without Optional Outputs -**Description**: Export only to PostgreSQL (no OPC or PI Web API) - -**Input**: -- `path_flag: None` -- `opc_output_config: None` or `{}` -- `pi_web_api_output_config: None` or `{}` - -**Expected Behavior**: -- Normal formatting -- Only PostgreSQL export executed -- OPC and PI Web API activities skipped -- Metrics written without OPC metrics - -**Assertions**: -- PI Web API activity NOT called -- OPC activity NOT called -- PostgreSQL export called -- Metrics written with empty `opc_metrics` - ---- - -#### Scenario 3.1.5: Export Without Transformed Data -**Description**: Only prediction exported, no transform table - -**Input**: -- `path_flag: None` -- `transformed_data: None` or `save_transform: False` -- `opc_output_config: None` or `{}` -- `pi_web_api_output_config: None` or `{}` - -**Expected Behavior**: -- Only prediction formatted and exported -- Transform export skipped -- Single PostgreSQL write - -**Assertions**: -- `format_transformed_data` NOT called -- One PostgreSQL export -- Transform table remains empty - ---- - -### 3.2 Error Scenarios - -These paths do **not** rely on Temporal activity retries for export failures: the write activities run once, errors are handled inside the activity, and the **workflow completes successfully** with degraded metadata on the persisted prediction (`prediction_confidence` and `comments`). - -#### Scenario 3.2.1: PI Web API Write Error -**Description**: PI Web API export fails - -**Input**: -- Valid prediction -- PI Web API service unavailable or invalid config - -**Expected Behavior**: -- `write_pi_web_api_data` surfaces the failure (exception handled in the activity layer) -- Notification may be sent -- Workflow **completes** (does not fail) -- Prediction row is still written to PostgreSQL with error confidence **13** and a comment describing the PI error -- Subsequent steps (e.g. OPC, Postgres) still run per workflow order with the updated prediction payload - -**Assertions**: -- PI Web API error notification sent (when applicable) -- Workflow completes -- PostgreSQL contains the prediction with `prediction_confidence` 13 and expected `comments` - ---- - -#### Scenario 3.2.2: OPC Write Error -**Description**: OPC server write fails - -**Input**: -- Valid prediction -- OPC server unavailable or invalid configuration - -**Expected Behavior**: -- `write_opc_data` reports failure without aborting the workflow -- Notification may be sent -- Workflow **completes** (does not fail) -- Prediction row is written to PostgreSQL with OPC error confidence **12** and a comment indicating OPC write issues - -**Assertions**: -- OPC error notification sent (when applicable) -- Workflow completes -- PostgreSQL contains the prediction with `prediction_confidence` 12 and expected `comments` - ---- - -#### Scenario 3.2.3: PI Web API Partial Write Error -**Description**: Two prediction tags attempt to be written to PI Web API, but only one succeeds - -**Input**: -- Valid prediction -- Two prediction tags configured -- PI Web API returns partial success (one tag succeeds, one fails) - -**Expected Behavior**: -- `write_pi_web_api_data` processes response -- `process_pi_web_api_response` detects partial failure -- Error confidence set (13) -- Notification sent for failed tag -- Workflow completes with error confidence (single activity attempt; no retry loop) - -**Assertions**: -- One tag written successfully -- One tag failed -- Error confidence set in prediction -- Error notification sent -- Workflow completes - ---- \ No newline at end of file +Optional outputs: +- `opc_output_config` +- `pi_web_api_output_config` \ No newline at end of file diff --git a/e2e/test_drift.py b/e2e/test_drift.py new file mode 100644 index 0000000..80d5664 --- /dev/null +++ b/e2e/test_drift.py @@ -0,0 +1,786 @@ +""" +End-to-end tests for the Drift workflow. + +Coverage focus: + +- Full pipeline persists drift rows with **all** columns expected by the + ``predictions_schema.drift`` table (model_id, feature, method, value, drift, + chunk, timestamp, timestamp_end, accurate, plus DB-managed id/created_at/updated_at). +- ``get_reference_data`` happy path (CSV downloaded from MLflow stub) and + fallback path (30% of target data when reference is unavailable). +- ``calculate_drift`` invariants: ``p_value`` dropped, duplicates removed, rows + outside target timestamps filtered, default ``drift_metrics`` propagated. +- Failure paths: ``ModelAnalysis.get_drift_metrics_dataframe`` raises ⇒ + workflow completes without writes; empty target data ⇒ workflow short-circuits; + invalid ``chunk_period`` ⇒ activity raises and workflow surfaces the error. +""" + +from datetime import datetime, timedelta, timezone +from decimal import Decimal +from pathlib import Path +from unittest.mock import MagicMock + +import pandas as pd +import pytest +from sqlalchemy import text +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Worker + +from e2e.helpers import ( + build_drift_dataframe, + insert_target_data_for_drift, + load_scenario_input, + make_workflow_id, + start_and_await_workflow, +) +from laborious.activities.activities import Activities +from laborious.workflows.drift import Drift + +# Drift columns that must be present (and non-null where required) on every row. +EXPECTED_DRIFT_COLUMNS = [ + 'id', + 'model_id', + 'feature', + 'method', + 'value', + 'drift', + 'chunk', + 'timestamp', + 'timestamp_end', + 'accurate', + 'created_at', + 'updated_at', +] + + +def _drift_input(model_id: int, **overrides) -> dict: + """Load the base drift scenario JSON and apply ad-hoc overrides.""" + input_data = load_scenario_input('drift_base.json', model_id=model_id) + input_data.update(overrides) + return input_data + + +def _five_minute_window(offset_minutes: int = 6) -> tuple[list[str], list[str]]: + """ + Build five consecutive UTC minute timestamps positioned in the recent past. + + The Drift workflow filters target rows with ``timestamp > NOW() - INTERVAL``, + so timestamps must be recent for tests to retrieve any data. We snap to + minute precision and back off ``offset_minutes`` minutes so all chunks land + well inside the default 60-minute interval defined in ``drift_base.json``. + + Args: + - offset_minutes: How many minutes ago the most recent chunk should be. + + Return: + Tuple ``(target_timestamps, chunk_timestamps)``: + + - ``target_timestamps``: ISO strings with ``+0000`` used as ``timestamp`` + and ``created_at`` columns when inserting target rows. + - ``chunk_timestamps``: ``YYYY-MM-DD HH:MM`` truncations matching what + ``calculate_drift`` filters on for ``chunk_period='min'``. + """ + base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta( + minutes=offset_minutes + ) + target_timestamps = [ + (base + timedelta(minutes=i)).strftime('%Y-%m-%d %H:%M:%S%z') for i in range(5) + ] + chunk_timestamps = [ + (base + timedelta(minutes=i)).strftime('%Y-%m-%d %H:%M') for i in range(5) + ] + return target_timestamps, chunk_timestamps + + +def _configure_reference_csv(mlflow_repository_stub, reference_rows: pd.DataFrame) -> None: + """ + Wire ``mlflow_repository_stub`` so ``get_reference_data`` returns + ``reference_rows`` by writing them to ``dst_path/evaluation_data.csv``. + """ + + def _download(run_id: str, artifact_path: str, dst_path: str, metadata=None): + target = Path(dst_path) / 'evaluation_data.csv' + reference_rows.to_csv(target, index=False) + + mlflow_repository_stub._client.get_model_version_by_alias.return_value = MagicMock( + run_id='fake-reference-run' + ) + mlflow_repository_stub.download_artifacts.side_effect = _download + + +def _force_reference_unavailable(mlflow_repository_stub) -> None: + """Make ``get_reference_data`` return ``None`` by failing alias resolution.""" + mlflow_repository_stub._client.get_model_version_by_alias.side_effect = Exception( + 'no production alias registered' + ) + + +def _assert_drift_columns_complete(rows, *, expected_count: int) -> None: + """Validate row count and that the canonical drift columns are populated.""" + assert len(rows) == expected_count, ( + f'Expected {expected_count} drift rows persisted, got {len(rows)}' + ) + seen_columns = set(rows[0]._mapping.keys()) if rows else set() + for column in EXPECTED_DRIFT_COLUMNS: + assert column in seen_columns, f'Missing drift column in postgres: {column}' + + for row in rows: + mapping = dict(row._mapping) + for column in EXPECTED_DRIFT_COLUMNS: + if column == 'value': + # ``value`` is nullable in the table; skip null check, only ensure key exists. + continue + assert mapping[column] is not None, f"Column '{column}' is NULL in {mapping}" + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_happy_path_persists_all_columns_with_reference_data( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.1.1: Happy path with reference data downloaded from MLflow. + + Validates the full drift pipeline produces one row per (chunk, feature, method) + combination plus the multivariate row block, with **all** columns required by + ``predictions_schema.drift`` populated. Uses a deterministic mocked drift + dataframe so the postgres assertions remain stable across runs. + """ + client = temporal_test_env.client + model_id = 411 + + target_timestamps, chunk_timestamps = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [10.0, 11.0, 12.0, 13.0, 14.0], + 'sensor_2': [20.0, 21.0, 22.0, 23.0, 24.0], + }, + ) + + reference_df = pd.DataFrame( + { + 'timestamp': ['2023-12-31 11:00:00+00:00', '2023-12-31 11:01:00+00:00'], + 'sensor_1': [9.5, 9.6], + 'sensor_2': [19.5, 19.6], + } + ) + _configure_reference_csv(mlflow_repository_stub, reference_df) + + features = ['sensor_1', 'sensor_2'] + methods = ['kolmogorov_smirnov', 'jensen_shannon', 'wasserstein'] + drift_flags = { + ('sensor_1', 'wasserstein'): True, + ('sensor_2', 'wasserstein'): True, + ('sensor_2', 'jensen_shannon'): True, + } + + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps, + features=features, + methods=methods, + statistic=0.42, + drift_flags=drift_flags, + multivariate_value=15.5, + multivariate_drift=True, + ) + ) + + input_data = _drift_input(model_id) + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-happy-path') + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT * FROM predictions_schema.drift ' + 'WHERE model_id = :model_id ORDER BY chunk, feature, method' + ), + {'model_id': model_id}, + ) + .mappings() + .all() + ) + + expected_rows = len(chunk_timestamps) * (len(features) * len(methods) + 1) # univariate + multivariate + assert len(rows) == expected_rows + for column in EXPECTED_DRIFT_COLUMNS: + assert column in rows[0], f'Missing drift column in postgres: {column}' + + for row in rows: + for column in EXPECTED_DRIFT_COLUMNS: + if column == 'value': + continue + assert row[column] is not None, f"Column '{column}' is NULL in {dict(row)}" + + multivariate_rows = [r for r in rows if r['feature'] == 'multivariate'] + univariate_rows = [r for r in rows if r['feature'] != 'multivariate'] + + assert len(multivariate_rows) == len(chunk_timestamps) + assert all(r['method'] == 'multivariate' for r in multivariate_rows) + assert all(r['drift'] is True for r in multivariate_rows) + assert all(Decimal(str(r['value'])) == Decimal('15.5') for r in multivariate_rows) + + assert len(univariate_rows) == len(features) * len(methods) * len(chunk_timestamps) + assert {r['method'] for r in univariate_rows} == set(methods) + assert {r['feature'] for r in univariate_rows} == set(features) + + drift_pairs = {(r['feature'], r['method']): r['drift'] for r in univariate_rows} + for (feature, method), expected_drift in drift_flags.items(): + assert drift_pairs[(feature, method)] is expected_drift + + assert all(r['accurate'] is True for r in rows), 'reference path should mark rows as accurate' + + assert all( + r['timestamp_end'].endswith(':59.999999999') and 'T' in r['timestamp_end'] + for r in rows + ), 'timestamp_end should preserve the high-precision ISO string from ModelAnalysis' + + assert all('p_value' not in r for r in rows), 'p_value must be dropped before postgres' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_uses_30pct_fallback_when_reference_unavailable( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, + notification_inserts, +): + """ + Scenario D.1.2: ``get_reference_data`` returns ``None`` (production alias missing), + so ``calculate_drift`` falls back to using the first 30% of target rows as + reference. Persisted rows must report ``accurate=False`` and a warning + notification must be emitted to mongo. + """ + client = temporal_test_env.client + model_id = 412 + + target_timestamps, chunk_timestamps = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [10.0, 10.1, 10.2, 10.3, 10.4], + 'sensor_2': [20.0, 20.1, 20.2, 20.3, 20.4], + }, + ) + + _force_reference_unavailable(mlflow_repository_stub) + + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps, + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + statistic=0.7, + include_multivariate=False, + ) + ) + + input_data = _drift_input(model_id) + + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-fallback') + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text('SELECT * FROM predictions_schema.drift WHERE model_id = :m'), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert len(rows) == len(chunk_timestamps) + assert all(r['accurate'] is False for r in rows), 'fallback path must mark rows as inaccurate' + assert all(r['feature'] == 'sensor_1' for r in rows) + + fallback_warnings = [ + call + for call in notification_inserts.call_args_list + if call.args + and isinstance(call.args[0], dict) + and call.args[0].get('notification_id') == 'MODEL_METRICS_REFERENCE_DATA_WARNING' + ] + assert len(fallback_warnings) >= 1, 'expected reference fallback warning notification' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_drops_p_value_and_dedupes_by_timestamp_method_feature( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.2.1: ModelAnalysis returns duplicate (timestamp, method, feature) + rows and a populated ``p_value`` column. ``calculate_drift`` must deduplicate + keeping the first occurrence and the persisted rows must not contain + ``p_value``. + """ + client = temporal_test_env.client + model_id = 421 + + target_timestamps, chunk_timestamps = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0], + 'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0], + }, + ) + + duplicates = [ + # Re-emit the first chunk for sensor_1/kolmogorov_smirnov with a different value + # so we can prove the dedup keeps the FIRST row. + { + 'timestamp': pd.Timestamp(chunk_timestamps[0]), + 'feature': 'sensor_1', + 'metric': 'kolmogorov_smirnov', + 'statistic': 0.99, + 'p_value': 0.02, + 'alert': True, + 'chunk_index': 0, + 'chunk_start_date': pd.Timestamp(chunk_timestamps[0]), + 'chunk_end_date': pd.Timestamp(f'{chunk_timestamps[0]}:59.999999999'), + } + ] + + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps, + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + statistic=0.5, + include_multivariate=False, + extra_rows=duplicates, + ) + ) + _force_reference_unavailable(mlflow_repository_stub) + + input_data = _drift_input(model_id) + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-dedupe') + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT chunk, feature, method, value FROM predictions_schema.drift ' + 'WHERE model_id = :m ORDER BY chunk' + ), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert len(rows) == len(chunk_timestamps), 'duplicates must be removed before persistence' + + first_chunk_rows = [r for r in rows if r['chunk'] == 0] + assert len(first_chunk_rows) == 1 + assert Decimal(str(first_chunk_rows[0]['value'])) == Decimal('0.5'), ( + 'dedup must keep the FIRST occurrence (statistic=0.5), not the duplicate (statistic=0.99)' + ) + + columns = set(rows[0].keys()) + assert 'p_value' not in columns + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_drops_rows_outside_target_timestamps( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.2.2: Drift rows whose timestamps are not present in the target + data (e.g. came from the reference distribution) must be discarded so the + saved drift only reflects analysis chunks. + """ + client = temporal_test_env.client + model_id = 422 + + target_timestamps, chunk_timestamps = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0], + 'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0], + }, + ) + + out_of_range_timestamp = pd.Timestamp('2099-12-31 23:59') + extra = [ + { + 'timestamp': out_of_range_timestamp, + 'feature': 'sensor_1', + 'metric': 'kolmogorov_smirnov', + 'statistic': 0.5, + 'p_value': 0.0, + 'alert': True, + 'chunk_index': 99, + 'chunk_start_date': out_of_range_timestamp, + 'chunk_end_date': pd.Timestamp('2099-12-31 23:59:59.999999999'), + } + ] + + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps, + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + include_multivariate=False, + extra_rows=extra, + ) + ) + _force_reference_unavailable(mlflow_repository_stub) + + input_data = _drift_input(model_id) + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-tts-filter') + ) + + with postgres_engine.connect() as conn: + chunks = [ + r[0] + for r in conn.execute( + text( + 'SELECT chunk FROM predictions_schema.drift WHERE model_id = :m ORDER BY chunk' + ), + {'m': model_id}, + ).all() + ] + + assert chunks == [0, 1, 2, 3, 4], 'out-of-range timestamps must be filtered out' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_empty_target_data_short_circuits_workflow( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.3.1: When ``load_custom_query`` returns no rows the workflow must + return early without invoking ModelAnalysis or writing any drift rows. + """ + client = temporal_test_env.client + model_id = 431 + + with postgres_engine.begin() as conn: + conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}')) + + _force_reference_unavailable(mlflow_repository_stub) + + drift_calls = {'count': 0} + + def _factory(univariate, multivariate): + drift_calls['count'] += 1 + return build_drift_dataframe( + timestamps=['2024-01-01 12:00'], + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + include_multivariate=False, + ) + + model_analysis_stub.set_drift_dataframe(_factory) + + input_data = _drift_input(model_id) + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-empty-target') + ) + + with postgres_engine.connect() as conn: + count = conn.execute( + text('SELECT COUNT(*) FROM predictions_schema.drift WHERE model_id = :m'), + {'m': model_id}, + ).scalar() + assert count == 0 + assert drift_calls['count'] == 0, 'ModelAnalysis must not be invoked when target data is empty' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_model_analysis_failure_keeps_workflow_alive_no_writes( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, + notification_inserts, +): + """ + Scenario D.3.2: ``ModelAnalysis.get_drift_metrics_dataframe`` raises. The + activity must catch the error, send an error notification and return ``[]`` + so the workflow completes without persisting drift rows. + """ + client = temporal_test_env.client + model_id = 432 + + target_timestamps, _ = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0], + 'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0], + }, + ) + + _force_reference_unavailable(mlflow_repository_stub) + model_analysis_stub.raise_on_get_drift_metrics_dataframe( + RuntimeError('drift analyzer crashed') + ) + + input_data = _drift_input(model_id) + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-analyzer-error') + ) + + with postgres_engine.connect() as conn: + count = conn.execute( + text('SELECT COUNT(*) FROM predictions_schema.drift WHERE model_id = :m'), + {'m': model_id}, + ).scalar() + assert count == 0 + + error_notifications = [ + call + for call in notification_inserts.call_args_list + if call.args + and isinstance(call.args[0], dict) + and call.args[0].get('notification_id') == 'MODEL_METRICS_GET_DRIFT_METRICS_ERROR' + ] + assert len(error_notifications) >= 1 + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_default_drift_metrics_propagated_to_analyzer( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.4.1: When ``drift_metrics`` is omitted from input the workflow + must default to ``['kolmogorov_smirnov', 'jensen_shannon', 'wasserstein']`` + and forward exactly that list to ``ModelAnalysis.detect_univariate_drift``. + """ + client = temporal_test_env.client + model_id = 441 + + target_timestamps, chunk_timestamps = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0], + 'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0], + }, + ) + + _force_reference_unavailable(mlflow_repository_stub) + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps, + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + include_multivariate=False, + ) + ) + + input_data = _drift_input(model_id) + input_data.pop('drift_metrics', None) + + await start_and_await_workflow( + client, Drift.run, input_data, make_workflow_id('test-drift-default-methods') + ) + + instance = model_analysis_stub.last_instance + assert instance is not None, 'ModelAnalysis must have been instantiated' + assert len(instance.detect_univariate_drift_calls) == 1 + forwarded_methods = instance.detect_univariate_drift_calls[0]['methods'] + assert forwarded_methods == ['kolmogorov_smirnov', 'jensen_shannon', 'wasserstein'] + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_invalid_chunk_period_raises_value_error( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, +): + """ + Scenario D.4.2: ``calculate_drift`` validates ``chunk_period`` and rejects + anything other than ``min`` / ``s``. The workflow must surface the + ``ValueError`` to the caller and persist nothing. + """ + client = temporal_test_env.client + model_id = 442 + + target_timestamps, _ = _five_minute_window() + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0], + 'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0], + }, + ) + _force_reference_unavailable(mlflow_repository_stub) + + input_data = _drift_input(model_id, chunk_period='hour') + + with pytest.raises(Exception) as excinfo: + await start_and_await_workflow( + client, + Drift.run, + input_data, + make_workflow_id('test-drift-bad-chunk-period'), + ) + + # Temporal wraps the activity ValueError in WorkflowFailureError; the message + # may live on ``.message`` or ``str(exc)`` depending on the SDK error class. + cause_descriptions = [] + current: BaseException | None = excinfo.value + while current is not None: + cause_descriptions.append( + getattr(current, 'message', None) or str(current) or repr(current) + ) + current = current.__cause__ + assert any('Invalid chunk period' in msg for msg in cause_descriptions), ( + f'Expected ValueError about chunk period in chain, got: {cause_descriptions}' + ) + + with postgres_engine.connect() as conn: + count = conn.execute( + text('SELECT COUNT(*) FROM predictions_schema.drift WHERE model_id = :m'), + {'m': model_id}, + ).scalar() + assert count == 0 + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_drift_chunk_period_seconds_preserves_seconds_in_timestamp_column( + temporal_test_env: WorkflowEnvironment, + temporal_worker_drift: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, + model_analysis_stub, +): + """ + Scenario D.4.3: With ``chunk_period='s'`` the persisted ``timestamp`` + column must preserve second-level precision instead of being flattened to + the start of the minute, and ``chunk_period`` must be propagated to the + analyzer so it actually chunks at second granularity. + """ + client = temporal_test_env.client + model_id = 443 + + base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(minutes=2) + target_timestamps = [ + base.strftime('%Y-%m-%d %H:%M:%S%z'), + (base + timedelta(seconds=30)).strftime('%Y-%m-%d %H:%M:%S%z'), + (base + timedelta(minutes=1)).strftime('%Y-%m-%d %H:%M:%S%z'), + ] + chunk_timestamps_full = [ + base.strftime('%Y-%m-%d %H:%M:%S'), + (base + timedelta(seconds=30)).strftime('%Y-%m-%d %H:%M:%S'), + (base + timedelta(minutes=1)).strftime('%Y-%m-%d %H:%M:%S'), + ] + + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [1.0, 2.0, 3.0], + 'sensor_2': [10.0, 20.0, 30.0], + }, + ) + _force_reference_unavailable(mlflow_repository_stub) + + model_analysis_stub.set_drift_dataframe( + lambda univariate, multivariate: build_drift_dataframe( + timestamps=chunk_timestamps_full, + features=['sensor_1'], + methods=['kolmogorov_smirnov'], + include_multivariate=False, + ) + ) + + input_data = _drift_input(model_id, chunk_period='s') + await start_and_await_workflow( + client, + Drift.run, + input_data, + make_workflow_id('test-drift-chunk-seconds'), + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT chunk, timestamp FROM predictions_schema.drift ' + 'WHERE model_id = :m ORDER BY chunk' + ), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert [r['chunk'] for r in rows] == [0, 1, 2] + seconds_present = {r['timestamp'].second for r in rows} + assert seconds_present == {0, 30}, ( + f'expected seconds 0 and 30 to be preserved, got {seconds_present}' + ) + + instance = model_analysis_stub.last_instance + assert instance is not None + assert instance.detect_univariate_drift_calls[0]['chunk_period'] == 's' diff --git a/e2e/test_minimal_retrain.py b/e2e/test_minimal_retrain.py new file mode 100644 index 0000000..32f4bb6 --- /dev/null +++ b/e2e/test_minimal_retrain.py @@ -0,0 +1,335 @@ +""" +End-to-end tests for the MinimalRetrain workflow. + +The MLflow registry is fully stubbed because no real artifacts exist in a +test container; we only validate that the workflow: + +- Loads training data via ``load_query_with_minio_offload``. +- Calls ``retrain_model`` with a payload pointing at MinIO. +- Calls ``update_production_model`` only when retrain succeeds. +- Persists ``retrain_reports`` rows with all required columns; success rows + carry the new ``version`` / ``mlflow_run_id`` / ``mlflow_experiment_id`` + while failure rows leave them ``NULL``. +""" + +from contextlib import contextmanager +from datetime import datetime, timedelta, timezone +from unittest.mock import MagicMock, patch + +import pytest +from sqlalchemy import text +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Worker + +from e2e.helpers import ( + insert_target_data_for_drift, + load_scenario_input, + make_workflow_id, + start_and_await_workflow, +) +from laborious.activities.activities import Activities +from laborious.workflows.minimal_retrain import MinimalRetrain + +EXPECTED_RETRAIN_REPORT_COLUMNS = [ + 'id', + 'model_id', + 'model_name', + 'timestamp', + 'status', + 'version', + 'mlflow_run_id', + 'mlflow_experiment_id', + 'created_at', +] + + +def _retrain_input(model_id: int, **overrides) -> dict: + """Load and override the minimal-retrain base scenario.""" + payload = load_scenario_input('minimal_retrain_base.json', model_id=model_id) + payload.update(overrides) + return payload + + +def _seed_retrain_training_rows(postgres_engine, model_id: int) -> None: + """ + Insert training rows in long format that pivot cleanly into + ``index=timestamp`` / ``columns={sensor_1, sensor_2}`` for ``retrain_model``. + """ + target_timestamps = [ + (datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(minutes=10 - i)) + .strftime('%Y-%m-%d %H:%M:%S%z') + for i in range(5) + ] + insert_target_data_for_drift( + postgres_engine, + model_id=model_id, + timestamps=target_timestamps, + variables_values={ + 'sensor_1': [10.0, 11.0, 12.0, 13.0, 14.0], + 'sensor_2': [20.0, 21.0, 22.0, 23.0, 24.0], + }, + ) + + +def _configure_retrain_happy_path(mlflow_repository_stub) -> None: + """ + Wire ``mlflow_repository_stub`` so retrain + update_production succeed. + + Mocks (in order of consumption): + + - ``_client.get_model_version_by_alias``: returns ``mv`` with a stable + ``run_id`` (used as ``source_run_id``). + - ``get_cached_model``: returns a wrapper exposing inert ``retrain`` and + ``store_model`` methods. + - ``start_run``: returns a context manager yielding a ``run_info`` with + run/experiment ids. + - ``log_params``: inert. + - ``_client.search_model_versions``: returns one registry entry whose + ``version`` is promoted by ``update_production_model``. + - ``promote_to_alias``: inert success. + """ + mv_src = MagicMock() + mv_src.run_id = 'source-run-id' + + new_version = MagicMock() + new_version.version = '7' + new_version.run_id = 'retrain-run-id' + + mlflow_repository_stub._client.get_model_version_by_alias.return_value = mv_src + + cached_wrapper = MagicMock() + cached_wrapper.retrain = MagicMock(return_value=None) + cached_wrapper.store_model = MagicMock(return_value=None) + mlflow_repository_stub.get_cached_model.return_value = cached_wrapper + + @contextmanager + def fake_start_run(**kwargs): + run_info = MagicMock() + run_info.run_id = 'retrain-run-id' + run_info.experiment_id = 'experiment-id' + yield run_info + + mlflow_repository_stub.start_run.side_effect = fake_start_run + mlflow_repository_stub.log_params = MagicMock(return_value=None) + mlflow_repository_stub._client.search_model_versions.return_value = [new_version] + mlflow_repository_stub.promote_to_alias = MagicMock(return_value=None) + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_minimal_retrain_happy_path_writes_success_report( + temporal_test_env: WorkflowEnvironment, + temporal_worker_minimal_retrain: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, +): + """ + Scenario MR.1.1: Retrain succeeds. ``retrain_reports`` must contain a + success row with version/mlflow_run_id/mlflow_experiment_id populated and + the registry must have been told to promote the new version to the + configured alias. + """ + client = temporal_test_env.client + model_id = 711 + + _seed_retrain_training_rows(postgres_engine, model_id) + _configure_retrain_happy_path(mlflow_repository_stub) + + input_data = _retrain_input(model_id) + + with patch('laborious.activities.mlflow.mlflow.log_artifact') as log_artifact_mock: + await start_and_await_workflow( + client, + MinimalRetrain.run, + input_data, + make_workflow_id('test-retrain-happy'), + ) + + assert log_artifact_mock.called, 'retrain_model should log the input CSV artifact' + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT * FROM predictions_schema.retrain_reports ' + 'WHERE model_id = :m' + ), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert len(rows) == 1 + for column in EXPECTED_RETRAIN_REPORT_COLUMNS: + assert column in rows[0], f'Missing retrain report column: {column}' + + row = rows[0] + assert row['model_id'] == model_id + assert row['model_name'] == 'test_model' + assert row['status'] == 'Model retrained successfully.' + assert row['version'] == '7' + assert row['mlflow_run_id'] == 'retrain-run-id' + assert row['mlflow_experiment_id'] == 'experiment-id' + assert row['timestamp'] is not None + assert row['created_at'] is not None + + mlflow_repository_stub.promote_to_alias.assert_called_once() + promote_kwargs = mlflow_repository_stub.promote_to_alias.call_args.kwargs + assert promote_kwargs['model_name'] == 'test_model' + assert promote_kwargs['version'] == '7' + assert promote_kwargs['alias'] == 'production' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_minimal_retrain_failure_writes_report_without_version_columns( + temporal_test_env: WorkflowEnvironment, + temporal_worker_minimal_retrain: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, +): + """ + Scenario MR.2.1: ``wrapper.retrain`` raises. The activity must catch the + error, return ``success=False`` so ``update_production_model`` is skipped, + and ``format_retrain_report`` must produce a row with the error message + and NULL version columns. + """ + client = temporal_test_env.client + model_id = 721 + + _seed_retrain_training_rows(postgres_engine, model_id) + _configure_retrain_happy_path(mlflow_repository_stub) + mlflow_repository_stub.get_cached_model.return_value.retrain.side_effect = RuntimeError( + 'training did not converge' + ) + + input_data = _retrain_input(model_id) + + with patch('laborious.activities.mlflow.mlflow.log_artifact'): + await start_and_await_workflow( + client, + MinimalRetrain.run, + input_data, + make_workflow_id('test-retrain-failure'), + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT * FROM predictions_schema.retrain_reports ' + 'WHERE model_id = :m' + ), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert len(rows) == 1 + row = rows[0] + assert row['model_id'] == model_id + assert row['model_name'] == 'test_model' + assert 'training did not converge' in row['status'] + assert row['version'] is None + assert row['mlflow_run_id'] is None + assert row['mlflow_experiment_id'] is None + + mlflow_repository_stub.promote_to_alias.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_minimal_retrain_missing_target_writes_failure_report( + temporal_test_env: WorkflowEnvironment, + temporal_worker_minimal_retrain: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, +): + """ + Scenario MR.2.2: ``model_config`` does not declare ``target``. The retrain + activity must short-circuit before any MLflow call and the report row must + carry the explicit guard message. + """ + client = temporal_test_env.client + model_id = 722 + + _seed_retrain_training_rows(postgres_engine, model_id) + _configure_retrain_happy_path(mlflow_repository_stub) + + input_data = _retrain_input(model_id, model_config={}) + + with patch('laborious.activities.mlflow.mlflow.log_artifact'): + await start_and_await_workflow( + client, + MinimalRetrain.run, + input_data, + make_workflow_id('test-retrain-missing-target'), + ) + + with postgres_engine.connect() as conn: + row = ( + conn.execute( + text( + 'SELECT * FROM predictions_schema.retrain_reports ' + 'WHERE model_id = :m' + ), + {'m': model_id}, + ) + .mappings() + .first() + ) + + assert row is not None + assert 'target' in row['status'].lower(), ( + f"expected target-missing message, got status={row['status']!r}" + ) + assert row['version'] is None + mlflow_repository_stub.get_cached_model.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_minimal_retrain_no_training_data_does_not_persist_report( + temporal_test_env: WorkflowEnvironment, + temporal_worker_minimal_retrain: Worker, + test_activities: Activities, + postgres_engine, + mlflow_repository_stub, +): + """ + Scenario MR.3.1: When the training query returns no rows the workflow must + not persist any report row. The workflow currently raises plain + ``ValueError`` which Temporal treats as a workflow-task failure (causing + indefinite retries until the test environment times out), so the assertion + here is constrained to the persistence side-effect. See ``e2e/CODE_ISSUES.md`` + issue MR-1 for the recommended ``ApplicationError`` fix. + """ + client = temporal_test_env.client + model_id = 731 + + with postgres_engine.begin() as conn: + conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}')) + + _configure_retrain_happy_path(mlflow_repository_stub) + + input_data = _retrain_input(model_id) + + with pytest.raises(Exception), patch('laborious.activities.mlflow.mlflow.log_artifact'): + await start_and_await_workflow( + client, + MinimalRetrain.run, + input_data, + make_workflow_id('test-retrain-no-data'), + ) + + with postgres_engine.connect() as conn: + count = conn.execute( + text('SELECT COUNT(*) FROM predictions_schema.retrain_reports WHERE model_id = :m'), + {'m': model_id}, + ).scalar() + assert count == 0 diff --git a/e2e/test_minio_offload.py b/e2e/test_minio_offload.py index 694e208..3b71ab7 100644 --- a/e2e/test_minio_offload.py +++ b/e2e/test_minio_offload.py @@ -9,7 +9,12 @@ from sqlalchemy import text from temporalio.testing import WorkflowEnvironment from temporalio.worker import Worker -from e2e.helpers import insert_sample_data, make_workflow_id, start_and_await_workflow +from e2e.helpers import ( + insert_sample_data, + load_scenario_input, + make_workflow_id, + start_and_await_workflow, +) from laborious.activities.activities import Activities from laborious.utils.models import minio_dataframe_payload as mdp from laborious.workflows.predictions_batch import PredictionsBatch @@ -32,30 +37,14 @@ async def test_load_query_with_minio_offload_writes_object_to_bucket( conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}')) insert_sample_data(postgres_engine, model_id, [1.0, 2.0]) - metadata = { - 'metadata': { - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': model_id, - 'workflow_name': 'predictions_batch', - } - } + scenario_input = load_scenario_input('minio_offload_load_query.json', model_id=model_id) + metadata = {'metadata': scenario_input['metadata']} with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1): - payload = await test_activities_real_minio.load_query_with_minio_offload( - { - **metadata, - 'query': ( - 'SELECT timestamp, variable, value, created_at ' - f'FROM predictions_schema.laborious_data WHERE model_id = {model_id}' - ), - 'model_name': 'test_model', - 'datetime_columns': ['timestamp', 'created_at'], - } - ) + payload = test_activities_real_minio.load_query_with_minio_offload(scenario_input) assert payload.object_key, 'offloaded payload must reference a MinIO object' assert payload.data is None or payload.data == {}, 'large payloads should not inline tabular dict' - df = await payload.retrieve(test_activities_real_minio.minio_repository, metadata['metadata']) + df = payload.retrieve(test_activities_real_minio.minio_repository, metadata['metadata']) assert len(df) >= 1 client = minio_container.get_client() @@ -82,31 +71,7 @@ async def test_predictions_batch_with_minio_offload_path( conn.execute(text(f'DELETE FROM predictions_schema.transformed_data WHERE model_id = {model_id}')) insert_sample_data(postgres_engine, model_id, [10.0, 20.0, 30.0]) - input_data = { - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': model_id, - 'query': ( - 'SELECT timestamp, variable, value, created_at ' - f'FROM predictions_schema.laborious_data WHERE model_id = {model_id}' - ), - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}}, - 'mlflow_transform_filters': {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}}, - 'mlflow_predict_filters': {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}}, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'retention_minutes': 0, - 'target': 'sensor_1', - }, - 'datetime_columns': ['timestamp', 'created_at'], - } + input_data = load_scenario_input('minio_offload_workflow.json', model_id=model_id) with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1): await start_and_await_workflow( @@ -121,3 +86,23 @@ async def test_predictions_batch_with_minio_offload_path( text(f'SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = {model_id}') ).scalar() assert count == 1 + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_load_query_with_inline_payload_when_below_threshold( + postgres_engine, + test_activities_real_minio: Activities, +): + """Scenario 4.2.1: payload stays inline when threshold is high enough.""" + model_id = 503 + with postgres_engine.begin() as conn: + conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}')) + insert_sample_data(postgres_engine, model_id, [1.0, 2.0]) + + scenario_input = load_scenario_input('minio_offload_load_query.json', model_id=model_id) + with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 10**9): + payload = test_activities_real_minio.load_query_with_minio_offload(scenario_input) + + assert payload.object_key is None + assert payload.data is not None diff --git a/e2e/test_predictions_batch_format_export.py b/e2e/test_predictions_batch_format_export.py index 91c795c..c09fc80 100644 --- a/e2e/test_predictions_batch_format_export.py +++ b/e2e/test_predictions_batch_format_export.py @@ -4,7 +4,7 @@ End-to-end tests for PredictionsBatch workflow - Format and Export scenarios. from decimal import Decimal from typing import Any, cast -from unittest.mock import ANY, AsyncMock, call +from unittest.mock import ANY, call import pytest from sientia_do.notifications.models import NotificationLevel @@ -12,47 +12,18 @@ from sqlalchemy import text from temporalio.testing import WorkflowEnvironment from temporalio.worker import Worker -from e2e.helpers import assert_prediction, insert_sample_data, make_workflow_id, start_and_await_workflow +from e2e.helpers import ( + assert_prediction, + insert_sample_data, + load_scenario_input, + make_workflow_id, + start_and_await_workflow, +) from laborious.activities.activities import Activities from laborious.workflows.predictions_batch import PredictionsBatch -base_input_data = { - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 301, - 'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 301', - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': { - 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_transform_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_predict_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'retention_minutes': 0, - 'target': 'sensor_1', - }, - 'datetime_columns': ['timestamp', 'created_at'], -} - -base_query = "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {model_id}" - def get_base_input_data(model_id): - return { - **base_input_data, - 'model_id': model_id, - 'query': base_query.format(model_id=model_id), - } + return load_scenario_input('format_export_base.json', model_id=model_id) @pytest.mark.asyncio @@ -668,8 +639,8 @@ async def test_scenario_3_2_3_pi_web_api_partial_write_error( insert_sample_data(postgres_engine, model_id, [23.5, 78.2]) - test_activities.pi_web_api_client.write_value = AsyncMock( - side_effect=[ + test_activities.pi_web_api_client.set_side_effect( + [ # Prediction batch: two web_ids requested, only one acknowledged. [{'WebId': 'web_id_1', 'Errors': []}], # Confidence write succeeds. @@ -707,4 +678,42 @@ async def test_scenario_3_2_3_pi_web_api_partial_write_error( prediction_confidence=13, comments="The number of written tags does not match the number of tag names: Expected ['tag_1', 'tag_3'] tags, but ['tag_1'] tags were written.", ) + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_scenario_3_3_1_combined_pi_and_opc_outputs( + temporal_test_env: WorkflowEnvironment, + temporal_worker: Worker, + test_activities: Activities, + postgres_engine, +): + """ + Scenario 3.3.1: PI and OPC enabled together. + """ + client = temporal_test_env.client + model_id = 333 + insert_sample_data(postgres_engine, model_id, [23.5, 78.2]) + + input_data = get_base_input_data(model_id) + input_data['pi_web_api_output_config'] = { + 'endpoint': 'test_endpoint', + 'prediction_tags': {'tag_1': 'web_id_1'}, + 'confidence_tags': {'tag_2': 'web_id_2'}, + } + input_data['opc_output_config'] = { + '1': { + 'prediction_tags': {'addr_1': {'data_type': 'float'}}, + 'confidence_tags': {'addr_2': {'data_type': 'float'}}, + } + } + + await start_and_await_workflow( + client, PredictionsBatch.run, input_data, make_workflow_id('test-pi-opc-combined') + ) + + assert test_activities.pi_web_api_client.write_value.call_count == 2 + opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data) + assert opc_write_data.call_count == 2 + assert_prediction(postgres_engine, model_id) diff --git a/e2e/test_predictions_batch_main_workflow.py b/e2e/test_predictions_batch_main_workflow.py index 1408031..89b81df 100644 --- a/e2e/test_predictions_batch_main_workflow.py +++ b/e2e/test_predictions_batch_main_workflow.py @@ -10,7 +10,7 @@ from temporalio.worker import Worker import pytest -from e2e.helpers import make_workflow_id, start_and_await_workflow +from e2e.helpers import load_scenario_input, make_workflow_id, start_and_await_workflow from laborious.activities.activities import Activities from laborious.workflows.predictions_batch import PredictionsBatch @@ -37,42 +37,7 @@ async def test_scenario_1_1_1_happy_path_complete_success( """ conn.execute(text(insert_sql)) - input_data = { - 'metadata': { - 'metadata': { - 'model_id': 123, - 'model_name': 'test_model', - 'schedule_name': 'test-schedule', - 'workflow_name': 'predictions_batch', - } - }, - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 123, - 'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 123', - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': { - 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_transform_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_predict_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'target': 'sensor_1', - 'retention_minutes': 0, - }, - 'datetime_columns': ['timestamp', 'created_at'], - } + input_data = load_scenario_input('main_happy_path.json', model_id=123) await start_and_await_workflow( client, @@ -125,41 +90,7 @@ async def test_scenario_1_2_1_sql_query_execution_error( """Invalid SQL: workflow may complete with early exit; no prediction rows.""" client = temporal_test_env.client - input_data = { - 'metadata': { - 'metadata': { - 'model_id': 128, - 'model_name': 'test_model', - 'schedule_name': 'test-schedule', - 'workflow_name': 'predictions_batch', - } - }, - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 128, - 'query': 'SELECT * FROM nonexistent_table WHERE invalid_syntax =', - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': { - 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_transform_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_predict_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'target': 'sensor_1', - 'retention_minutes': 0, - }, - } + input_data = load_scenario_input('main_sql_error.json', model_id=128) await start_and_await_workflow( client, @@ -186,22 +117,7 @@ async def test_scenario_1_2_2_missing_required_parameters( """Missing query: workflow does not produce predictions and is terminated explicitly.""" client = temporal_test_env.client - input_data = { - 'metadata': { - 'metadata': { - 'model_id': 129, - 'model_name': 'test_model', - 'schedule_name': 'test-schedule', - 'workflow_name': 'predictions_batch', - } - }, - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 129, - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - } + input_data = load_scenario_input('main_missing_required.json', model_id=129) handle = await client.start_workflow( PredictionsBatch.run, @@ -244,42 +160,7 @@ async def test_scenario_1_2_3_invalid_datetime_column_specification( ) ) - input_data = { - 'metadata': { - 'metadata': { - 'model_id': 130, - 'model_name': 'test_model', - 'schedule_name': 'test-schedule', - 'workflow_name': 'predictions_batch', - } - }, - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 130, - 'query': 'SELECT timestamp, variable, value FROM predictions_schema.laborious_data WHERE model_id = 130', - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': { - 'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_transform_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_predict_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'target': 'sensor_1', - 'retention_minutes': 0, - }, - 'datetime_columns': ['nonexistent_column'], - } + input_data = load_scenario_input('main_invalid_datetime.json', model_id=130) handle = await client.start_workflow( PredictionsBatch.run, diff --git a/e2e/test_predictions_batch_prediction_process.py b/e2e/test_predictions_batch_prediction_process.py index 80249f0..5148670 100644 --- a/e2e/test_predictions_batch_prediction_process.py +++ b/e2e/test_predictions_batch_prediction_process.py @@ -14,6 +14,7 @@ from temporalio.worker import Worker from e2e.helpers import ( assert_continue, + load_scenario_input, assert_repeat, assert_stop, insert_sample_data, @@ -23,47 +24,8 @@ from e2e.helpers import ( from laborious.activities.activities import Activities from laborious.workflows.predictions_batch import PredictionsBatch -base_input_data = { - 'schedule_name': 'test-schedule', - 'model_name': 'test_model', - 'model_id': 201, - 'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 201', - 'schema': 'predictions_schema', - 'table_name': 'predictions', - 'transform_table_name': 'transformed_data', - 'input_filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'CONTINUE', - 'CONFIG': {'variables': ['sensor_1']}, - }, - }, - 'mlflow_transform_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'mlflow_predict_filters': { - 'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}, - }, - 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], - 'opc_output_config': {}, - 'pi_web_api_output_config': {}, - 'save_transform': True, - 'prediction_store_policy': 'lts:1', - 'model_config': { - 'retention_minutes': 0, - 'target': 'sensor_1', - }, - 'datetime_columns': ['timestamp', 'created_at'], -} - -base_query = 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {model_id}' - - def get_base_input_data(model_id): - return { - **base_input_data, - 'model_id': model_id, - 'query': base_query.format(model_id=model_id), - } + return load_scenario_input('prediction_process_base.json', model_id=model_id) def insert_sample_prediction(postgres_engine, model_id): @@ -383,3 +345,32 @@ async def test_scenario_2_4_1_input_empty_data_stop( client, PredictionsBatch.run, input_data, make_workflow_id('test-empty-data-stop') ) assert_stop(postgres_engine, model_id) + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_scenario_2_4_1_priority_conflict_resolution( + temporal_test_env: WorkflowEnvironment, + temporal_worker: Worker, + test_activities: Activities, + postgres_engine, + bad_data_model, +): + """Conflicting filter outputs must honor configured path_priority order.""" + client = temporal_test_env.client + model_id = 242 + 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', 'CONFIG': {}}} + input_data['mlflow_predict_filters'] = {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}} + input_data['path_priority'] = ['STOP', 'CONTINUE', 'REPEAT'] + + await start_and_await_workflow( + client, PredictionsBatch.run, input_data, make_workflow_id('test-priority-conflict') + ) + assert_continue( + postgres_engine=postgres_engine, + model_id=model_id, + prediction_confidence=Decimal(10), + comments='Unknown MLFlow API error', + ) diff --git a/e2e/test_simple_metrics.py b/e2e/test_simple_metrics.py new file mode 100644 index 0000000..2413e81 --- /dev/null +++ b/e2e/test_simple_metrics.py @@ -0,0 +1,322 @@ +""" +End-to-end tests for the SimpleMetrics workflow. + +Coverage focus: + +- Happy path computes rmse/mse/mae/r2 from predictions joined against ``laborious_data`` + and persists rows to ``predictions_schema.simple_metrics_data`` with all required columns. +- Subset metric selection (only rmse) writes exactly the requested rows. +- Zero-variance target produces ``r2=0`` per division-by-zero guard. +- Empty join (no overlapping data) short-circuits without persisting anything. +""" + +from datetime import datetime, timedelta, timezone +from decimal import Decimal + +import math +import pytest +from sqlalchemy import text +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Worker + +from e2e.helpers import load_scenario_input, make_workflow_id, start_and_await_workflow +from laborious.activities.activities import Activities +from laborious.workflows.simple_metrics import SimpleMetrics + +EXPECTED_SIMPLE_METRICS_COLUMNS = [ + 'id', + 'model_id', + 'metric', + 'value', + 'timestamp', + 'data_size', + 'interval_minutes', + 'created_at', +] + + +def _simple_metrics_input(model_id: int, **overrides) -> dict: + """Load and override the simple-metrics base scenario.""" + payload = load_scenario_input('simple_metrics_base.json', model_id=model_id) + payload.update(overrides) + return payload + + +def _seed_predictions_and_targets( + postgres_engine, + model_id: int, + pairs: list[tuple[float, float]], + target_name: str = 'sensor_target', + offset_minutes: int = 6, +) -> list[str]: + """ + Insert matching prediction/target rows used by the SimpleMetrics SQL JOIN. + + For each ``(prediction, target)`` pair we write a row in ``predictions`` and + a matching row in ``laborious_data`` with ``variable=target_name`` so the + inner join in the workflow query yields one row per pair. + + Args: + - postgres_engine: SQLAlchemy engine bound to the test container. + - model_id: Model id stamped on every row. + - pairs: ``(prediction, target)`` pairs, one per minute. + - target_name: Variable name in ``laborious_data`` representing the target. + - offset_minutes: Earliest row sits this many minutes ago so timestamps fall + inside the workflow's recent-data window. + + Return: + List of timestamp strings written for the inserted rows. + """ + base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta( + minutes=offset_minutes + ) + timestamps = [ + (base + timedelta(minutes=i)).strftime('%Y-%m-%d %H:%M:%S%z') + for i in range(len(pairs)) + ] + + prediction_rows = [] + target_rows = [] + for index, (prediction, target_value) in enumerate(pairs): + ts = timestamps[index] + prediction_rows.append( + f"({model_id}, {prediction}, 0, 0, 'Good', '{ts}', '{ts}')" + ) + target_rows.append( + f"({model_id}, '{target_name}', {target_value}, '{ts}', '{ts}')" + ) + + with postgres_engine.begin() as conn: + conn.execute(text(f'DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}')) + conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}')) + # The SimpleMetrics SQL JOIN only filters ``predictions.model_id``; it does + # NOT filter ``laborious_data.model_id`` (see ``e2e/CODE_ISSUES.md`` issue + # SM-1). Without this cross-model cleanup, a previous test's target rows + # under the same variable name would join into this test's predictions + # whenever timestamps happened to overlap. + conn.execute( + text( + "DELETE FROM predictions_schema.laborious_data " + "WHERE variable IN (:sensor_default, :target_name) " + "AND timestamp >= NOW() - INTERVAL '120 minutes'" + ), + {'sensor_default': 'sensor_target', 'target_name': target_name}, + ) + if prediction_rows: + conn.execute( + text( + 'INSERT INTO predictions_schema.predictions ' + '(model_id, prediction, prediction_confidence, response_time, ' + 'prediction_status, "timestamp", created_at) VALUES ' + + ', '.join(prediction_rows) + ) + ) + conn.execute( + text( + 'INSERT INTO predictions_schema.laborious_data ' + '(model_id, variable, value, "timestamp", created_at) VALUES ' + + ', '.join(target_rows) + ) + ) + return timestamps + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_simple_metrics_happy_path_persists_all_metrics_and_columns( + temporal_test_env: WorkflowEnvironment, + temporal_worker_simple_metrics: Worker, + test_activities: Activities, + postgres_engine, +): + """ + Scenario S.1.1: rmse/mse/mae/r2 are calculated from a deterministic + prediction/target pair set and written one row per metric. Every column + expected by ``predictions_schema.simple_metrics_data`` must be populated and + the numerical values must match closed-form expectations. + """ + client = temporal_test_env.client + model_id = 511 + + pairs = [ + (1.0, 2.0), + (2.0, 4.0), + (3.0, 5.0), + (4.0, 9.0), + (5.0, 12.0), + ] + diffs = [target - prediction for prediction, target in pairs] + n = len(diffs) + expected_rmse = math.sqrt(sum(d * d for d in diffs) / n) + expected_mse = sum(d * d for d in diffs) / n + expected_mae = sum(abs(d) for d in diffs) / n + target_mean = sum(t for _, t in pairs) / n + ss_res = sum((target - prediction) ** 2 for prediction, target in pairs) + ss_tot = sum((t - target_mean) ** 2 for _, t in pairs) + expected_r2 = 1.0 - (ss_res / ss_tot) + + _seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs) + + input_data = _simple_metrics_input(model_id) + await start_and_await_workflow( + client, + SimpleMetrics.run, + input_data, + make_workflow_id('test-simple-metrics-happy'), + ) + + with postgres_engine.connect() as conn: + rows = ( + conn.execute( + text( + 'SELECT * FROM predictions_schema.simple_metrics_data ' + 'WHERE model_id = :m ORDER BY metric' + ), + {'m': model_id}, + ) + .mappings() + .all() + ) + + assert len(rows) == 4, f'Expected 4 metric rows, got {len(rows)}' + for column in EXPECTED_SIMPLE_METRICS_COLUMNS: + assert column in rows[0], f'Missing simple_metrics column: {column}' + for row in rows: + for column in EXPECTED_SIMPLE_METRICS_COLUMNS: + assert row[column] is not None, f"Column '{column}' is NULL in {dict(row)}" + + by_metric = {row['metric']: row for row in rows} + assert set(by_metric) == {'rmse', 'mse', 'mae', 'r2'} + + def _decimal_close(actual, expected, places: int = 6) -> bool: + return abs(float(actual) - expected) < 10 ** (-places) + + assert _decimal_close(by_metric['rmse']['value'], expected_rmse) + assert _decimal_close(by_metric['mse']['value'], expected_mse) + assert _decimal_close(by_metric['mae']['value'], expected_mae) + assert _decimal_close(by_metric['r2']['value'], expected_r2) + + assert all(row['data_size'] == n for row in rows), 'data_size must equal target row count' + assert all(row['interval_minutes'] == 60 for row in rows) + assert all(row['model_id'] == model_id for row in rows) + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_simple_metrics_subset_metrics_writes_only_requested_rows( + temporal_test_env: WorkflowEnvironment, + temporal_worker_simple_metrics: Worker, + test_activities: Activities, + postgres_engine, +): + """ + Scenario S.1.2: Requesting ``metrics=['rmse']`` must persist exactly one row + with metric ``rmse`` and skip mse/mae/r2. + """ + client = temporal_test_env.client + model_id = 512 + + pairs = [(1.0, 2.0), (2.0, 4.0), (3.0, 6.0)] + _seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs) + + input_data = _simple_metrics_input(model_id, metrics=['rmse']) + await start_and_await_workflow( + client, + SimpleMetrics.run, + input_data, + make_workflow_id('test-simple-metrics-subset'), + ) + + with postgres_engine.connect() as conn: + metrics = [ + r[0] + for r in conn.execute( + text( + 'SELECT metric FROM predictions_schema.simple_metrics_data ' + 'WHERE model_id = :m' + ), + {'m': model_id}, + ).all() + ] + assert metrics == ['rmse'] + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_simple_metrics_zero_variance_target_returns_zero_r2( + temporal_test_env: WorkflowEnvironment, + temporal_worker_simple_metrics: Worker, + test_activities: Activities, + postgres_engine, +): + """ + Scenario S.2.1: When the target column has zero variance the activity must + return ``r2 = 0`` (division-by-zero guard) and still persist all four metrics. + """ + client = temporal_test_env.client + model_id = 521 + + pairs = [(0.0, 5.0), (1.0, 5.0), (2.0, 5.0), (3.0, 5.0)] + _seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs) + + input_data = _simple_metrics_input(model_id) + await start_and_await_workflow( + client, + SimpleMetrics.run, + input_data, + make_workflow_id('test-simple-metrics-zero-variance'), + ) + + with postgres_engine.connect() as conn: + r2_value = conn.execute( + text( + "SELECT value FROM predictions_schema.simple_metrics_data " + "WHERE model_id = :m AND metric = 'r2'" + ), + {'m': model_id}, + ).scalar() + assert r2_value is not None + assert Decimal(str(r2_value)) == Decimal('0'), f'expected r2=0, got {r2_value!r}' + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_simple_metrics_no_overlapping_data_short_circuits( + temporal_test_env: WorkflowEnvironment, + temporal_worker_simple_metrics: Worker, + test_activities: Activities, + postgres_engine, +): + """ + Scenario S.3.1: When the join produces no rows (no matching laborious_data + row for the configured ``target``), the workflow returns early without + invoking ``calculate_simple_metrics`` and writes nothing. + """ + client = temporal_test_env.client + model_id = 531 + + # Insert predictions but no matching target rows for the configured variable. + _seed_predictions_and_targets( + postgres_engine, + model_id=model_id, + pairs=[(1.0, 1.0)], + target_name='wrong_variable_name', + ) + + input_data = _simple_metrics_input(model_id) + await start_and_await_workflow( + client, + SimpleMetrics.run, + input_data, + make_workflow_id('test-simple-metrics-empty-join'), + ) + + with postgres_engine.connect() as conn: + count = conn.execute( + text( + 'SELECT COUNT(*) FROM predictions_schema.simple_metrics_data ' + 'WHERE model_id = :m' + ), + {'m': model_id}, + ).scalar() + assert count == 0, 'Empty target data must short-circuit and skip persistence' diff --git a/input_sample.json b/input_sample.json new file mode 100644 index 0000000..da3ef69 --- /dev/null +++ b/input_sample.json @@ -0,0 +1,72 @@ +{ + "models": [ + { + "id": "1001", + "name": "test-runtime", + "active": false, + "model_config": { + "alias": "production", + "retention_minutes": 60, + "target": "Square" + } + } + ], + "pipelines": [ + { + "schedule_name": "laborious-test-runtime", + "model_id": "1001", + "workflow_type": "predictions_batch", + "frequency": "60s", + "max_retry_policy": 1, + "query": "select * from sientia_data.laborious_data where model_id = 1 and \"timestamp\" > NOW() - INTERVAL '5 minutes' order by \"timestamp\" desc limit 30;", + "retention_time": 60, + "write_tags": [], + "input_filters": [ + { + "filter_name": "EMPTY_DATA", + "policy": "STOP" + }, + { + "filter_name": "SPECIFIC_VARIABLES_NULL_VALUES", + "policy": "CONTINUE", + "config": { + "variables": [ + "Counter" + ] + } + } + ], + "mlflow_transform_filters": [ + { + "filter_name": "API_ERROR", + "policy": "REPEAT" + }, + { + "filter_name": "NAN_VALUES", + "policy": "STOP" + } + ], + "mlflow_predict_filters": [ + { + "filter_name": "API_ERROR", + "policy": "CONTINUE" + } + ], + "path_priority": [ + "STOP", + "CONTINUE", + "REPEAT" + ], + "active": true, + "updated_at": { + "$date": "2026-05-07T23:35:01.600Z" + }, + "save_transform": false, + "pi_web_api_output_config": {}, + "datetime_columns": [ + "timestamp", + "created_at" + ] + } + ] +} diff --git a/laborious/activities/activities.py b/laborious/activities/activities.py index e92b12e..45e9541 100644 --- a/laborious/activities/activities.py +++ b/laborious/activities/activities.py @@ -6,7 +6,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.observability.logger import Logger - from sientia_do.repository.minio_repository import MinioRepository + from sientia_do.repository.minio_repository_sync import MinioRepository from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.plugin_store import PluginStore @@ -159,7 +159,7 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): metrics_controller=mc, ) - async def shutdown(self): + def shutdown(self) -> None: """ Close database pools, sync clients, and OPC sessions in a defined order. @@ -172,6 +172,6 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): Storage.close(self) MLFlow.close(self) Gates.close(self) - await OPC.close(self) + OPC.close(self) ModelMetrics.close(self) API.close(self) diff --git a/laborious/activities/api.py b/laborious/activities/api.py index 0e4114f..8b651d9 100644 --- a/laborious/activities/api.py +++ b/laborious/activities/api.py @@ -11,7 +11,7 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController from sientia_do.observability.sientia_monitoring import SientiaMonitoring - from sientia_do.repository.pi_web_api_client import PIWebAPIClient + from sientia_do.repository.pi_web_api_client_sync import PIWebAPIClient from laborious import metrics @@ -105,7 +105,7 @@ class API(SientiaMonitoring): self.pi_web_api_client.close() SientiaMonitoring.shutdown(self) - async def process_pi_web_api_response( + def process_pi_web_api_response( self, response_data: list[dict[str, Any]], tags: dict[str, str], @@ -156,7 +156,7 @@ class API(SientiaMonitoring): self.error( f'Error writing tag {tag_name}:{web_id} to PI Web API: {errors}', metadata ) - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PI_WEB_API_PREDICTION_WRITTEN_ERROR_COUNT, tags={ **core_labels, @@ -165,7 +165,7 @@ class API(SientiaMonitoring): ) confidence = PI_WEB_API_PREDICTION_ERROR_CONFIDENCE else: - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PI_WEB_API_PREDICTION_WRITTEN_COUNT, tags={ **core_labels, @@ -182,7 +182,7 @@ class API(SientiaMonitoring): metadata, ) - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR', message=f'The number of written tags does not match the number of tag names: Expected {tag_names} tags, but {written_tags} tags were written.\nResponse:\n {json.dumps(response_data, indent=4)}\nTags:\n {json.dumps(tags, indent=4)}', @@ -194,7 +194,7 @@ class API(SientiaMonitoring): return confidence, message @activity.defn(name='write_pi_web_api_data') - async def write_pi_web_api_data(self, input_data: dict[str, Any]) -> dict[Any, Any]: + def write_pi_web_api_data(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Write prediction and confidence data to PI Web API. @@ -232,7 +232,7 @@ class API(SientiaMonitoring): confidence_value = data.head(1)['prediction_confidence'].values[0] try: - prediction_response = await self.pi_web_api_client.write_value( + prediction_response = self.pi_web_api_client.write_value( web_ids=prediction_tags, value={ 'Timestamp': data.head(1)['timestamp'].values[0], @@ -241,7 +241,7 @@ class API(SientiaMonitoring): metadata=metadata, ) - confidence, message = await self.process_pi_web_api_response( + confidence, message = self.process_pi_web_api_response( response_data=prediction_response, tags=raw_prediction_tags, core_labels=core_labels, @@ -258,7 +258,7 @@ class API(SientiaMonitoring): except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR', message=f'Error writing prediction data to PI Web API: {e}\n Tags: {raw_prediction_tags}', @@ -275,7 +275,7 @@ class API(SientiaMonitoring): return data.to_dict() try: - confidence_response = await self.pi_web_api_client.write_value( + confidence_response = self.pi_web_api_client.write_value( web_ids=confidence_tags, value={ 'Timestamp': data.head(1)['timestamp'].values[0], @@ -284,7 +284,7 @@ class API(SientiaMonitoring): metadata=metadata, ) - await self.process_pi_web_api_response( + self.process_pi_web_api_response( response_data=confidence_response, tags=raw_confidence_tags, core_labels=core_labels, @@ -293,7 +293,7 @@ class API(SientiaMonitoring): except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR', message=f'Error writing confidence data to PI Web API: {e}\n Tags: {raw_confidence_tags}', diff --git a/laborious/activities/gates.py b/laborious/activities/gates.py index ad88ca1..f72bd2e 100644 --- a/laborious/activities/gates.py +++ b/laborious/activities/gates.py @@ -1,8 +1,5 @@ -from sientia_do.repository.minio_repository import MinioRepository from temporalio import activity, workflow -from laborious.utils.repository.minio_manager import MinioManager - with workflow.unsafe.imports_passed_through(): import traceback from collections.abc import Callable, Mapping @@ -13,6 +10,8 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController + from sientia_do.observability.sientia_monitoring import SientiaMonitoring + from sientia_do.repository.minio_repository_sync import MinioRepository from sientia_do.utils.formatters import create_sample_dict from laborious import metrics @@ -66,7 +65,7 @@ mlflow_content_path_confidence: Mapping[str, int] = { } -class Gates(MinioManager): +class Gates(SientiaMonitoring): """ Data quality gates and filtering activities for the Laborious system. @@ -106,8 +105,12 @@ class Gates(MinioManager): Raises: Exception: If BaseActivity initialization fails """ - MinioManager.__init__( - self, minio_repository, logger, notification_handler, metrics_controller + self.minio_repository = minio_repository + SientiaMonitoring.__init__( + self, + logger=logger, + notification_handler=notification_handler, + metrics_controller=metrics_controller, ) def close(self) -> None: @@ -115,7 +118,12 @@ class Gates(MinioManager): Close the gates activity and clean up resources. """ - MinioManager.close(self) + if self.minio_repository is not None: + try: + self.minio_repository.close() + finally: + self.minio_repository = None + SientiaMonitoring.shutdown(self) def __del__(self): self.close() @@ -155,7 +163,7 @@ class Gates(MinioManager): return policy, filter_config @activity.defn(name='input_gate') - async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: + def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Apply input data quality filters and validation. @@ -194,7 +202,7 @@ class Gates(MinioManager): filters = input_data['filters'] payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) path_priority = input_data['path_priority'] filter_output = [] @@ -214,7 +222,7 @@ class Gates(MinioManager): filter_output.append(policy) except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=f'INTPUT_GATE_ERROR__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', @@ -235,7 +243,7 @@ class Gates(MinioManager): return None, 0, '' @activity.defn(name='mlflow_response_gate') - async def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: + def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Validate MLFlow API response quality and integrity. @@ -279,7 +287,7 @@ class Gates(MinioManager): self.debug(f'Filters: {filters}', metadata) payload = MinioDataFramePayload.from_dict(raw_data) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) gate_type = input_data['type'] path_priority = input_data['path_priority'] @@ -298,7 +306,7 @@ class Gates(MinioManager): 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( + self.send_notification( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}', message=status.get('message', 'Unknown MLFlow API error'), @@ -308,7 +316,7 @@ class Gates(MinioManager): ) except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', @@ -329,7 +337,7 @@ class Gates(MinioManager): return None, 0, '' @activity.defn(name='mlflow_content_gate') - async def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: + def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Validate MLFlow prediction content quality and integrity. @@ -368,7 +376,7 @@ class Gates(MinioManager): filters = input_data['filters'] payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) gate_type = input_data['type'] path_priority = input_data['path_priority'] @@ -385,7 +393,7 @@ class Gates(MinioManager): try: if mlflow_content_filter_functions[fil](data, filter_config): filter_output.append(policy) - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}', message=f'Data not passed the content filter {fil}:{config}', @@ -395,7 +403,7 @@ class Gates(MinioManager): ) except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', @@ -471,7 +479,7 @@ class Gates(MinioManager): return policy_type, int(policy_value) @activity.defn(name='format_transformed_data') - async def format_transformed_data(self, input_data: dict[str, Any]) -> MinioDataFramePayload: + def format_transformed_data(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ Format transformed data for storage and export operations. @@ -507,7 +515,7 @@ class Gates(MinioManager): self.info('Formatting transformed data...', metadata) payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) data['timestamp'] = data.index data = data.reset_index(drop=True) @@ -515,7 +523,7 @@ class Gates(MinioManager): data = data.melt(id_vars='timestamp', var_name='variable', value_name='value') data['model_id'] = model_id - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe=data, minio_repo=self.minio_repository, model_name=input_data['model_name'], @@ -526,7 +534,7 @@ class Gates(MinioManager): ) @activity.defn(name='format_prediction') - async def format_prediction(self, input_data: dict[str, Any]) -> dict: + def format_prediction(self, input_data: dict[str, Any]) -> dict: """ Format prediction data according to configured storage policies. @@ -558,7 +566,7 @@ class Gates(MinioManager): self.info('Formatting prediction...', metadata) payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) # Create timestamp column from index and reset index data['timestamp'] = data.index @@ -607,7 +615,7 @@ class Gates(MinioManager): return data.to_dict() @activity.defn(name='format_default_prediction') - async def format_default_prediction(self, input_data: dict[str, Any]) -> dict: + def format_default_prediction(self, input_data: dict[str, Any]) -> dict: """ Create and format default prediction data for error conditions. @@ -652,7 +660,7 @@ class Gates(MinioManager): return data.to_dict() @activity.defn(name='format_retrain_report') - async def format_retrain_report(self, input_data: dict[str, Any]) -> dict: + def format_retrain_report(self, input_data: dict[str, Any]) -> dict: """ Format retrain report data for storage and audit trail maintenance. @@ -722,7 +730,7 @@ class Gates(MinioManager): return report.to_dict() @activity.defn(name='write_metrics') - async def write_metrics(self, input_data: dict[str, Any]): + def write_metrics(self, input_data: dict[str, Any]): """ Write prediction performance metrics to Prometheus monitoring system. @@ -759,19 +767,19 @@ class Gates(MinioManager): 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], } - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PREDICTIONS_WRITTEN_COUNT, tags=core_tags, ) - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PREDICTION_CONFIDENCE_MONITOR, method='set', tags=core_tags, value=prediction_confidence, ) - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PREDICTION_RESPONSE_TIME_MONITOR, method='observe', tags=core_tags, @@ -781,7 +789,7 @@ class Gates(MinioManager): for server_id, tags in opc_metrics.items(): for tag, response_time in tags.items(): if response_time is not None: - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR, method='observe', tags={ @@ -792,7 +800,7 @@ class Gates(MinioManager): value=response_time, ) - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.PREDICTION_OPC_WRITING_COUNT, tags={ **core_tags, diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index b995907..6cf311d 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -16,7 +16,8 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController - from sientia_do.repository.minio_repository import MinioRepository + from sientia_do.observability.sientia_monitoring import SientiaMonitoring + from sientia_do.repository.minio_repository_sync import MinioRepository from sientia_do.temporal.constants import ( DATETIME_FORMAT, DATETIME_FORMAT_MS_WITH_TZ, @@ -29,10 +30,9 @@ with workflow.unsafe.imports_passed_through(): from laborious.utils.dataframe_debug import build_dataframe_debug_message from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload - from laborious.utils.repository.minio_manager import MinioManager -class MLFlow(MinioManager): +class MLFlow(SientiaMonitoring): """ Temporal activities that talk to MLflow through ``SientiaMLflowRepository`` and ``SientiaModel`` wrappers. @@ -78,8 +78,12 @@ class MLFlow(MinioManager): None """ - MinioManager.__init__( - self, minio_repository, logger, notification_handler, metrics_controller + self.minio_repository = minio_repository + SientiaMonitoring.__init__( + self, + logger=logger, + notification_handler=notification_handler, + metrics_controller=metrics_controller, ) self.mlflow_repository = mlflow_repository self.plugin_store = plugin_store @@ -91,7 +95,12 @@ class MLFlow(MinioManager): Return: None """ - MinioManager.close(self) + if self.minio_repository is not None: + try: + self.minio_repository.close() + finally: + self.minio_repository = None + SientiaMonitoring.shutdown(self) def __del__(self): self.close() @@ -207,7 +216,7 @@ class MLFlow(MinioManager): return alias or self._DEFAULT_MODEL_ALIAS @activity.defn(name='request_transform') - async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload: + def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ Pivot long-format sensor rows, load the production wrapper, and run ``wrapper.transform``. @@ -228,7 +237,7 @@ class MLFlow(MinioManager): self.info('Transforming data...', metadata) payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) @@ -284,7 +293,7 @@ class MLFlow(MinioManager): self.info('Data transformed successfully', metadata) if not response_data.get('success', False): - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe=None, minio_repo=self.minio_repository, model_name=model_name, @@ -295,7 +304,7 @@ class MLFlow(MinioManager): logger=self.logger, ) - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe=response_data['content'], minio_repo=self.minio_repository, model_name=model_name, @@ -309,7 +318,7 @@ class MLFlow(MinioManager): ) @activity.defn(name='request_predict') - async def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload: + def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ Load the production wrapper and call ``wrapper.predict`` on the prepared feature frame. @@ -329,7 +338,7 @@ class MLFlow(MinioManager): self.info('Predicting data...', metadata) payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) @@ -390,7 +399,7 @@ class MLFlow(MinioManager): self.info('Data predicted successfully', metadata) if not response_data.get('success', False): - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe=None, minio_repo=self.minio_repository, model_name=model_name, @@ -401,7 +410,7 @@ class MLFlow(MinioManager): logger=self.logger, ) - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe=response_data['content'], minio_repo=self.minio_repository, model_name=model_name, @@ -415,7 +424,7 @@ class MLFlow(MinioManager): ) @activity.defn(name='retrain_model') - async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]: + def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Fit an updated wrapper from historical data, then log and register in MLflow. @@ -441,11 +450,11 @@ class MLFlow(MinioManager): try: payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='ERROR_LOADING_RETRAIN_DATA', message=f'Error loading retrain data: {e}', @@ -576,7 +585,7 @@ class MLFlow(MinioManager): } @activity.defn(name='update_production_model') - async def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]: + def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Point the ``production`` alias at the model version registered for the retrain run. @@ -622,7 +631,7 @@ class MLFlow(MinioManager): except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='UPDATE_PRODUCTION_MODEL_ERROR', message=f'Error updating production model {model_name}: {e}', @@ -634,7 +643,7 @@ class MLFlow(MinioManager): raise e @activity.defn(name='get_reference_data') - async def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None: + def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None: """ Download ``evaluation_data.csv`` from the MLflow run linked to ``production`` and parse it. diff --git a/laborious/activities/model_metrics.py b/laborious/activities/model_metrics.py index 79b140c..8b17ad1 100644 --- a/laborious/activities/model_metrics.py +++ b/laborious/activities/model_metrics.py @@ -7,14 +7,15 @@ with workflow.unsafe.imports_passed_through(): from typing import Any import numpy as np - from pandas import DataFrame, Index, to_datetime - from sientia.ModelAnalysis import ModelAnalysis + import pandas as pd + from pandas import DataFrame, Index, Series, to_datetime from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController from sientia_do.observability.sientia_monitoring import SientiaMonitoring - from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ + from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ + from sientia_model.analytics.model_analysis import ModelAnalysis from laborious import metrics from laborious.utils.dataframe_debug import build_dataframe_debug_message @@ -27,9 +28,12 @@ warnings.filterwarnings( class ModelMetrics(SientiaMonitoring): """ - Metrics activities for the Laborious system. + Metrics and statistical analysis activities for the Laborious pipeline. - This class provides activities for writing metrics to the Prometheus monitoring system. + This class centralizes drift/statistical computations and model-quality + aggregates used by scheduled workflows. Besides producing tabular outputs + for persistence, it also emits operational metrics (count, lag, error) + through ``SientiaMonitoring`` so execution health is observable in runtime. """ _MAX_DEBUG_DATAFRAME_ROWS = 100 @@ -44,7 +48,10 @@ class ModelMetrics(SientiaMonitoring): def close(self) -> None: """ - Close the model metrics activity and clean up resources. + Shutdown monitoring resources associated with model metrics activities. + + This is invoked during worker teardown to flush/close metric controller + internals and prevent dangling telemetry tasks. """ SientiaMonitoring.shutdown(self) @@ -69,7 +76,7 @@ class ModelMetrics(SientiaMonitoring): metadata, ) - async def get_drift_metrics( + def get_drift_metrics( self, reference_data: DataFrame, target_data: DataFrame, @@ -80,14 +87,23 @@ class ModelMetrics(SientiaMonitoring): metadata: dict[str, Any], ) -> DataFrame: """ - Calculate univariate drift metrics for a model. + Compute univariate and multivariate drift outputs and merge them into one dataframe. + + The method orchestrates three analysis stages (univariate drift, + multivariate drift, and dataframe projection), emitting lag/count/error + metrics for each stage independently so failures are attributable. + Args: - model_analysis (ModelAnalysis): Model analysis object - reference_data (DataFrame): Reference data - target_data (DataFrame): Target data - reference_columns (list[str]): Reference columns - drift_metrics (list[str]): Drift metrics - metadata (dict[str, Any]): Workflow execution metadata + - reference_data (DataFrame): Baseline dataset representing expected behavior. + - target_data (DataFrame): Current analysis dataset to compare against reference. + - target_name (str): Target column name used by ``ModelAnalysis`` config. + - reference_columns (Index): Feature columns evaluated for drift. + - drift_metrics (list[str]): Enabled univariate methods. + - chunk_period (str): Time bucket granularity used by analysis methods. + - metadata (dict[str, Any]): Workflow metadata for logs and notifications. + + Return: + DataFrame: Consolidated drift dataframe ready for downstream formatting/persistence. """ config = { @@ -118,12 +134,10 @@ class ModelMetrics(SientiaMonitoring): ) except Exception as e: self.error(f'Error detecting univariate drift: {e}', metadata) - await self.emit_metric( - metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels - ) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels) raise e - await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) + self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) core_labels = self.get_core_labels(metadata, operation_type='detect_multivariate_drift') start_time = time.time() @@ -137,12 +151,10 @@ class ModelMetrics(SientiaMonitoring): ) except Exception as e: self.error(f'Error detecting multivariate drift: {e}', metadata) - await self.emit_metric( - metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels - ) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels) raise e - await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) + self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) start_time = time.time() core_labels = self.get_core_labels(metadata, operation_type='get_drift_metrics_dataframe') @@ -153,19 +165,40 @@ class ModelMetrics(SientiaMonitoring): ) except Exception as e: self.error(f'Error getting drift metrics: {e}', metadata) - await self.emit_metric( - metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels - ) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels) raise e - await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) + self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels) + self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels) self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata) return drift_df + @staticmethod + def _to_naive_utc(series: Series) -> Series: + """ + Parse ``series`` as datetime and return a TZ-naive UTC copy. + + ``sientia_model.analytics.model_analysis.ModelAnalysis`` preserves the + timezone of the input dataframe in its outputs, while target rows + loaded from PostgreSQL come in with ``+00:00``. Forcing both sides of + a comparison to TZ-naive UTC keeps ``isin`` / ``floor`` operations + deterministic regardless of how the analyzer (or a test double) + constructs its timestamps. + + Args: + - series (Series): Input series containing datetime-parseable values. + + Return: + Series: Datetime64 series with ``tz=None`` representing UTC instants. + """ + parsed = to_datetime(series) + if getattr(parsed.dt, 'tz', None) is not None: + parsed = parsed.dt.tz_convert('UTC').dt.tz_localize(None) + return parsed + @activity.defn(name='calculate_drift') - async def calculate_drift(self, input_data: dict[str, Any]) -> list[dict]: + def calculate_drift(self, input_data: dict[str, Any]) -> list[dict]: """ Calculate drift metrics for a model. @@ -195,14 +228,17 @@ class ModelMetrics(SientiaMonitoring): target_data = target_data.pivot(index='timestamp', columns='variable', values='value') target_data['timestamp'] = target_data.index + # Keep timestamps as datetime: ModelAnalysis._chunk_dataframe relies on + # ``pd.Grouper(freq=...)`` which rejects string timestamp columns. target_data['timestamp'] = to_datetime(target_data['timestamp']) - target_data['timestamp'] = target_data['timestamp'].dt.strftime(DATETIME_FORMAT) target_data = target_data.reset_index(drop=True) target_data.dropna(inplace=True) if reference_raw_data is not None: self.info('Using reference data', metadata) reference_data = DataFrame(reference_raw_data) + if 'timestamp' in reference_data.columns: + reference_data['timestamp'] = to_datetime(reference_data['timestamp']) accurate = True else: # Get 30% first rows of target_data @@ -211,7 +247,7 @@ class ModelMetrics(SientiaMonitoring): reference_data = target_data.head(int(len(target_data) * 0.3)) accurate = False - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING', message='Using 30% first rows of target data as reference data', @@ -225,7 +261,7 @@ class ModelMetrics(SientiaMonitoring): ).columns try: - drift_df = await self.get_drift_metrics( + drift_df = self.get_drift_metrics( reference_data=reference_data, target_data=target_data, target_name=target_name, @@ -236,7 +272,7 @@ class ModelMetrics(SientiaMonitoring): ) except Exception as e: self.error(f'Error getting drift metrics: {e}', metadata) - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR', message=f'Error getting drift metrics: {e}', @@ -250,17 +286,13 @@ class ModelMetrics(SientiaMonitoring): self.warning('No drift metrics found', metadata) return [] - # Drop unnecessary columns - drift_df.drop(columns=['p_value'], inplace=True) - - # Extract timestamps only until minutes - if chunk_period == 'min': - target_timestamps = target_data['timestamp'].apply(lambda x: x[:16]) - else: - target_timestamps = target_data['timestamp'] - - # Drop rows where timestamp is not in target data, to avoid save drift from reference - drift_df = drift_df[drift_df['timestamp'].isin(target_timestamps)] + # Defense-in-depth: drop chunks whose floored timestamp does not appear + # in the analysis window. ``ModelAnalysis`` already chunks only over + # ``analysis_df`` so this only excludes rows injected by upstream + # callers that pre-merge reference data into the result. + target_floor = self._to_naive_utc(target_data['timestamp']).dt.floor(chunk_period) + drift_floor = self._to_naive_utc(drift_df['timestamp']).dt.floor(chunk_period) + drift_df = drift_df[drift_floor.isin(target_floor)] if drift_df.empty: self.warning( @@ -269,14 +301,21 @@ class ModelMetrics(SientiaMonitoring): ) return [] - # Rename columns to match database columns - drift_df.rename( + # Map ``sientia_model.analytics.model_analysis`` schema onto the drift + # table columns: ``metric -> method``, ``statistic -> value``, + # ``alert -> drift``, ``chunk_index -> chunk``, + # ``chunk_end_date -> timestamp_end``. ``p_value`` and + # ``chunk_start_date`` are not persisted. + drift_df = drift_df.rename( columns={ 'metric': 'method', 'statistic': 'value', - }, - inplace=True, + 'alert': 'drift', + 'chunk_index': 'chunk', + 'chunk_end_date': 'timestamp_end', + } ) + drift_df.drop(columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore') # Drop duplicates drift_df.drop_duplicates( @@ -286,16 +325,23 @@ class ModelMetrics(SientiaMonitoring): drift_df['model_id'] = model_id drift_df['accurate'] = accurate - drift_df['timestamp'] = to_datetime(drift_df['timestamp']) + drift_df['timestamp'] = self._to_naive_utc(drift_df['timestamp']) drift_df['timestamp'] = drift_df['timestamp'].dt.tz_localize('UTC') drift_df['timestamp'] = drift_df['timestamp'].dt.strftime(DATETIME_FORMAT_WITH_TZ) + # ``timestamp_end`` may carry nanosecond precision (beyond + # ``timestamptz`` microseconds), so serialize as ISO text for the + # ``text`` Postgres column. + drift_df['timestamp_end'] = drift_df['timestamp_end'].apply( + lambda value: pd.Timestamp(value).isoformat() if pd.notna(value) else None + ) + self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata) return drift_df.to_dict(orient='records') @activity.defn(name='calculate_simple_metrics') - async def calculate_simple_metrics(self, input_data: dict[str, Any]) -> list[dict]: + def calculate_simple_metrics(self, input_data: dict[str, Any]) -> list[dict]: """ Calculate simple metrics for a model. Metrics available are: - rmse diff --git a/laborious/activities/opc.py b/laborious/activities/opc.py index da36793..277d935 100644 --- a/laborious/activities/opc.py +++ b/laborious/activities/opc.py @@ -51,7 +51,7 @@ class OPC(SientiaMonitoring): self.opc_repository: dict[str, OpcRepository] = {} - async def init_opc(self): + def init_opc(self) -> None: """ Initialize OPC server connections and establish communication channels. @@ -59,58 +59,43 @@ class OPC(SientiaMonitoring): establish secure connections using certificate-based authentication. Each server connection is managed independently, and connection failures are reported through the notification system. - - The method performs the following operations: - 1. Creates OpcRepository instances for each configured server - 2. Establishes secure connections with certificate validation - 3. Reports connection success/failure through notifications - 4. Logs connection status for operational visibility - - Raises: - Exception: If OPC repository initialization fails or connection - establishment encounters critical errors - - Note: - Connection failures are logged and reported but do not prevent - the initialization of other OPC servers. Each server is handled - independently to ensure maximum availability. """ + self.logger.info('Initializing OPC servers...') for opc_id, server in self.opc_servers.items(): self.opc_repository[opc_id] = OpcRepository( - opc_id=server['id'], - server_name=server['server_name'], + opc_id=opc_id, url=server['url'], + server_name=server['server_name'], logger=self.logger, + notification_handler=self.notification_handler, + metrics_controller=self.metrics_controller, server_uri=server['server_uri'], cert_path=server['cert_path'], private_key_path=server['private_key_path'], server_cert_path=server['server_cert_path'], - notification_handler=self.notification_handler, - reconnection_interval=server['reconnection_interval'], - metrics_controller=self.metrics_controller, ) - is_connected, error_data = await self.opc_repository[opc_id].connect() - if not is_connected: - await self.send_notification_async( + ok, err = self.opc_repository[opc_id].connect() + if not ok: + self.send_notification( metadata={ 'model_id': '-', 'model_name': '-', 'workflow_name': '-', 'schedule_name': 'INITIALIZATION', }, - notification_id=error_data['notification_id'], - message=error_data['message'], - block=error_data['block'], - level=error_data.get('level', NotificationLevel.ERROR), - attachment_content=error_data.get('attachment_content', None), + notification_id=f'OPC_CONNECTION_ERROR_{server.get("id", opc_id)}', + message=err.get('message', 'Failed to connect to OPC server'), + block='opc_repository', + level=NotificationLevel.ERROR, + attachment_content=err.get('attachment_content', traceback.format_exc()), ) else: self.logger.info( f'OPC server {opc_id}:{server["server_name"]} connected successfully.' ) - async def write_data( + def write_data( self, server_id: str, tag: str, @@ -122,11 +107,6 @@ class OPC(SientiaMonitoring): """ Write data to a specific OPC server tag with comprehensive error handling. - This method provides a secure and reliable way to write data to OPC servers - with automatic error handling, notification integration, and detailed logging. - It validates server availability before attempting write operations and - provides comprehensive error reporting for operational monitoring. - Args: - server_id (str): The id of the OPC server. - tag (str): The tag to write to. @@ -135,15 +115,15 @@ class OPC(SientiaMonitoring): - tag_type (str): The tag type. Returns: - - bool: True if the data was written successfully, False otherwise. + - float | None: Response time in seconds if successful, None otherwise. """ try: - is_success, info_data = await self.opc_repository[server_id].write_data( + is_success, info_data = self.opc_repository[server_id].write_data( tag, data, data_type, self.logger, metadata ) if not is_success: - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=info_data['notification_id'], message=info_data['message'], @@ -155,7 +135,7 @@ class OPC(SientiaMonitoring): return info_data['response_time'] except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id=f'WRITE_OPC_{tag_type.upper()}_ERROR', message=f'Error writing data to OPC server: {e}', @@ -165,30 +145,25 @@ class OPC(SientiaMonitoring): ) raise e - async def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool: + def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool: """ - Validate that an OPC server is available and configured for write operations. + Validate that an OPC repository exists for the requested server identifier. - This method checks if the specified OPC server exists in the active - repository and is available for data writing operations. It provides - immediate feedback for server availability and logs validation failures - for operational monitoring. + This guard prevents write attempts against unknown/uninitialized servers. + When the server is missing, it emits an error notification with the list + of available repositories to help operators diagnose configuration drift. Args: - server_id (str): Unique identifier for the OPC server to validate - metadata (dict[str, Any]): Context metadata for logging and notifications + - server_id (str): OPC server identifier from workflow output config. + - metadata (dict[str, Any]): Workflow metadata used for logs/alerts. - Returns: - bool: True if server is available, False otherwise - - Note: - Server validation failures are automatically reported through the - notification system with detailed information about available servers. - This helps operators quickly identify configuration issues. + Return: + bool: ``True`` when the server repository is available; ``False`` otherwise. """ + if self.opc_repository.get(server_id) is None: message = f'OPC server {server_id} not found to perform write operation.' - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='OPC_SERVER_NOT_FOUND', message=message, @@ -199,7 +174,7 @@ class OPC(SientiaMonitoring): return False return True - async def manage_output_tags( + def manage_output_tags( self, server_id: str, config: dict[str, Any], @@ -207,37 +182,30 @@ class OPC(SientiaMonitoring): metadata: dict[str, Any], ) -> tuple[bool, dict[str, float | None]]: """ - Manage the writing of prediction and confidence data to OPC server tags. + Write prediction and confidence values for one OPC server configuration. - This method orchestrates the writing of multiple data types to OPC servers - based on configuration. It handles both prediction data and confidence - values independently, allowing for flexible tag configuration and - comprehensive error handling. - - The method supports two main tag types: - 1. Prediction tags: Write actual prediction values to configured OPC tags - 2. Confidence tags: Write confidence scores to separate OPC tags + The method iterates through optional ``prediction_tags`` and + ``confidence_tags``, performs synchronous writes for each tag, collects + per-tag response times, and returns an aggregate success flag + (all tags successful) with a metrics-friendly response map. Args: - server_id (str): Unique identifier for the target OPC server - config (dict[str, Any]): OPC tag configuration containing: - - prediction_tags (dict, optional): Prediction tag configurations - - confidence_tags (dict, optional): Confidence tag configurations - data (DataFrame): DataFrame containing prediction and confidence data - metadata (dict[str, Any]): Context metadata for logging and notifications - success (bool): Current success status to maintain across operations + - server_id (str): Target OPC server id. + - config (dict[str, Any]): Server output configuration containing optional + ``prediction_tags`` and ``confidence_tags`` sections. + - data (DataFrame): Prediction dataframe used as source values. + - metadata (dict[str, Any]): Workflow metadata for logging/notifications. - Returns: - tuple[bool, int]: (overall_success, total_tags_written) - - overall_success: True if all configured tags were written successfully - - total_tags_written: Count of successfully written tags + Return: + tuple[bool, dict[str, float | None]]: Global success flag and response-time + map per tag (``None`` for failed writes). """ response_times: dict[str, float | None] = {} if 'prediction_tags' in config: for tag, tag_config in config['prediction_tags'].items(): - response_time = await self.write_data( + response_time = self.write_data( server_id=server_id, tag=tag, data=data.head(1)['prediction'].values[0], @@ -254,7 +222,7 @@ class OPC(SientiaMonitoring): if 'confidence_tags' in config: for tag, tag_config in config['confidence_tags'].items(): - response_time = await self.write_data( + response_time = self.write_data( server_id=server_id, tag=tag, data=data.head(1)['prediction_confidence'].values[0], @@ -274,25 +242,25 @@ class OPC(SientiaMonitoring): return success, response_times @activity.defn(name='write_opc_data') - async def write_opc_data( + def write_opc_data( self, input_data: dict[str, Any] ) -> tuple[dict[Hashable, Any], dict[str, dict[str, float | None]]]: """ - Write prediction and confidence data to OPC servers. The two writing - operations are optional and independent of each other. + Execute OPC writes across all configured servers and collect per-tag metrics. + + For each server in ``opc_output_config``, this activity validates server + availability, writes enabled prediction/confidence tags, accumulates + response-time metrics, and then normalizes confidence/comments in the + returned prediction payload when at least one write fails. Args: - - input_data(dict[str, Any]): The input data. Contains the following keys: - - data(dict[str, Any]): The dataframe that contains the data to write - to the OPC servers. - - opc_output_config(dict[str, Any]): The OPC writing configuration. - The keys are the OPC server names and the values contain: - - prediction_tags(dict[str, Any]): The tags to write to the OPC servers. - - confidence_tags(dict[str, Any]): The tags to write to the OPC servers. - - Returns: - - dict[Any, Any]: The data that was written to the OPC servers. + - input_data (dict[str, Any]): Payload containing workflow metadata, data + to write, and ``opc_output_config`` server/tag definitions. + Return: + tuple[dict[Hashable, Any], dict[str, dict[str, float | None]]]: Updated + prediction payload dict and nested metrics + ``{server_id: {tag_name: response_time_or_none}}``. """ metadata = input_data['metadata'] self.info('Writing data to OPC servers...', metadata) @@ -305,19 +273,21 @@ class OPC(SientiaMonitoring): metrics: dict[str, dict[str, float | None]] = {} for server_id, config in opc_output_config.items(): - if not await self.validate_server(server_id, metadata): + if not self.validate_server(server_id, metadata): success = False continue - local_success, local_response_times = await self.manage_output_tags( + local_success, local_response_times = self.manage_output_tags( server_id, config, data, metadata ) metrics[server_id] = local_response_times local_count = len(local_response_times) success = success and local_success + n_pred = len(config.get('prediction_tags') or {}) + n_conf = len(config.get('confidence_tags') or {}) self.info( - f'Process completed for OPC server {server_id}: {local_count} of {len(config.get("prediction_tags", []))} prediction tags and {len(config.get("confidence_tags", []))} confidence tags', + f'Process completed for OPC server {server_id}: {local_count} of {n_pred} prediction tags and {n_conf} confidence tags', metadata, ) @@ -327,29 +297,16 @@ class OPC(SientiaMonitoring): self, data: DataFrame, success: bool, metadata: dict[str, Any] ) -> dict[Hashable, Any]: """ - Process prediction confidence based on OPC write operation success. - - This method updates the prediction confidence values in the DataFrame - based on the success status of OPC server write operations. If any - write operations failed, it sets the confidence to a predefined error - value to indicate data quality issues. - - The method implements a confidence degradation strategy: - - Success: Maintains original confidence values - - Failure: Sets confidence to error value for operational awareness + Apply fallback confidence/comment values when OPC writes are not fully successful. Args: - data (DataFrame): DataFrame containing prediction and confidence data - success (bool): Overall success status of OPC write operations - metadata (dict[str, Any]): Context metadata for logging and notifications + - data (DataFrame): Prediction dataframe to be returned to downstream steps. + - success (bool): Aggregate write status across all attempted OPC tags. + - metadata (dict[str, Any]): Workflow metadata used for debug logs. - Returns: - dict[Any, Any]: Processed data as a dictionary with updated confidence values - - Note: - The error confidence value (OPC_WRITTING_ERROR_CONFIDENCE = 12) is - used to indicate that data was not successfully exported to OPC servers. - This allows downstream systems to handle data quality appropriately. + Return: + dict[Hashable, Any]: Serialized dataframe dict with original values on success, + or downgraded confidence/comment fields on failure. """ message = 'Some data could not be written to OPC servers' @@ -367,25 +324,14 @@ class OPC(SientiaMonitoring): return data.to_dict() - async def close(self): + def close(self) -> None: """ - Gracefully shutdown all OPC server connections and cleanup resources. + Disconnect all tracked OPC repositories and clear in-memory references. - This method ensures proper cleanup of all active OPC server connections - by calling the disconnect method on each repository instance. It's - designed to be called during application shutdown to prevent resource - leaks and ensure clean termination. - - The method performs the following cleanup operations: - 1. Iterates through all active OPC repository connections - 2. Calls disconnect() on each repository instance - 3. Allows for graceful connection termination - 4. Prevents resource leaks and connection hanging - - Note: - This method should be called during application shutdown to ensure - proper cleanup. It handles all active connections regardless of - their current state and provides a clean shutdown experience. + This method should be called during worker shutdown to ensure every + synchronous OPC session is explicitly closed before process exit. """ + for opc in self.opc_repository.values(): - await opc.disconnect() + opc.disconnect() + self.opc_repository.clear() diff --git a/laborious/activities/storage.py b/laborious/activities/storage.py index d64cc52..6a996a5 100644 --- a/laborious/activities/storage.py +++ b/laborious/activities/storage.py @@ -1,7 +1,5 @@ from temporalio import activity, workflow -from laborious.utils.repository.minio_manager import MinioManager - with workflow.unsafe.imports_passed_through(): # Extend the Temporal Postgres activities for convenient query -> MinIO export import traceback @@ -13,8 +11,9 @@ with workflow.unsafe.imports_passed_through(): from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController - from sientia_do.repository.minio_repository import MinioRepository - from sientia_do.temporal.activities.postgres import Postgres + from sientia_do.observability.sientia_monitoring import SientiaMonitoring + from sientia_do.repository.minio_repository_sync import MinioRepository + from sientia_do.temporal.activities.postgres_sync import Postgres from sientia_do.temporal.constants import now from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload @@ -22,7 +21,7 @@ with workflow.unsafe.imports_passed_through(): _LOAD_QUERY_OFFLOAD_SKIP_KEYS = frozenset({'model_name', 'key_prefix', 'size_threshold_bytes'}) -class Storage(Postgres, MinioManager): +class Storage(Postgres, SientiaMonitoring): """ Extensions for Postgres activities with a helper to export query results directly to MinIO as Parquet and return the object name. @@ -60,14 +59,16 @@ class Storage(Postgres, MinioManager): metrics_controller=metrics_controller, ) - MinioManager.__init__( - self, minio_repository, logger, notification_handler, metrics_controller + self.minio_repository = minio_repository + SientiaMonitoring.__init__( + self, + logger=logger, + notification_handler=notification_handler, + metrics_controller=metrics_controller, ) @activity.defn(name='load_query_with_minio_offload') - async def load_query_with_minio_offload( - self, input_data: dict[str, Any] - ) -> MinioDataFramePayload: + def load_query_with_minio_offload(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ Run the custom SQL load, then return a MinIO-aware dataframe wire dict. @@ -88,7 +89,7 @@ class Storage(Postgres, MinioManager): metadata: dict = input_data.get('metadata', {}) model_name = input_data['model_name'] - rows = await self.load_custom_query( + rows = self.load_custom_query( input_data, ) if not rows: @@ -99,7 +100,7 @@ class Storage(Postgres, MinioManager): else: dataframe = pd.DataFrame(rows) - return await MinioDataFramePayload.from_dataframe( + return MinioDataFramePayload.from_dataframe( dataframe, minio_repo=self.minio_repository, workflow_metadata=metadata, @@ -109,15 +110,29 @@ class Storage(Postgres, MinioManager): ) @activity.defn(name='export_payload_to_postgres') - async def export_payload_to_postgres(self, input_data: dict[str, Any]) -> dict: + def export_payload_to_postgres(self, input_data: dict[str, Any]) -> dict: """ - Export a payload to PostgreSQL. + Resolve a MinIO-aware payload into a DataFrame and persist it into PostgreSQL. + + This activity accepts the serialized payload produced by previous steps + (inline dict or MinIO object reference), reconstructs the tabular data, + and delegates the final write to ``export_data_to_postgres`` using the + same input contract expected by the Postgres activity mixin. + + Args: + - input_data (dict[str, Any]): Activity input containing ``data`` as a + ``MinioDataFramePayload``-compatible dict plus database write options + (schema/table/on_conflict/metadata and related fields). + + Return: + dict: Result dictionary returned by ``export_data_to_postgres``, including + success status and optional write diagnostics. """ metadata = input_data.get('metadata') payload = MinioDataFramePayload.from_dict(input_data['data']) - data = await payload.retrieve(self.minio_repository, metadata) + data = payload.retrieve(self.minio_repository, metadata) - return await self.export_data_to_postgres( + return self.export_data_to_postgres( { **input_data, 'data': data, @@ -125,7 +140,7 @@ class Storage(Postgres, MinioManager): ) @activity.defn(name='cleanup_minio_objects_expired') - async def cleanup_minio_objects_expired(self, input_data: dict[str, Any]) -> dict[str, Any]: + def cleanup_minio_objects_expired(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Delete objects under the given prefixes that are older than the retention window. @@ -154,7 +169,7 @@ class Storage(Postgres, MinioManager): 'deleted_count': 0, } try: - keys = await self.minio_repository.list_objects( + keys = self.minio_repository.list_objects( prefix=prefix, recursive=True, metadata=metadata, @@ -166,7 +181,7 @@ class Storage(Postgres, MinioManager): continue if ts >= cutoff: continue - await self.minio_repository.delete_file( + self.minio_repository.delete_file( object_name=key, metadata=metadata, ) @@ -184,7 +199,7 @@ class Storage(Postgres, MinioManager): report['deleted_count'] += 1 except Exception as e: trace = traceback.format_exc() - await self.send_notification_async( + self.send_notification( metadata=metadata, notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED', message=f'Error cleaning up MinIO objects: {e}', @@ -201,9 +216,16 @@ class Storage(Postgres, MinioManager): return report def close(self) -> None: - """Close Storage resources (MinIO client and Postgres engine).""" - Postgres.close(self) - MinioManager.close(self) + """ + Shutdown Storage resources in deterministic order. - def __del__(self): - self.close() + The method first closes Postgres resources via ``Postgres.close`` (engine, + sessions, and monitoring hooks), then closes the optional MinIO repository + and clears the local reference to avoid accidental reuse after shutdown. + """ + Postgres.close(self) + if self.minio_repository is not None: + try: + self.minio_repository.close() + finally: + self.minio_repository = None diff --git a/laborious/utils/models/minio_dataframe_payload.py b/laborious/utils/models/minio_dataframe_payload.py index 0521f5e..9b75f8f 100644 --- a/laborious/utils/models/minio_dataframe_payload.py +++ b/laborious/utils/models/minio_dataframe_payload.py @@ -21,7 +21,7 @@ from typing import Any, Literal from pandas import DataFrame, read_parquet from sientia_do.observability.logger import Logger -from sientia_do.repository.minio_repository import MinioRepository +from sientia_do.repository.minio_repository_sync import MinioRepository from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, DATETIME_FORMAT_WITH_TZ, now # Keys that are part of the serialized wire format (not arbitrary metadata). @@ -89,7 +89,7 @@ class MinioDataFramePayload: metadata: dict[str, Any] | None = None, ) -> None: """ - Emit debug logs only when logger is provided + Emit a debug message only when a logger instance is available. Args: - logger (Logger | None): Logger instance used for debug messages @@ -182,7 +182,14 @@ class MinioDataFramePayload: def cleanup_prefix(self) -> str | None: """ - Return True if cleanup is enabled for this payload. + Return the MinIO prefix eligible for retention cleanup. + + Cleanup is only applicable when payload data was offloaded to MinIO + (``object_key`` present and inline ``data`` absent). Inline-only payloads + return ``None`` because there is no object tree to prune. + + Return: + str | None: Prefix used by cleanup listing, or ``None`` when cleanup does not apply. """ if self.object_key is not None and self.data is None: return self.object_prefix @@ -190,12 +197,19 @@ class MinioDataFramePayload: def has_data(self) -> bool: """ - Return True if the payload has some data internally or in MinIO. + Indicate whether the payload contains retrievable tabular content. + + A payload is considered non-empty when either inline ``data`` exists + (and is not an empty dict) or an ``object_key`` is available for MinIO + download. + + Return: + bool: ``True`` when data can be retrieved, ``False`` otherwise. """ return (self.data is not None and self.data != {}) or self.object_key is not None @classmethod - async def from_dataframe( + def from_dataframe( cls, dataframe: DataFrame | None, minio_repo: MinioRepository, @@ -274,7 +288,7 @@ class MinioDataFramePayload: dataframe.to_parquet(parquet_buffer, engine='pyarrow', index=True) file_bytes = parquet_buffer.getvalue() - upload_result = await minio_repo.upload_file( + upload_result = minio_repo.upload_file( file_bytes=file_bytes, relative_key=object_key, metadata=workflow_metadata, @@ -299,7 +313,7 @@ class MinioDataFramePayload: status=status, ) - async def retrieve( + def retrieve( self, minio_repo: MinioRepository, workflow_metadata: dict[str, Any] | None = None, @@ -336,7 +350,7 @@ class MinioDataFramePayload: f'MinioDataFramePayload.retrieve downloading object from MinIO: {self.object_key}', workflow_metadata, ) - file_bytes = await minio_repo.download_file( + file_bytes = minio_repo.download_file( object_name=self.object_key, metadata=workflow_metadata ) df = read_parquet(BytesIO(file_bytes)) diff --git a/laborious/utils/repository/minio_manager.py b/laborious/utils/repository/minio_manager.py deleted file mode 100644 index 17c85e6..0000000 --- a/laborious/utils/repository/minio_manager.py +++ /dev/null @@ -1,32 +0,0 @@ -from sientia_do.notifications.handlers import NotificationHandler -from sientia_do.observability.logger import Logger -from sientia_do.observability.metrics_controller import MetricsController -from sientia_do.observability.sientia_monitoring import SientiaMonitoring -from sientia_do.repository.minio_repository import MinioRepository - - -class MinioManager(SientiaMonitoring): - minio_repository: MinioRepository | None = None - - def __init__( - self, - minio_repository: MinioRepository | None = None, - logger: Logger | None = None, - notification_handler: NotificationHandler | None = None, - metrics_controller: MetricsController | None = None, - ): - if self.minio_repository is None: - self.minio_repository = minio_repository - SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) - - def close(self) -> None: - """ - Close the MinioManager and clean up resources. - """ - if self.minio_repository is not None: - try: - self.minio_repository.close() - finally: - self.minio_repository = None - - SientiaMonitoring.shutdown(self) diff --git a/laborious/utils/repository/opc_repository.py b/laborious/utils/repository/opc_repository.py index 385b256..2e797f3 100644 --- a/laborious/utils/repository/opc_repository.py +++ b/laborious/utils/repository/opc_repository.py @@ -1,4 +1,10 @@ -import asyncio +""" +Synchronous OPC UA client repository using python-opcua (opcua package). + +Connects to OPC UA servers, optionally configures Basic256 security, validates sessions, +and writes node values with typed variants and Prometheus-compatible metrics. +""" + import json import time import traceback @@ -6,9 +12,8 @@ from datetime import datetime from pathlib import Path from typing import Any -from asyncua import Client -from asyncua.crypto.security_policies import SecurityPolicyBasic256 -from asyncua.ua import DataValue, Variant, VariantType +from opcua import Client, ua +from opcua.crypto import security_policies from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger @@ -20,28 +25,38 @@ from laborious import metrics data_type_map = { 'float': { 'converter': float, - 'opc_type': VariantType.Float, + 'opc_type': ua.VariantType.Float, }, 'double': { 'converter': float, - 'opc_type': VariantType.Double, + 'opc_type': ua.VariantType.Double, }, 'int': { 'converter': int, - 'opc_type': VariantType.Int32, + 'opc_type': ua.VariantType.Int32, }, 'bool': { 'converter': bool, - 'opc_type': VariantType.Boolean, + 'opc_type': ua.VariantType.Boolean, }, 'str': { 'converter': str, - 'opc_type': VariantType.String, + 'opc_type': ua.VariantType.String, }, } class OpcRepository(SientiaMonitoring): + """ + Synchronous OPC UA repository for connect/disconnect and typed writes. + + Attributes: + url: OPC UA endpoint URL. + id: Server identifier used in metrics and notifications. + server_name: Human-readable server name for labels. + client: Active opcua.Client instance while connected. + """ + def __init__( self, opc_id: str, @@ -66,10 +81,10 @@ class OpcRepository(SientiaMonitoring): self.logger = logger self.error_count = 0 self.reconnection_interval = reconnection_interval - self.last_reconnection_time: None | datetime = None + self.last_reconnection_time: datetime | None = None self.disconnection_interval = 10.0 self.notification_handler = notification_handler - self.client: None | Client = None + self.client: Client | None = None SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) @@ -80,24 +95,12 @@ class OpcRepository(SientiaMonitoring): 'schedule_name': '-', } - async def set_security(self): + def set_security(self) -> None: """ - Configures the security settings for the OPC UA client. - This method sets up the security policy, certificates, and timeouts - required for establishing a secure connection with the OPC UA server. + Configure Basic256 security policy, certificates, and long channel/session timeouts. + Raises: - ValueError: If either the certificate path or private key path is not provided. - Attributes: - - cert_path (str): Path to the client's certificate file. - - private_key_path (str): Path to the client's private key file. - - server_cert_path (str, optional): Path to the server's certificate file. - - server_uri (str): The URI of the server to be used as the application URI. - - client (opcua.Client): The OPC UA client instance. - - logger (logging.Logger): Logger instance for logging information. - Security Settings: - - Security Policy: Basic256 - - Secure Channel Timeout: 10,000,000 ms - - Session Timeout: 10,000,000 ms + ValueError: If certificate paths are missing or client is not initialized. """ if self.cert_path is None or self.private_key_path is None: @@ -114,26 +117,24 @@ class OpcRepository(SientiaMonitoring): self.client.application_uri = self.server_uri self.logger.custom_info('Setting security...', self.metadata) - await self.client.set_security( - SecurityPolicyBasic256, - certificate=str(cert), - private_key=str(private_key), - server_certificate=str(server_cert) if server_cert else None, + self.client.set_security( + security_policies.SecurityPolicyBasic256, + str(cert), + str(private_key), + str(server_cert) if server_cert else None, ) self.client.secure_channel_timeout = 10000000 self.client.session_timeout = 10000000 - async def connect(self) -> tuple[bool, dict[str, Any]]: + def connect(self) -> tuple[bool, dict[str, Any]]: """ - Establishes a connection to the OPC server. - This method initializes the OPC client using the provided URL and - sets up security if a certificate path is specified. It then - attempts to connect to the server and logs the connection status. - Raises: - Exception: If the connection to the OPC server fails. + Create the synchronous client, optionally apply security, and connect to the server. + + Return: + tuple[bool, dict[str, Any]]: Success flag and error payload when False. """ - self.client = Client(self.url, timeout=10, watchdog_intervall=3600000) # type: ignore[attr-defined] + self.client = Client(self.url, timeout=10) self.client.name = self.pod_id self.client.application_name = self.pod_id @@ -142,32 +143,25 @@ class OpcRepository(SientiaMonitoring): self.client.product_uri = pod_uri if self.cert_path: - await self.set_security() + self.set_security() self.logger.custom_info( f'Starting connection to OPC server {self.id}:{self.server_name}...', self.metadata ) - return await self.try_connect() + return self.try_connect() - async def try_connect(self) -> tuple[bool, dict[str, Any]]: + def try_connect(self) -> tuple[bool, dict[str, Any]]: """ - Attempt to establish connection to the OPC server. + Perform the TCP/session handshake and emit connection metrics. - This method performs the actual connection attempt to the OPC server - and handles connection failures with comprehensive error reporting. - It updates reconnection timing and provides detailed error information - for operational monitoring and debugging. - - Returns: - tuple[bool, dict[str, Any]]: Connection result - - bool: True if connection successful, False otherwise - - dict: Error information if connection failed + Return: + tuple[bool, dict[str, Any]]: Success flag and structured error when False. """ tags = { 'pod_id': self.pod_id, 'server_name': self.server_name, } - await self.emit_metric(metrics.OPC_CONNECTIONS_TOTAL, tags) + self.emit_metric_sync(metrics.OPC_CONNECTIONS_TOTAL, tags) try: self.last_reconnection_time = datetime.now() if self.client is None: @@ -177,9 +171,9 @@ class OpcRepository(SientiaMonitoring): 'block': 'opc_repository', 'level': NotificationLevel.ERROR, } - await self.client.connect() + self.client.connect() - await self.emit_metric( + self.emit_metric_sync( metric_object=metrics.OPC_CONNECTION_STATUS, method='set', tags={ @@ -191,12 +185,12 @@ class OpcRepository(SientiaMonitoring): return True, {} except Exception as e: - await self.disconnect() + self.disconnect() trace = traceback.format_exc() self.logger.custom_error(trace, self.metadata) - await self.emit_metric(metrics.OPC_CONNECTIONS_FAILED, tags) + self.emit_metric_sync(metrics.OPC_CONNECTIONS_FAILED, tags) return False, { 'notification_id': f'OPC_CONNECTION_ERROR_{self.id}', @@ -206,21 +200,27 @@ class OpcRepository(SientiaMonitoring): 'attachment_content': trace, } - async def disconnection_fallback(self) -> list: + def disconnection_fallback(self) -> list[dict[str, Any]]: """ - Tries 5 times to disconnect from the OPC UA server, with a delay of 100ms x try. + Retry disconnect up to five times with linear backoff. + + Return: + list[dict[str, Any]]: Empty on success, otherwise error records per attempt. """ assert self.client is not None error_stack = [] for i in range(5): try: - self.logger.info(f'Disconnecting from OPC UA server, attempt {i + 1} of 5') - await self.client.disconnect() + self.logger.custom_info( + f'Disconnecting from OPC UA server, attempt {i + 1} of 5', self.metadata + ) + self.client.disconnect() return [] except Exception as e: - self.logger.error( - f'Failed to disconnect from OPC UA server in attempt {i + 1} of 5: {e}' + self.logger.custom_error( + f'Failed to disconnect from OPC UA server in attempt {i + 1} of 5: {e}', + self.metadata, ) error_stack.append( { @@ -229,23 +229,20 @@ class OpcRepository(SientiaMonitoring): 'traceback': traceback.format_exc(), } ) - await asyncio.sleep(self.disconnection_interval * i) + time.sleep(self.disconnection_interval * i) return error_stack - async def disconnect(self): + def disconnect(self) -> None: + """ + Tear down the UA session and reset connection metrics. """ - Gracefully disconnect from the OPC server. - This method safely terminates the connection to the OPC server - and cleans up client resources. It handles disconnection errors - gracefully and ensures proper resource cleanup. - """ if self.client is None: return - errors = await self.disconnection_fallback() + errors = self.disconnection_fallback() if errors: - await self.send_notification_async( + self.send_notification( metadata=self.metadata, notification_id=f'OPC_DISCONNECTION_ERROR_{self.id}', message='Failed to disconnect from OPC server in 5 attempts.', @@ -254,8 +251,10 @@ class OpcRepository(SientiaMonitoring): attachment_content=json.dumps(errors, indent=4), ) else: - self.logger.warning(f'Disconnected from OPC server {self.id} successfully') - await self.emit_metric( + self.logger.custom_warning( + f'Disconnected from OPC server {self.id} successfully', self.metadata + ) + self.emit_metric_sync( metric_object=metrics.OPC_CONNECTION_STATUS, method='set', tags={ @@ -268,77 +267,52 @@ class OpcRepository(SientiaMonitoring): self.client = None - async def validate_connection(self) -> tuple[bool, dict[str, Any]]: + def _session_alive(self) -> bool: """ - Validate and maintain OPC server connection health. + Best-effort check that the synchronous client still has a working session. - This method performs comprehensive connection validation and - implements automatic reconnection logic for production reliability. - It handles various connection states and implements intelligent - reconnection strategies with error counting and timing controls. - - Connection Validation: - 1. Checks client existence and connection state - 2. Implements error counting with automatic disconnection - 3. Enforces reconnection timing windows - 4. Provides detailed error reporting and notifications - - Reconnection Strategy: - - Error Count Threshold: Disconnects after 5 consecutive errors - - Reconnection Window: Enforces minimum intervals between attempts - - Automatic Recovery: Attempts reconnection when conditions allow - - State Monitoring: Continuously monitors connection health - - Args: - None - - Returns: - tuple[bool, dict[str, Any]]: Connection validation result - - bool: True if connection is healthy, False otherwise - - dict: Error information if validation fails + Return: + bool: True if a root browse succeeds, False otherwise. """ + if self.client is None: - return await self.connect() - - # if self.error_count > 5: # NOSONAR - # self.logger.custom_warning( - # f'OPC server {self.id} will be disconnected due to multiple errors', self.metadata - # ) - # try: - # await self.disconnect() - # except Exception as e: - # trace = traceback.format_exc() - # self.logger.custom_error( - # f'Failed to disconnect from OPC server: {e}', self.metadata - # ) - # self.logger.custom_error(trace, self.metadata) - # self.logger.custom_info( - # f'Attempting to reconnect to OPC server {self.id}...', self.metadata - # ) - # return await self.connect() - - # Check if client is connected using asyncua's connection state + return False try: - if ( - self.client.uaclient.protocol is None - or self.client.uaclient.protocol.state == 'closed' - ): - # OPC server is not connected + self.client.get_root_node() + return True + except Exception: + return False + + def validate_connection(self) -> tuple[bool, dict[str, Any]]: + """ + Ensure the UA session is usable; reconnect when outside the backoff window. + + Return: + tuple[bool, dict[str, Any]]: Whether the session is ready and optional error payload. + """ + + if self.client is None: + return self.connect() + + try: + if not self._session_alive(): self.logger.custom_error(f'OPC server {self.id} is not connected', self.metadata) if ( self.last_reconnection_time is None or (datetime.now() - self.last_reconnection_time).total_seconds() > self.reconnection_interval ): - await self.disconnect() + self.disconnect() self.logger.custom_info( f'Trying to reconnect to OPC server {self.id}...', self.metadata ) - return await self.connect() + return self.connect() return False, { 'notification_id': f'OPC_CONNECTION_AWAITING_RECONNECTION_WINDOW_{self.id}', - 'message': f'OPC server {self.id} is not connected, waiting for next reconnection window...', + 'message': ( + f'OPC server {self.id} is not connected, waiting for next reconnection window...' + ), 'block': 'opc_repository', 'level': NotificationLevel.WARNING, } @@ -355,39 +329,24 @@ class OpcRepository(SientiaMonitoring): 'attachment_content': trace, } - async def write_data( + def write_data( self, node: str, value: Any, data_type: str, logger: Logger, metadata: dict[str, Any] ) -> tuple[bool, dict[str, Any]]: """ - Write data to OPC server with comprehensive validation and monitoring. - - This method provides secure and reliable data writing to OPC servers - with automatic connection validation, data type conversion, and - comprehensive error handling. It implements performance monitoring - and metrics collection for operational visibility. - - Data Writing Process: - 1. Connection validation and automatic reconnection - 2. Node validation and error handling - 3. Data type conversion and validation - 4. OPC data writing with timestamp - 5. Performance metrics collection - 6. Error handling and notification + Write a typed value to an OPC UA node after validating connectivity. Args: - node (str): OPC node identifier to write data to - value (Any): Data value to write to the OPC node - data_type (str): Data type for OPC conversion - logger (Logger): Logger instance for operation logging - metadata (dict[str, Any]): Context metadata for logging and metrics + node: Node id string accepted by opcua Client.get_node. + value: Scalar value to encode. + data_type: Key into ``data_type_map`` (e.g. float, str). + logger: Caller logger for per-write traces. + metadata: Workflow metadata for error context. - Returns: - tuple[bool, dict[str, Any]]: Write operation result - - bool: True if write successful, False otherwise - - dict: Error information if write failed + Return: + tuple[bool, dict[str, Any]]: Success flag and either ``response_time`` or error fields. """ - is_connected, error = await self.validate_connection() + is_connected, error = self.validate_connection() if not is_connected: return False, error @@ -395,8 +354,8 @@ class OpcRepository(SientiaMonitoring): start_time = time.time() try: - # ignored because self.validate_connection is called before, so we know self.client is not None - node_obj = self.client.get_node(node) # type: ignore[union-attr] + assert self.client is not None + node_obj = self.client.get_node(node) except Exception as e: trace = traceback.format_exc() logger.custom_error(trace, metadata.get('schedule_name', 'N/A')) @@ -419,16 +378,10 @@ class OpcRepository(SientiaMonitoring): data = data_type_map[data_type]['converter'](value) logger.custom_info(f'Writing {data} - {type(data)} to {node}', metadata) - # now = datetime.now() # NOSONAR - ua_data = DataValue( - Variant(data, data_type_map[data_type]['opc_type']), - # SourceTimestamp=DateTime( # NOSONAR - # now.year, now.month, now.day, now.hour, now.minute, now.second, now.microsecond # NOSONAR - # ), # NOSONAR - ) + variant_type = data_type_map[data_type]['opc_type'] try: - await node_obj.write_value(ua_data) + node_obj.set_value(data, variant_type) end_time = time.time() response_time = end_time - start_time diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index 06ceecf..2597864 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -162,7 +162,7 @@ async def main(): ) logger.custom_info('Initializing OPC...', metadata) - await activities.init_opc() + activities.init_opc() logger.custom_info(f'Starting SDK Metrics Server on port {SDK_METRICS_PORT}...', metadata) @@ -265,7 +265,7 @@ async def main(): exit_code = 1 finally: notification_handler.shutdown() - await activities.shutdown() + activities.shutdown() metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN sys.exit(exit_code) diff --git a/laborious/workflows/sub_workflows/format_and_export_prediction.py b/laborious/workflows/sub_workflows/format_and_export_prediction.py index 68552b2..cf950ce 100644 --- a/laborious/workflows/sub_workflows/format_and_export_prediction.py +++ b/laborious/workflows/sub_workflows/format_and_export_prediction.py @@ -161,7 +161,7 @@ class FormatAndExportPrediction: write_transformed_handler = None - opc_metrics = {} + opc_metrics: dict[str, dict[str, float | None]] = {} # write to pi web api if pi_web_api_output_config: diff --git a/laborious/workflows/sub_workflows/prediction_process.py b/laborious/workflows/sub_workflows/prediction_process.py index da4125b..0f86fa1 100644 --- a/laborious/workflows/sub_workflows/prediction_process.py +++ b/laborious/workflows/sub_workflows/prediction_process.py @@ -250,7 +250,7 @@ class PredictionProcess: async def path_flag_handler( self, data: dict[str, Any], - path_flag: str, + path_flag: str | None, input_data: dict, confidence: int, last_timestamp: str, diff --git a/requirements-light.txt b/requirements-light.txt deleted file mode 100644 index f105052..0000000 --- a/requirements-light.txt +++ /dev/null @@ -1,12 +0,0 @@ -temporalio -psycopg2-binary -sqlalchemy -asyncua -redis -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 -prometheus-client -botocore -boto3 -s3fs -pyarrow -mlflow \ No newline at end of file diff --git a/requirements-local.txt b/requirements-local.txt index 0cc6e78..8fbf9b2 100644 --- a/requirements-local.txt +++ b/requirements-local.txt @@ -1,12 +1,13 @@ temporalio psycopg2-binary sqlalchemy -asyncua +opcua 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 +#git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 +/home/grezewave/Documents/projects/sientia/sientia-dataops-library +#git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.3 /home/grezewave/Documents/projects/sientia/sientia-model-library -prometheus-client +prometheus-client botocore boto3 s3fs diff --git a/requirements.txt b/requirements.txt index 2b9206f..c848080 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,7 @@ temporalio psycopg2-binary sqlalchemy -asyncua +opcua redis sientia_do==1.12.0 sientia_model==0.8.2 diff --git a/tests/conftest.py b/tests/conftest.py index 7671dce..dd8f543 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,6 +2,19 @@ import os import sys from unittest.mock import MagicMock +from sientia_do.temporal.activities.postgres_sync import Postgres + + +def _noop_postgres_del(_self): + """ + Unit tests use MagicMock metrics controllers; postgres_sync.Postgres.__del__ calls + close() during GC and triggers async shutdown. Explicit ``close()`` is covered in tests. + """ + return None + + +Postgres.__del__ = _noop_postgres_del # type: ignore[method-assign] + # The production code converts SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES to int at import-time. # Tests must set it to a valid integer string to avoid import errors. os.environ.setdefault('SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES', '1') diff --git a/tests/laborious/activities/test_activities.py b/tests/laborious/activities/test_activities.py index 0a025bd..4f39f9b 100644 --- a/tests/laborious/activities/test_activities.py +++ b/tests/laborious/activities/test_activities.py @@ -1,6 +1,4 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, patch - -from pytest import mark +from unittest.mock import ANY, MagicMock, patch from laborious.activities.activities import Activities from laborious.activities.api import API @@ -156,7 +154,6 @@ def test___init__( ) -@mark.asyncio @patch('laborious.activities.activities.Storage') @patch('laborious.activities.activities.MLFlow') @patch('laborious.activities.activities.OPC') @@ -164,7 +161,7 @@ def test___init__( @patch('laborious.activities.activities.ModelMetrics') @patch('laborious.activities.activities.API') @patch('laborious.activities.activities.MinioRepository') -async def test_shutdown( +def test_shutdown( _mock_minio_repository, mock_api_init, mock_model_metrics_init, @@ -173,7 +170,7 @@ async def test_shutdown( mock_mlflow_init, mock_storage_init, ): - mock_opc_init.close = AsyncMock() + mock_opc_init.close = MagicMock() postgres_config = { 'host': 'localhost', 'port': 5432, @@ -222,7 +219,7 @@ async def test_shutdown( mlflow_repository=mlflow_repository, ) - await activities.shutdown() + activities.shutdown() mock_opc_init.close.assert_called_once() mock_storage_init.close.assert_called_once() mock_mlflow_init.close.assert_called_once() diff --git a/tests/laborious/activities/test_api.py b/tests/laborious/activities/test_api.py index d281eb1..be5ff84 100644 --- a/tests/laborious/activities/test_api.py +++ b/tests/laborious/activities/test_api.py @@ -1,7 +1,6 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, MagicMock, call, patch -import pytest_asyncio -from pytest import fixture, mark +from pytest import fixture from sientia_do.notifications.models import NotificationLevel from laborious.activities.api import API, PI_WEB_API_PREDICTION_ERROR_CONFIDENCE @@ -71,7 +70,7 @@ def test_get_pi_web_api_core_labels_without_operation_type(mock_pi_web_api_clien auth_token='test_token', logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) with patch.object( SientiaMonitoring, @@ -105,7 +104,7 @@ def test_get_pi_web_api_core_labels_with_operation_type(mock_pi_web_api_client): auth_token='test_token', logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) with patch.object( SientiaMonitoring, @@ -132,17 +131,17 @@ def test__init__(): auth_token='test_token', logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) assert api.pi_web_api_client is not None -@pytest_asyncio.fixture +@fixture @patch('laborious.activities.api.PIWebAPIClient') def api(mock_pi_web_api_client): mock_client = MagicMock() - mock_client.write_value = AsyncMock() + mock_client.write_value = MagicMock() mock_client.close = MagicMock() mock_client.base_url = 'https://test-pi-server.com' mock_pi_web_api_client.return_value = mock_client @@ -153,12 +152,12 @@ def api(mock_pi_web_api_client): auth_token='test_token', logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) - api_instance.send_notification_async = AsyncMock() + api_instance.send_notification = MagicMock() api_instance.info = MagicMock() api_instance.error = MagicMock() - api_instance.emit_metric = AsyncMock() + api_instance.emit_metric_sync = MagicMock() api_instance.get_core_labels = MagicMock( return_value={ 'pod_id': 'test_pod', @@ -170,9 +169,8 @@ def api(mock_pi_web_api_client): return api_instance -@mark.asyncio @patch('laborious.activities.api.DataFrame') -async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data): +def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data): input_data = { **base_input_data, 'pi_web_api_output_config': { @@ -190,7 +188,7 @@ async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_dat [{'WebId': 'web_id_3', 'Errors': []}, {'WebId': 'web_id_4', 'Errors': []}], ] - result = await api.write_pi_web_api_data(input_data) + result = api.write_pi_web_api_data(input_data) api.pi_web_api_client.write_value.assert_has_calls( [ @@ -220,9 +218,8 @@ async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_dat } -@mark.asyncio @patch('laborious.activities.api.DataFrame') -async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data): +def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data): mock_dataframe.return_value = _create_mock_dataframe( { 'prediction': [0.75], @@ -233,9 +230,9 @@ async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_ api.pi_web_api_client.write_value.side_effect = Exception('Prediction write failed') - result = await api.write_pi_web_api_data(base_input_data) + result = api.write_pi_web_api_data(base_input_data) - api.send_notification_async.assert_called_once_with( + api.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR', message="Error writing prediction data to PI Web API: Prediction write failed\n Tags: {'tag1': 'web_id_1'}", @@ -248,9 +245,8 @@ async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_ assert api.pi_web_api_client.write_value.call_count == 1 -@mark.asyncio @patch('laborious.activities.api.DataFrame') -async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data): +def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data): mock_dataframe.return_value = _create_mock_dataframe() # First call succeeds, second fails @@ -259,9 +255,9 @@ async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_ Exception('Confidence write failed'), ] - result = await api.write_pi_web_api_data(base_input_data) + result = api.write_pi_web_api_data(base_input_data) - api.send_notification_async.assert_called_once_with( + api.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR', message="Error writing confidence data to PI Web API: Confidence write failed\n Tags: {'tag2': 'web_id_2'}", @@ -278,9 +274,8 @@ async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_ assert api.pi_web_api_client.write_value.call_count == 2 -@mark.asyncio @patch('laborious.activities.api.DataFrame') -async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data): +def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data): input_data = { **base_input_data, 'pi_web_api_output_config': { @@ -298,7 +293,7 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_ [], ] - result = await api.write_pi_web_api_data(input_data) + result = api.write_pi_web_api_data(input_data) api.pi_web_api_client.write_value.assert_has_calls( [ @@ -328,9 +323,8 @@ 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( +def test_write_pi_web_api_data_updates_confidence_and_comments( mock_dataframe, api, base_input_data ): mock_dataframe.return_value = _create_mock_dataframe() @@ -341,23 +335,23 @@ async def test_write_pi_web_api_data_updates_confidence_and_comments( with patch.object( api, 'process_pi_web_api_response', - new=AsyncMock(side_effect=[(0.33, 'PI warning'), (0, '')]), + new=MagicMock(side_effect=[(0.33, 'PI warning'), (0, '')]), ) as process_mock: - result = await api.write_pi_web_api_data(base_input_data) + result = api.write_pi_web_api_data(base_input_data) - assert process_mock.await_count == 2 + assert process_mock.call_count == 2 assert result is not None -@mark.asyncio -async def test_close(api): +@patch('laborious.activities.api.SientiaMonitoring.shutdown') +def test_close(mock_shutdown, api): api.close() api.pi_web_api_client.close.assert_called_once() + mock_shutdown.assert_called_once_with(api) -@mark.asyncio -async def test_process_pi_web_api_response_success(api): +def test_process_pi_web_api_response_success(api): """Test successful processing of PI Web API response with all tags written.""" response_data = [ {'WebId': 'web_id_1', 'Errors': []}, @@ -371,7 +365,7 @@ async def test_process_pi_web_api_response_success(api): 'workflow_name': 'test_workflow', } - confidence, message = await api.process_pi_web_api_response( + confidence, message = api.process_pi_web_api_response( response_data=response_data, tags=tags, core_labels=core_labels, @@ -380,9 +374,9 @@ async def test_process_pi_web_api_response_success(api): assert confidence == 0 assert message == '' - assert api.emit_metric.call_count == 2 - # Verify that emit_metric was called with correct tags structure - call_args_list = api.emit_metric.call_args_list + assert api.emit_metric_sync.call_count == 2 + # Verify that emit_metric_sync was called with correct tags structure + call_args_list = api.emit_metric_sync.call_args_list assert len(call_args_list) == 2 # Check that all calls include core_labels and tag_name for call_args in call_args_list: @@ -390,8 +384,7 @@ async def test_process_pi_web_api_response_success(api): assert call_args.kwargs['tags']['tag_name'] in ['tag1', 'tag2'] -@mark.asyncio -async def test_process_pi_web_api_response_with_errors(api): +def test_process_pi_web_api_response_with_errors(api): """Test processing response with errors in some tags.""" response_data = [ {'WebId': 'web_id_1', 'Errors': ['Error writing tag']}, @@ -405,7 +398,7 @@ async def test_process_pi_web_api_response_with_errors(api): 'workflow_name': 'test_workflow', } - confidence, message = await api.process_pi_web_api_response( + confidence, message = api.process_pi_web_api_response( response_data=response_data, tags=tags, core_labels=core_labels, @@ -417,11 +410,10 @@ async def test_process_pi_web_api_response_with_errors(api): message == "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag2'] tags were written." ) - assert api.emit_metric.call_count == 2 + assert api.emit_metric_sync.call_count == 2 -@mark.asyncio -async def test_process_pi_web_api_response_missing_tags(api): +def test_process_pi_web_api_response_missing_tags(api): """Test processing response when number of written tags doesn't match expected.""" response_data = [ {'WebId': 'web_id_1', 'Errors': []}, @@ -434,7 +426,7 @@ async def test_process_pi_web_api_response_missing_tags(api): 'workflow_name': 'test_workflow', } - confidence, message = await api.process_pi_web_api_response( + confidence, message = api.process_pi_web_api_response( response_data=response_data, tags=tags, core_labels=core_labels, @@ -446,14 +438,13 @@ async def test_process_pi_web_api_response_missing_tags(api): message == "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag1'] tags were written." ) - api.send_notification_async.assert_called_once() - call_args = api.send_notification_async.call_args + api.send_notification.assert_called_once() + call_args = api.send_notification.call_args assert call_args.kwargs['notification_id'] == 'WRITE_PI_WEB_API_PREDICTION_ERROR' assert call_args.kwargs['level'] == NotificationLevel.ERROR -@mark.asyncio -async def test_process_pi_web_api_response_missing_webid(api): +def test_process_pi_web_api_response_missing_webid(api): """Test processing response when WebId is missing in response item.""" response_data = [ {'Errors': []}, @@ -467,7 +458,7 @@ async def test_process_pi_web_api_response_missing_webid(api): 'workflow_name': 'test_workflow', } - confidence, message = await api.process_pi_web_api_response( + confidence, message = api.process_pi_web_api_response( response_data=response_data, tags=tags, core_labels=core_labels, @@ -482,8 +473,7 @@ async def test_process_pi_web_api_response_missing_webid(api): api.error.assert_any_call('The response did not contain some WebIds', metadata['metadata']) -@mark.asyncio -async def test_process_pi_web_api_response_missing_tag_name(api): +def test_process_pi_web_api_response_missing_tag_name(api): """Test processing response when tag name is not found for WebId.""" response_data = [ {'WebId': 'unknown_web_id', 'Errors': []}, @@ -496,7 +486,7 @@ async def test_process_pi_web_api_response_missing_tag_name(api): 'workflow_name': 'test_workflow', } - confidence, message = await api.process_pi_web_api_response( + confidence, message = api.process_pi_web_api_response( response_data=response_data, tags=tags, core_labels=core_labels, diff --git a/tests/laborious/activities/test_gates.py b/tests/laborious/activities/test_gates.py index c144068..2ea921d 100644 --- a/tests/laborious/activities/test_gates.py +++ b/tests/laborious/activities/test_gates.py @@ -1,8 +1,9 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, MagicMock, call, patch from pandas import DataFrame -from pytest import fixture, mark +from pytest import fixture from sientia_do.notifications.models import NotificationLevel +from sientia_do.observability.sientia_monitoring import SientiaMonitoring from laborious.activities.gates import Gates @@ -17,17 +18,17 @@ def _passthrough_from_dict(): def _minio_payload(retrieve_return, status=None): """ - Build a MinioDataFramePayload-like test double with async retrieve. + Build a MinioDataFramePayload-like test double with retrieve. Args: - retrieve_return: Value returned from await retrieve(minio_repo, metadata). + retrieve_return: Value returned from retrieve(minio_repo, metadata). status: Optional status dict for MLflow response gate (payload.status). Return: MagicMock: Object with async retrieve and optional status. """ p = MagicMock() - p.retrieve = AsyncMock(return_value=retrieve_return) + p.retrieve = MagicMock(return_value=retrieve_return) p.status = status return p @@ -37,7 +38,7 @@ def gates_activity(): gates = Gates( logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) gates.error = MagicMock() gates.debug = MagicMock() @@ -45,8 +46,7 @@ def gates_activity(): gates.warning = MagicMock() gates.critical = MagicMock() gates.send_notification = MagicMock() - gates.send_notification_async = AsyncMock() - gates.emit_metric = AsyncMock() + gates.emit_metric_sync = MagicMock() return gates @@ -60,8 +60,7 @@ metadata = { } -@mark.asyncio -async def test_input_gate_invalid_filter(gates_activity): +def test_input_gate_invalid_filter(gates_activity): # Arrange input_data = { **metadata, @@ -71,7 +70,7 @@ async def test_input_gate_invalid_filter(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = gates_activity.input_gate(input_data) # Assert assert result == (None, 0, '') @@ -80,9 +79,8 @@ async def test_input_gate_invalid_filter(gates_activity): ) -@mark.asyncio @patch('laborious.activities.gates.input_filter_functions') -async def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity): +def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity): # Arrange mock_input_filter_functions.__contains__.return_value = True mock_input_filter_functions.__getitem__.return_value = MagicMock( @@ -96,11 +94,11 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac } # Act - result = await gates_activity.input_gate(input_data) + result = gates_activity.input_gate(input_data) # Assert assert result == (None, 0, '') - gates_activity.send_notification_async.assert_called_once_with( + gates_activity.send_notification.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", @@ -110,8 +108,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac ) -@mark.asyncio -async def test_input_gate_no_filters(gates_activity): +def test_input_gate_no_filters(gates_activity): # Arrange input_data = { **metadata, @@ -121,15 +118,14 @@ async def test_input_gate_no_filters(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = gates_activity.input_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_input_gate_with_filter(gates_activity): +def test_input_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, @@ -139,15 +135,14 @@ async def test_input_gate_with_filter(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = gates_activity.input_gate(input_data) # Assert assert result == ('STOP', -1, 'Input data with bad quality') gates_activity.debug.assert_called() -@mark.asyncio -async def test_input_gate_with_filter_lowercase_keys(gates_activity): +def test_input_gate_with_filter_lowercase_keys(gates_activity): # Arrange input_data = { **metadata, @@ -157,14 +152,13 @@ async def test_input_gate_with_filter_lowercase_keys(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = 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): +def test_input_gate_with_filter_capitalized_keys(gates_activity): # Arrange input_data = { **metadata, @@ -174,14 +168,13 @@ async def test_input_gate_with_filter_capitalized_keys(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = 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): +def test_input_gate_with_filter_not_caught(gates_activity): # Arrange input_data = { **metadata, @@ -191,15 +184,14 @@ async def test_input_gate_with_filter_not_caught(gates_activity): } # Act - result = await gates_activity.input_gate(input_data) + result = gates_activity.input_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_mlflow_response_gate_invalid_filter(gates_activity): +def test_mlflow_response_gate_invalid_filter(gates_activity): # Arrange input_data = { **metadata, @@ -213,15 +205,14 @@ async def test_mlflow_response_gate_invalid_filter(gates_activity): } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = gates_activity.mlflow_response_gate(input_data) # Assert assert result == (None, 0, '') -@mark.asyncio @patch('laborious.activities.gates.mlflow_response_filter_functions') -async def test_mlflow_response_gate_filter_exception( +def test_mlflow_response_gate_filter_exception( mock_mlflow_response_filter_functions, gates_activity ): # Arrange @@ -241,11 +232,11 @@ async def test_mlflow_response_gate_filter_exception( } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = gates_activity.mlflow_response_gate(input_data) # Assert assert result == (None, 0, '') - gates_activity.send_notification_async.assert_called_once_with( + gates_activity.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER', message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error", @@ -255,8 +246,7 @@ async def test_mlflow_response_gate_filter_exception( ) -@mark.asyncio -async def test_mlflow_response_gate_no_filters(gates_activity): +def test_mlflow_response_gate_no_filters(gates_activity): # Arrange input_data = { **metadata, @@ -270,15 +260,14 @@ async def test_mlflow_response_gate_no_filters(gates_activity): } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = gates_activity.mlflow_response_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_mlflow_response_gate_with_filter(gates_activity): +def test_mlflow_response_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, @@ -292,16 +281,15 @@ async def test_mlflow_response_gate_with_filter(gates_activity): } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = gates_activity.mlflow_response_gate(input_data) # Assert assert result == ('STOP', -1, 'API error occurred') gates_activity.debug.assert_called() - gates_activity.send_notification_async.assert_called() + gates_activity.send_notification.assert_called() -@mark.asyncio -async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity): +def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity): # Arrange input_data = { **metadata, @@ -315,14 +303,13 @@ async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity) } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = 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): +def test_mlflow_response_gate_with_filter_not_caught(gates_activity): # Arrange input_data = { **metadata, @@ -336,15 +323,14 @@ async def test_mlflow_response_gate_with_filter_not_caught(gates_activity): } # Act - result = await gates_activity.mlflow_response_gate(input_data) + result = gates_activity.mlflow_response_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_mlflow_content_gate_invalid_filter(gates_activity): +def test_mlflow_content_gate_invalid_filter(gates_activity): # Arrange input_data = { **metadata, @@ -355,17 +341,14 @@ async def test_mlflow_content_gate_invalid_filter(gates_activity): } # Act - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) # Assert assert result == (None, 0, '') -@mark.asyncio @patch('laborious.activities.gates.mlflow_content_filter_functions') -async def test_mlflow_content_gate_filter_exception( - mock_mlflow_content_filter_functions, gates_activity -): +def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functions, gates_activity): # Arrange mock_mlflow_content_filter_functions.__contains__.return_value = True mock_mlflow_content_filter_functions.__getitem__.return_value = MagicMock( @@ -380,12 +363,12 @@ async def test_mlflow_content_gate_filter_exception( } # Act - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() - gates_activity.send_notification_async.assert_called_once_with( + gates_activity.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR', message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test error", @@ -395,8 +378,7 @@ async def test_mlflow_content_gate_filter_exception( ) -@mark.asyncio -async def test_mlflow_content_gate_no_filters(gates_activity): +def test_mlflow_content_gate_no_filters(gates_activity): # Arrange input_data = { **metadata, @@ -407,15 +389,14 @@ async def test_mlflow_content_gate_no_filters(gates_activity): } # Act - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_mlflow_content_gate_with_filter(gates_activity): +def test_mlflow_content_gate_with_filter(gates_activity): # Arrange input_data = { **metadata, @@ -426,16 +407,15 @@ async def test_mlflow_content_gate_with_filter(gates_activity): } # Act - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) # Assert assert result == ('STOP', -1, 'Transformed data not passed the content filter') gates_activity.debug.assert_called() - gates_activity.send_notification_async.assert_called() + gates_activity.send_notification.assert_called() -@mark.asyncio -async def test_mlflow_content_gate_with_filter_not_caught(gates_activity): +def test_mlflow_content_gate_with_filter_not_caught(gates_activity): # Arrange input_data = { **metadata, @@ -446,15 +426,14 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity): } # Act - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) # Assert assert result == (None, 0, '') gates_activity.debug.assert_called() -@mark.asyncio -async def test_mlflow_content_gate_filter_returns_false(gates_activity): +def test_mlflow_content_gate_filter_returns_false(gates_activity): input_data = { **metadata, 'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}}, @@ -463,7 +442,7 @@ async def test_mlflow_content_gate_filter_returns_false(gates_activity): 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], } - result = await gates_activity.mlflow_content_gate(input_data) + result = gates_activity.mlflow_content_gate(input_data) assert result == (None, 0, '') gates_activity.debug.assert_called() @@ -525,8 +504,7 @@ def test_get_prediction_store_policy_valid_policy(gates_activity): assert policy_value == 1 -@mark.asyncio -async def test_format_prediction_no_timestamp(gates_activity): +def test_format_prediction_no_timestamp(gates_activity): # Arrange input_data = { **metadata, @@ -545,7 +523,7 @@ async def test_format_prediction_no_timestamp(gates_activity): } # Act - result = await gates_activity.format_prediction(input_data) + result = gates_activity.format_prediction(input_data) # Assert assert result['prediction'] == {0: 1} @@ -557,8 +535,7 @@ async def test_format_prediction_no_timestamp(gates_activity): assert result['comments'] == {0: ''} -@mark.asyncio -async def test_format_prediction_with_timestamp_erl(gates_activity): +def test_format_prediction_with_timestamp_erl(gates_activity): # Arrange input_data = { **metadata, @@ -585,7 +562,7 @@ async def test_format_prediction_with_timestamp_erl(gates_activity): } # Act - result = await gates_activity.format_prediction(input_data) + result = gates_activity.format_prediction(input_data) # Assert assert result['prediction'] == {0: 2, 1: 1} @@ -597,8 +574,7 @@ async def test_format_prediction_with_timestamp_erl(gates_activity): assert result['comments'] == {0: '', 1: ''} -@mark.asyncio -async def test_format_prediction_with_timestamp_lts(gates_activity): +def test_format_prediction_with_timestamp_lts(gates_activity): # Arrange input_data = { **metadata, @@ -625,7 +601,7 @@ async def test_format_prediction_with_timestamp_lts(gates_activity): } # Act - result = await gates_activity.format_prediction(input_data) + result = gates_activity.format_prediction(input_data) # Assert assert result['prediction'] == {0: 3, 1: 2} @@ -637,8 +613,7 @@ async def test_format_prediction_with_timestamp_lts(gates_activity): assert result['comments'] == {0: '', 1: ''} -@mark.asyncio -async def test_format_prediction_with_timestamp_invalid_policy(gates_activity): +def test_format_prediction_with_timestamp_invalid_policy(gates_activity): # Arrange input_data = { **metadata, @@ -663,16 +638,15 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity): gates_activity.get_prediction_store_policy = MagicMock(return_value=('invalid', 1)) try: - await gates_activity.format_prediction(input_data) + gates_activity.format_prediction(input_data) except ValueError as e: assert str(e) == 'Invalid policy type: invalid' else: raise AssertionError('Expected ValueError') -@mark.asyncio -@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock) -async def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity): +@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock) +def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity): # Arrange payload_result = MagicMock() mock_from_dataframe.return_value = payload_result @@ -691,7 +665,7 @@ async def test_format_transformed_data_single_row(mock_from_dataframe, gates_act } # Act - result = await gates_activity.format_transformed_data(input_data) + result = gates_activity.format_transformed_data(input_data) # Assert assert result is payload_result @@ -705,9 +679,8 @@ async def test_format_transformed_data_single_row(mock_from_dataframe, gates_act gates_activity.info.assert_called() -@mark.asyncio -@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock) -async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity): +@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock) +def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity): # Arrange payload_result = MagicMock() mock_from_dataframe.return_value = payload_result @@ -732,7 +705,7 @@ async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_ } # Act - result = await gates_activity.format_transformed_data(input_data) + result = gates_activity.format_transformed_data(input_data) # Assert assert result is payload_result @@ -746,9 +719,8 @@ async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_ gates_activity.info.assert_called() -@mark.asyncio -@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock) -async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity): +@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock) +def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity): # Arrange payload_result = MagicMock() mock_from_dataframe.return_value = payload_result @@ -760,7 +732,7 @@ async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_act } # Act - result = await gates_activity.format_transformed_data(input_data) + result = gates_activity.format_transformed_data(input_data) # Assert assert result is payload_result @@ -774,8 +746,7 @@ async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_act gates_activity.info.assert_called() -@mark.asyncio -async def test_format_default_prediction(gates_activity): +def test_format_default_prediction(gates_activity): # Arrange input_data = { **metadata, @@ -786,7 +757,7 @@ async def test_format_default_prediction(gates_activity): } # Act - result = await gates_activity.format_default_prediction(input_data) + result = gates_activity.format_default_prediction(input_data) # Assert assert result['prediction'] == {0: 0} @@ -799,8 +770,7 @@ async def test_format_default_prediction(gates_activity): gates_activity.debug.assert_called() -@mark.asyncio -async def test_format_retrain_report(gates_activity): +def test_format_retrain_report(gates_activity): # Arrange input_data = { **metadata, @@ -819,7 +789,7 @@ async def test_format_retrain_report(gates_activity): } # Act - result = await gates_activity.format_retrain_report(input_data) + result = gates_activity.format_retrain_report(input_data) # Assert assert result['model_id'] == {0: 'test_model'} @@ -831,8 +801,7 @@ async def test_format_retrain_report(gates_activity): assert result['mlflow_experiment_id'] == {0: 'test_mlflow_experiment_id'} -@mark.asyncio -async def test_format_retrain_report_failure(gates_activity): +def test_format_retrain_report_failure(gates_activity): # Arrange input_data = { **metadata, @@ -851,7 +820,7 @@ async def test_format_retrain_report_failure(gates_activity): } # Act - result = await gates_activity.format_retrain_report(input_data) + result = gates_activity.format_retrain_report(input_data) # Assert assert result['model_id'] == {0: 'test_model'} @@ -865,9 +834,8 @@ async def test_format_retrain_report_failure(gates_activity): gates_activity.debug.assert_called() -@mark.asyncio @patch('laborious.activities.gates.metrics') -async def test_write_metrics(mock_metrics, gates_activity): +def test_write_metrics(mock_metrics, gates_activity): """Test write_metrics method.""" input_data = { **metadata, @@ -878,7 +846,7 @@ async def test_write_metrics(mock_metrics, gates_activity): }, 'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}}, } - await gates_activity.write_metrics(input_data) + gates_activity.write_metrics(input_data) core_tags = { 'pod_id': gates_activity.pod_id, 'runtime': gates_activity.runtime, @@ -886,7 +854,7 @@ async def test_write_metrics(mock_metrics, gates_activity): 'model_name': metadata['metadata']['model_name'], 'workflow_name': metadata['metadata']['workflow_name'], } - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT, @@ -894,7 +862,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR, @@ -904,7 +872,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR, @@ -914,7 +882,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT, @@ -926,7 +894,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR, @@ -940,7 +908,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT, @@ -952,7 +920,7 @@ async def test_write_metrics(mock_metrics, gates_activity): ), ] ) - gates_activity.emit_metric.assert_has_calls( + gates_activity.emit_metric_sync.assert_has_calls( [ call( metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR, @@ -968,9 +936,8 @@ async def test_write_metrics(mock_metrics, gates_activity): ) -@mark.asyncio @patch('laborious.activities.gates.metrics') -async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity): +def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity): """Test write_metrics method with None response_time in opc_metrics.""" input_data = { **metadata, @@ -981,10 +948,10 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act }, 'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': None}}, } - await gates_activity.write_metrics(input_data) + gates_activity.write_metrics(input_data) # Verify that metrics for tag1 are emitted - gates_activity.emit_metric.assert_any_call( + gates_activity.emit_metric_sync.assert_any_call( metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR, method='observe', tags={ @@ -1002,7 +969,38 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act # Verify that metrics for tag2 (with None response_time) are NOT emitted calls = [ c - for c in gates_activity.emit_metric.call_args_list + for c in gates_activity.emit_metric_sync.call_args_list if len(c[1].get('tags', {})) > 0 and c[1]['tags'].get('tag') == 'tag2' ] assert len(calls) == 0, 'Metrics should not be emitted for None response_time' + + +@patch.object(SientiaMonitoring, 'shutdown') +def test_close_disposes_minio_repository(mock_shutdown): + """ + ``Gates.close`` should close the optional MinIO client and clear the repository reference. + """ + minio = MagicMock() + gates = Gates( + minio_repository=minio, + logger=MagicMock(), + notification_handler=MagicMock(), + metrics_controller=MagicMock(), + ) + gates.close() + minio.close.assert_called_once() + assert gates.minio_repository is None + mock_shutdown.assert_called_once_with(gates) + + +@patch.object(SientiaMonitoring, 'shutdown') +def test_close_without_minio_repository(mock_shutdown): + """When no MinIO repository is configured, ``close`` only shuts down monitoring.""" + gates = Gates( + minio_repository=None, + logger=MagicMock(), + notification_handler=MagicMock(), + metrics_controller=MagicMock(), + ) + gates.close() + mock_shutdown.assert_called_once_with(gates) diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index be17a12..30f792a 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -1,9 +1,9 @@ from datetime import datetime -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import numpy as np import pandas as pd -from pytest import fixture, mark, raises +from pytest import fixture, raises from sientia_do.notifications.models import NotificationLevel from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ @@ -22,7 +22,7 @@ def _passthrough_from_dict(): def test___init__(mock_minio_repository): logger = MagicMock() notification_handler = MagicMock() - metrics_controller = AsyncMock() + metrics_controller = MagicMock() mlflow_repo = MagicMock() plugin_store = MagicMock() @@ -63,7 +63,7 @@ def test___init__(mock_minio_repository): def mlflow(mock_minio_repository): logger = MagicMock() notification_handler = MagicMock() - metrics_controller = AsyncMock() + metrics_controller = MagicMock() mlflow_repo = MagicMock() plugin_store = MagicMock() @@ -85,11 +85,10 @@ def mlflow(mock_minio_repository): metrics_controller=metrics_controller, ) - mlflow.minio_repository = AsyncMock() + mlflow.minio_repository = MagicMock() mlflow.send_notification = MagicMock() - mlflow.emit_metric = AsyncMock() - mlflow.send_notification_async = AsyncMock() + mlflow.emit_metric = MagicMock() mlflow.error = MagicMock() mlflow.debug = MagicMock() mlflow.info = MagicMock() @@ -150,12 +149,11 @@ def test_detect_and_parse_datetime_index_timestamp_with_tz_success(mlflow): assert out.index[0].endswith('+0000') -@mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, ) -async def test_request_transform_success(mock_from_dataframe, mlflow): +def test_request_transform_success(mock_from_dataframe, mlflow): ts = pd.Timestamp('2020-01-01', tz='UTC') raw = pd.DataFrame( { @@ -180,8 +178,8 @@ async def test_request_transform_success(mock_from_dataframe, mlflow): wrapper.transform.return_value = (out_df, {'meta': True}) mlflow.mlflow_repository.get_cached_model.return_value = wrapper - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw) input_data = { **metadata, @@ -190,7 +188,7 @@ async def test_request_transform_success(mock_from_dataframe, mlflow): 'model_config': {}, } - response_data = await mlflow.request_transform(input_data) + response_data = mlflow.request_transform(input_data) mlflow.mlflow_repository.get_cached_model.assert_called_once_with( model_name='test_model', @@ -202,12 +200,11 @@ 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, + new_callable=MagicMock, ) -async def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow): +def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow): ts = pd.Timestamp('2020-01-01', tz='UTC') raw = pd.DataFrame( { @@ -223,8 +220,8 @@ async def test_request_transform_success_without_transform_meta(mock_from_datafr wrapper.transform.return_value = (out_df, {}) mlflow.mlflow_repository.get_cached_model.return_value = wrapper - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw) input_data = { **metadata, @@ -233,21 +230,20 @@ async def test_request_transform_success_without_transform_meta(mock_from_datafr 'model_config': {}, } - await mlflow.request_transform(input_data) + mlflow.request_transform(input_data) mock_from_dataframe.assert_called_once() -@mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, ) -async def test_request_transform_failure(mock_from_dataframe, mlflow): +def test_request_transform_failure(mock_from_dataframe, mlflow): mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('boom') data_mock = MagicMock() - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=data_mock) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=data_mock) input_data = { **metadata, @@ -260,7 +256,7 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow): data_mock.drop_duplicates.return_value = data_mock data_mock.pivot.return_value = data_mock - await mlflow.request_transform(input_data) + mlflow.request_transform(input_data) mock_from_dataframe.assert_called_once_with( dataframe=None, @@ -274,13 +270,12 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow): ) -@mark.asyncio @patch( 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, ) @patch('laborious.activities.mlflow.to_datetime') -async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): +def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): wrapper = MagicMock() pred_df = MagicMock() wrapper.predict.return_value = (pred_df, {}) @@ -288,8 +283,8 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): data_mock = MagicMock() data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')]) - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=data_mock) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=data_mock) input_data = { **metadata, @@ -301,7 +296,7 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): pred_df.columns = MagicMock() pred_df.__setitem__ = MagicMock() - response_data = await mlflow.request_predict(input_data) + response_data = mlflow.request_predict(input_data) data_mock.replace.assert_called_once_with(np.nan, None, inplace=True) mock_to_datetime.assert_called() @@ -310,15 +305,12 @@ 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, + new_callable=MagicMock, ) @patch('laborious.activities.mlflow.to_datetime') -async def test_request_predict_success_dataframe_and_meta( - mock_to_datetime, mock_from_dataframe, mlflow -): +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}) @@ -326,8 +318,8 @@ async def test_request_predict_success_dataframe_and_meta( data_mock = MagicMock() data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')]) - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=data_mock) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=data_mock) input_data = { **metadata, @@ -336,25 +328,24 @@ async def test_request_predict_success_dataframe_and_meta( 'model_config': {}, } - await mlflow.request_predict(input_data) + 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', - new_callable=AsyncMock, + new_callable=MagicMock, ) @patch('laborious.activities.mlflow.to_datetime') -async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, mlflow): +def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, mlflow): mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('predict boom') data_mock = MagicMock() - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=data_mock) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=data_mock) input_data = { **metadata, @@ -363,7 +354,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml 'model_config': {}, } - await mlflow.request_predict(input_data) + mlflow.request_predict(input_data) mock_from_dataframe.assert_called_once_with( dataframe=None, @@ -377,12 +368,11 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml ) -@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_data_success_retrain( +def test_retrain_model_success_data_success_retrain( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): mock_mkdtemp.return_value = 'tmp' @@ -408,10 +398,10 @@ async def test_retrain_model_success_data_success_retrain( } ) - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw_data) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw_data) - response = await mlflow.retrain_model( + response = mlflow.retrain_model( { **metadata, 'data': payload, @@ -428,12 +418,11 @@ async def test_retrain_model_success_data_success_retrain( assert response['experiment']['run_id'] == 'new-run' -@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_with_payload_data( +def test_retrain_model_success_with_payload_data( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): mock_mkdtemp.return_value = 'tmp' @@ -448,8 +437,8 @@ async def test_retrain_model_success_with_payload_data( raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at']) raw_data.__getitem__.return_value.max.return_value = 'ts' - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw_data) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw_data) pivoted = MagicMock() raw_data.sort_values.return_value = raw_data @@ -460,7 +449,7 @@ async def test_retrain_model_success_with_payload_data( pivoted.index = MagicMock() pivoted.__setitem__ = MagicMock() - response = await mlflow.retrain_model( + response = mlflow.retrain_model( { **metadata, 'data': payload, @@ -472,12 +461,11 @@ 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_always_uses_retrain_even_with_full_retrain_flag( +def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag( mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow ): mock_mkdtemp.return_value = 'tmp' @@ -498,10 +486,10 @@ async def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag( 'value': [1.0, 2.0, 3.0, 4.0], } ) - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw_data) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw_data) - response = await mlflow.retrain_model( + response = mlflow.retrain_model( { **metadata, 'data': payload, @@ -515,17 +503,16 @@ async def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag( 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): +def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow): mv_alias = MagicMock(run_id='src') mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('retrain failed') raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at']) raw_data.__getitem__.return_value.max.return_value = 'tsmax' - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw_data) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw_data) pivoted = MagicMock() raw_data.sort_values.return_value = raw_data @@ -536,7 +523,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow) pivoted.index = MagicMock() pivoted.__setitem__ = MagicMock() - response = await mlflow.retrain_model( + response = mlflow.retrain_model( { **metadata, 'data': payload, @@ -549,9 +536,8 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow) assert 'retrain failed' in response['message'] -@mark.asyncio -async def test_retrain_model_data_error(mlflow): - response = await mlflow.retrain_model( +def test_retrain_model_data_error(mlflow): + response = mlflow.retrain_model( { **metadata, 'model_name': 'test_model', @@ -565,8 +551,7 @@ async def test_retrain_model_data_error(mlflow): assert 'data' in response['message'].lower() or 'loading' in response['message'].lower() -@mark.asyncio -async def test_retrain_model_missing_target(mlflow): +def test_retrain_model_missing_target(mlflow): ts = pd.Timestamp('2020-01-01', tz='UTC') raw_data = pd.DataFrame( { @@ -575,10 +560,10 @@ async def test_retrain_model_missing_target(mlflow): 'value': [1.0, 2.0], } ) - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=raw_data) + payload = MagicMock() + payload.retrieve = MagicMock(return_value=raw_data) - response = await mlflow.retrain_model( + response = mlflow.retrain_model( { **metadata, 'data': payload, @@ -591,12 +576,11 @@ async def test_retrain_model_missing_target(mlflow): assert 'target' in response['message'] -@mark.asyncio -async def test_retrain_model_data_error_no_minio_repository(mlflow): +def test_retrain_model_data_error_no_minio_repository(mlflow): mlflow.minio_repository = None with raises(ValueError) as e: - await mlflow.retrain_model( + mlflow.retrain_model( { **metadata, 'object_key': 'test_object_key', @@ -610,8 +594,7 @@ async def test_retrain_model_data_error_no_minio_repository(mlflow): assert str(e.value) == 'Minio repository not initialized' -@mark.asyncio -async def test_update_production_model(mlflow): +def test_update_production_model(mlflow): mlflow.mlflow_repository._client.search_model_versions.return_value = [ MagicMock(version='3', run_id='run-x'), MagicMock(version='2', run_id='run-x'), @@ -626,7 +609,7 @@ async def test_update_production_model(mlflow): 'status': 'success', } - response = await mlflow.update_production_model(input_data) + response = mlflow.update_production_model(input_data) mlflow.mlflow_repository.promote_to_alias.assert_called_once_with( model_name='test_model', @@ -639,8 +622,7 @@ async def test_update_production_model(mlflow): assert response['version'] == '3' -@mark.asyncio -async def test_update_production_model_error(mlflow): +def test_update_production_model_error(mlflow): mlflow.mlflow_repository._client.search_model_versions.return_value = [] input_data = { @@ -653,9 +635,9 @@ async def test_update_production_model_error(mlflow): } try: - await mlflow.update_production_model(input_data) + mlflow.update_production_model(input_data) except Exception: - mlflow.send_notification_async.assert_called_once_with( + mlflow.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='UPDATE_PRODUCTION_MODEL_ERROR', message=ANY, @@ -667,9 +649,8 @@ async def test_update_production_model_error(mlflow): raise AssertionError('Expected exception') -@mark.asyncio @patch('laborious.activities.mlflow.to_datetime') -async def test_get_reference_data_success(mock_to_datetime, mlflow): +def test_get_reference_data_success(mock_to_datetime, mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -690,14 +671,13 @@ async def test_get_reference_data_success(mock_to_datetime, mlflow): with patch('laborious.activities.mlflow.rmtree'): with patch('laborious.activities.mlflow.Path') as mp: mp.return_value.rglob.return_value = [MagicMock()] - result = await mlflow.get_reference_data(input_data) + result = mlflow.get_reference_data(input_data) mock_reference_data.to_dict.assert_called_once_with(orient='records') assert result == mock_reference_data.to_dict.return_value -@mark.asyncio -async def test_get_reference_data_not_found(mlflow): +def test_get_reference_data_not_found(mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -705,14 +685,13 @@ async def test_get_reference_data_not_found(mlflow): mlflow.mlflow_repository._client.get_model_version_by_alias.side_effect = Exception('missing') - result = await mlflow.get_reference_data(input_data) + result = mlflow.get_reference_data(input_data) mlflow.warning.assert_called() assert result is None -@mark.asyncio -async def test_get_reference_data_missing_csv_file_returns_none(mlflow): +def test_get_reference_data_missing_csv_file_returns_none(mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -724,13 +703,12 @@ async def test_get_reference_data_missing_csv_file_returns_none(mlflow): 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) + result = mlflow.get_reference_data(input_data) assert result is None -@mark.asyncio -async def test_get_reference_data_exception(mlflow): +def test_get_reference_data_exception(mlflow): input_data = { **metadata, 'model_name': 'test_model', @@ -740,6 +718,6 @@ async def test_get_reference_data_exception(mlflow): mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail') - result = await mlflow.get_reference_data(input_data) + result = mlflow.get_reference_data(input_data) assert result is None diff --git a/tests/laborious/activities/test_model_metrics.py b/tests/laborious/activities/test_model_metrics.py index b2b5e0b..53c7ebb 100644 --- a/tests/laborious/activities/test_model_metrics.py +++ b/tests/laborious/activities/test_model_metrics.py @@ -1,7 +1,7 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import ANY, MagicMock, patch from pandas import DataFrame -from pytest import fixture, mark +from pytest import fixture from sientia_do.notifications.models import NotificationLevel from laborious.activities.model_metrics import ModelMetrics @@ -12,7 +12,7 @@ def model_metrics_activity(): model_metrics = ModelMetrics( logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) model_metrics.error = MagicMock() model_metrics.debug = MagicMock() @@ -20,8 +20,7 @@ def model_metrics_activity(): model_metrics.warning = MagicMock() model_metrics.critical = MagicMock() model_metrics.send_notification = MagicMock() - model_metrics.send_notification_async = AsyncMock() - model_metrics.emit_metric = AsyncMock() + model_metrics.emit_metric_sync = MagicMock() model_metrics.get_core_labels = MagicMock( return_value={ 'pod_id': 'test_pod', @@ -29,7 +28,7 @@ def model_metrics_activity(): 'workflow_name': 'test_workflow', } ) - model_metrics.observe_lag = AsyncMock() + model_metrics.observe_lag_sync = MagicMock() model_metrics.pod_id = 'test_pod' return model_metrics @@ -44,8 +43,7 @@ metadata = { } -@mark.asyncio -async def test_calculate_drift_invalid_chunk_period(model_metrics_activity): +def test_calculate_drift_invalid_chunk_period(model_metrics_activity): # Arrange input_data = { **metadata, @@ -64,7 +62,7 @@ async def test_calculate_drift_invalid_chunk_period(model_metrics_activity): # Act & Assert try: - await model_metrics_activity.calculate_drift(input_data) + model_metrics_activity.calculate_drift(input_data) except ValueError as e: assert str(e) == 'Invalid chunk period: invalid, must be "min" or "s"' model_metrics_activity.error.assert_called_once_with( @@ -74,10 +72,9 @@ async def test_calculate_drift_invalid_chunk_period(model_metrics_activity): raise AssertionError('Expected ValueError') -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_with_reference_data( +def test_calculate_drift_with_reference_data( mock_to_datetime, mock_dataframe, model_metrics_activity ): # Arrange @@ -107,7 +104,7 @@ async def test_calculate_drift_with_reference_data( } ] - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df) reference_data = DataFrame( { @@ -142,7 +139,7 @@ async def test_calculate_drift_with_reference_data( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert isinstance(result, list) @@ -150,10 +147,18 @@ async def test_calculate_drift_with_reference_data( model_metrics_activity.info.assert_called() model_metrics_activity.get_drift_metrics.assert_called_once() # Verify transformations were called - mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True) + mock_drift_df.drop.assert_called_once_with( + columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore' + ) mock_drift_df.__getitem__.assert_called() mock_drift_df.rename.assert_called_once_with( - columns={'metric': 'method', 'statistic': 'value'}, inplace=True + columns={ + 'metric': 'method', + 'statistic': 'value', + 'alert': 'drift', + 'chunk_index': 'chunk', + 'chunk_end_date': 'timestamp_end', + } ) mock_drift_df.drop_duplicates.assert_called_once_with( subset=['timestamp', 'method', 'feature'], keep='first', inplace=True @@ -161,10 +166,9 @@ async def test_calculate_drift_with_reference_data( mock_drift_df.to_dict.assert_called_once_with(orient='records') -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_without_reference_data( +def test_calculate_drift_without_reference_data( mock_to_datetime, mock_dataframe, model_metrics_activity ): # Arrange @@ -194,7 +198,7 @@ async def test_calculate_drift_without_reference_data( } ] - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df) target_data_dict = { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'], @@ -232,13 +236,13 @@ async def test_calculate_drift_without_reference_data( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert isinstance(result, list) assert result == mock_drift_df.to_dict.return_value model_metrics_activity.warning.assert_called() - model_metrics_activity.send_notification_async.assert_called_once_with( + model_metrics_activity.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING', message='Using 30% first rows of target data as reference data', @@ -247,10 +251,18 @@ async def test_calculate_drift_without_reference_data( attachment_content=ANY, ) # Verify transformations were called - mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True) + mock_drift_df.drop.assert_called_once_with( + columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore' + ) mock_drift_df.__getitem__.assert_called() mock_drift_df.rename.assert_called_once_with( - columns={'metric': 'method', 'statistic': 'value'}, inplace=True + columns={ + 'metric': 'method', + 'statistic': 'value', + 'alert': 'drift', + 'chunk_index': 'chunk', + 'chunk_end_date': 'timestamp_end', + } ) mock_drift_df.drop_duplicates.assert_called_once_with( subset=['timestamp', 'method', 'feature'], keep='first', inplace=True @@ -258,16 +270,13 @@ async def test_calculate_drift_without_reference_data( mock_drift_df.to_dict.assert_called_once_with(orient='records') -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_empty_drift_df( - mock_to_datetime, mock_dataframe, model_metrics_activity -): +def test_calculate_drift_empty_drift_df(mock_to_datetime, mock_dataframe, model_metrics_activity): # Arrange mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27' - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=DataFrame()) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=DataFrame()) reference_data = DataFrame( { @@ -302,7 +311,7 @@ async def test_calculate_drift_empty_drift_df( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert result == [] @@ -311,10 +320,9 @@ async def test_calculate_drift_empty_drift_df( ) -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_empty_after_timestamp_filter( +def test_calculate_drift_empty_after_timestamp_filter( mock_to_datetime, mock_dataframe, model_metrics_activity ): # Arrange @@ -340,7 +348,7 @@ async def test_calculate_drift_empty_after_timestamp_filter( mock_drift_df.__getitem__.side_effect = getitem_side_effect - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df) reference_data = DataFrame( { @@ -375,7 +383,7 @@ async def test_calculate_drift_empty_after_timestamp_filter( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert result == [] @@ -383,17 +391,16 @@ async def test_calculate_drift_empty_after_timestamp_filter( 'No drift metrics found after dropping rows where timestamp is not in target data', metadata['metadata'], ) - # Verify transformations were called - mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True) + # When the timestamp filter empties the dataframe, the rename/drop pipeline + # is short-circuited, so neither ``drop`` nor ``rename`` should run. + mock_drift_df.drop.assert_not_called() + mock_drift_df.rename.assert_not_called() mock_drift_df.__getitem__.assert_called() -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_success_min( - mock_to_datetime, mock_dataframe, model_metrics_activity -): +def test_calculate_drift_success_min(mock_to_datetime, mock_dataframe, model_metrics_activity): # Arrange mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27' mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = ( @@ -421,7 +428,7 @@ async def test_calculate_drift_success_min( } ] - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df) reference_data = DataFrame( { @@ -456,7 +463,7 @@ async def test_calculate_drift_success_min( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert isinstance(result, list) @@ -464,10 +471,18 @@ async def test_calculate_drift_success_min( model_metrics_activity.info.assert_called() model_metrics_activity.get_drift_metrics.assert_called_once() # Verify transformations were called - mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True) + mock_drift_df.drop.assert_called_once_with( + columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore' + ) mock_drift_df.__getitem__.assert_called() mock_drift_df.rename.assert_called_once_with( - columns={'metric': 'method', 'statistic': 'value'}, inplace=True + columns={ + 'metric': 'method', + 'statistic': 'value', + 'alert': 'drift', + 'chunk_index': 'chunk', + 'chunk_end_date': 'timestamp_end', + } ) mock_drift_df.drop_duplicates.assert_called_once_with( subset=['timestamp', 'method', 'feature'], keep='first', inplace=True @@ -475,10 +490,9 @@ async def test_calculate_drift_success_min( mock_drift_df.to_dict.assert_called_once_with(orient='records') -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model_metrics_activity): +def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model_metrics_activity): # Arrange mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27' mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = ( @@ -506,7 +520,7 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model } ] - model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df) + model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df) reference_data = DataFrame( { @@ -541,7 +555,7 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert isinstance(result, list) @@ -549,10 +563,18 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model model_metrics_activity.info.assert_called() model_metrics_activity.get_drift_metrics.assert_called_once() # Verify transformations were called - mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True) + mock_drift_df.drop.assert_called_once_with( + columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore' + ) mock_drift_df.__getitem__.assert_called() mock_drift_df.rename.assert_called_once_with( - columns={'metric': 'method', 'statistic': 'value'}, inplace=True + columns={ + 'metric': 'method', + 'statistic': 'value', + 'alert': 'drift', + 'chunk_index': 'chunk', + 'chunk_end_date': 'timestamp_end', + } ) mock_drift_df.drop_duplicates.assert_called_once_with( subset=['timestamp', 'method', 'feature'], keep='first', inplace=True @@ -560,16 +582,15 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model mock_drift_df.to_dict.assert_called_once_with(orient='records') -@mark.asyncio @patch('laborious.activities.model_metrics.DataFrame') @patch('laborious.activities.model_metrics.to_datetime') -async def test_calculate_drift_get_drift_metrics_error( +def test_calculate_drift_get_drift_metrics_error( mock_to_datetime, mock_dataframe, model_metrics_activity ): # Arrange mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27' - model_metrics_activity.get_drift_metrics = AsyncMock( + model_metrics_activity.get_drift_metrics = MagicMock( side_effect=Exception('Get drift metrics error') ) @@ -606,14 +627,14 @@ async def test_calculate_drift_get_drift_metrics_error( } # Act - result = await model_metrics_activity.calculate_drift(input_data) + result = model_metrics_activity.calculate_drift(input_data) # Assert assert result == [] model_metrics_activity.error.assert_called_once_with( 'Error getting drift metrics: Get drift metrics error', metadata['metadata'] ) - model_metrics_activity.send_notification_async.assert_called_once_with( + model_metrics_activity.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR', message='Error getting drift metrics: Get drift metrics error', @@ -623,12 +644,11 @@ async def test_calculate_drift_get_drift_metrics_error( ) -@mark.asyncio @patch('laborious.activities.model_metrics.to_datetime') @patch('laborious.activities.model_metrics.time.time') @patch('laborious.activities.model_metrics.ModelAnalysis') @patch('laborious.activities.model_metrics.metrics') -async def test_get_drift_metrics_success( +def test_get_drift_metrics_success( mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity ): # Arrange @@ -668,7 +688,7 @@ async def test_get_drift_metrics_success( ).columns # Act - result = await model_metrics_activity.get_drift_metrics( + result = model_metrics_activity.get_drift_metrics( reference_data=reference_data, target_data=target_data, target_name='target', @@ -681,16 +701,15 @@ async def test_get_drift_metrics_success( # Assert assert isinstance(result, DataFrame) model_metrics_activity.debug.assert_called() - model_metrics_activity.observe_lag.assert_called() - model_metrics_activity.emit_metric.assert_called() + model_metrics_activity.observe_lag_sync.assert_called() + model_metrics_activity.emit_metric_sync.assert_called() -@mark.asyncio @patch('laborious.activities.model_metrics.to_datetime') @patch('laborious.activities.model_metrics.time.time') @patch('laborious.activities.model_metrics.ModelAnalysis') @patch('laborious.activities.model_metrics.metrics') -async def test_get_drift_metrics_univariate_error( +def test_get_drift_metrics_univariate_error( mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity ): # Arrange @@ -722,7 +741,7 @@ async def test_get_drift_metrics_univariate_error( # Act & Assert try: - await model_metrics_activity.get_drift_metrics( + model_metrics_activity.get_drift_metrics( reference_data=reference_data, target_data=target_data, target_name='target', @@ -736,19 +755,18 @@ async def test_get_drift_metrics_univariate_error( model_metrics_activity.error.assert_called_once_with( 'Error detecting univariate drift: Univariate drift error', metadata['metadata'] ) - model_metrics_activity.emit_metric.assert_called_with( + model_metrics_activity.emit_metric_sync.assert_called_with( metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY ) else: raise AssertionError('Expected Exception') -@mark.asyncio @patch('laborious.activities.model_metrics.to_datetime') @patch('laborious.activities.model_metrics.time.time') @patch('laborious.activities.model_metrics.ModelAnalysis') @patch('laborious.activities.model_metrics.metrics') -async def test_get_drift_metrics_multivariate_error( +def test_get_drift_metrics_multivariate_error( mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity ): mock_time.return_value = 1000.0 @@ -769,7 +787,7 @@ async def test_get_drift_metrics_multivariate_error( ).columns try: - await model_metrics_activity.get_drift_metrics( + model_metrics_activity.get_drift_metrics( reference_data=reference_data, target_data=target_data, target_name='target', @@ -783,19 +801,18 @@ async def test_get_drift_metrics_multivariate_error( model_metrics_activity.error.assert_called_once_with( 'Error detecting multivariate drift: Multivariate drift error', metadata['metadata'] ) - model_metrics_activity.emit_metric.assert_called_with( + model_metrics_activity.emit_metric_sync.assert_called_with( metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY ) else: raise AssertionError('Expected Exception') -@mark.asyncio @patch('laborious.activities.model_metrics.to_datetime') @patch('laborious.activities.model_metrics.time.time') @patch('laborious.activities.model_metrics.ModelAnalysis') @patch('laborious.activities.model_metrics.metrics') -async def test_get_drift_metrics_dataframe_error( +def test_get_drift_metrics_dataframe_error( mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity ): mock_time.return_value = 1000.0 @@ -817,7 +834,7 @@ async def test_get_drift_metrics_dataframe_error( ).columns try: - await model_metrics_activity.get_drift_metrics( + model_metrics_activity.get_drift_metrics( reference_data=reference_data, target_data=target_data, target_name='target', @@ -831,15 +848,14 @@ async def test_get_drift_metrics_dataframe_error( model_metrics_activity.error.assert_called_once_with( 'Error getting drift metrics: Dataframe error', metadata['metadata'] ) - model_metrics_activity.emit_metric.assert_called_with( + model_metrics_activity.emit_metric_sync.assert_called_with( metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY ) else: raise AssertionError('Expected Exception') -@mark.asyncio -async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity): +def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -858,7 +874,7 @@ async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activi } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 4 @@ -877,8 +893,7 @@ async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activi model_metrics_activity.debug.assert_called_once() -@mark.asyncio -async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity): +def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -897,7 +912,7 @@ async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 1 @@ -911,8 +926,7 @@ async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity ) -@mark.asyncio -async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity): +def test_calculate_simple_metrics_success_mse_only(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -931,7 +945,7 @@ async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity) } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 1 @@ -945,8 +959,7 @@ async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity) ) -@mark.asyncio -async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity): +def test_calculate_simple_metrics_success_mae_only(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -965,7 +978,7 @@ async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity) } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 1 @@ -979,8 +992,7 @@ async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity) ) -@mark.asyncio -async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): +def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -999,7 +1011,7 @@ async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 1 @@ -1013,8 +1025,7 @@ async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity): ) -@mark.asyncio -async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): +def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): # Arrange # All target values are the same, so ss_tot will be 0 target_data = DataFrame( @@ -1034,7 +1045,7 @@ async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 1 @@ -1049,8 +1060,7 @@ async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity): ) -@mark.asyncio -async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity): +def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity): # Arrange target_data = DataFrame( { @@ -1069,7 +1079,7 @@ async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_me } # Act - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) # Assert assert len(result['metric']) == 2 @@ -1084,8 +1094,7 @@ async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_me ) -@mark.asyncio -async def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity): +def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity): target_data = DataFrame( { 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'], @@ -1102,7 +1111,7 @@ async def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_act 'interval_minutes': 5, } - result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data)) + result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data)) assert len(result['metric']) == 1 assert result['metric'].values[0] == 'rmse' diff --git a/tests/laborious/activities/test_opc.py b/tests/laborious/activities/test_opc.py index 0da8214..3fd2dac 100644 --- a/tests/laborious/activities/test_opc.py +++ b/tests/laborious/activities/test_opc.py @@ -1,6 +1,6 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, MagicMock, call, patch -import pytest_asyncio +import pytest from pandas import DataFrame from pytest import mark from sientia_do.notifications.models import NotificationLevel @@ -23,39 +23,32 @@ def test__init__(): opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) assert opc.opc_servers == servers assert opc.opc_repository == {} -@mark.asyncio @patch('laborious.activities.opc.OpcRepository') -@patch('laborious.activities.opc.OPC.send_notification_async') -async def test_init_opc(mock_send_notification, mock_opc_repository): +@patch('laborious.activities.opc.OPC.send_notification') +def test_init_opc(mock_send_notification, mock_opc_repository): mock_logger = MagicMock() - mock_metrics_controller = AsyncMock() - server1 = MagicMock( - connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {})) - ) - server2 = MagicMock( - connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {})) - ) - server3 = MagicMock( - connect=AsyncMock( - return_value=( - False, - { - 'notification_id': 'OPC_CONNECTION_ERROR_server3', - 'message': 'Failed to connect to OPC server: Test error', - 'block': 'opc_repository', - 'level': NotificationLevel.ERROR, - 'attachment_content': 'Test error', - }, - ) - ), - write_data=AsyncMock(return_value=(True, {})), + mock_metrics_controller = MagicMock() + server1 = MagicMock() + server1.connect.return_value = (True, {}) + server2 = MagicMock() + server2.connect.return_value = (True, {}) + server3 = MagicMock() + server3.connect.return_value = ( + False, + { + 'notification_id': 'OPC_CONNECTION_ERROR_server3', + 'message': 'Failed to connect to OPC server: Test error', + 'block': 'opc_repository', + 'level': NotificationLevel.ERROR, + 'attachment_content': 'Test error', + }, ) mock_opc_repository.side_effect = [server1, server2, server3] mock_notification_handler = MagicMock() @@ -97,7 +90,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository): notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) - await opc.init_opc() + opc.init_opc() assert opc.opc_servers == servers assert opc.logger == mock_logger @@ -112,13 +105,12 @@ async def test_init_opc(mock_send_notification, mock_opc_repository): server_name='server1', url='http://localhost:8080', logger=mock_logger, + notification_handler=mock_notification_handler, + metrics_controller=mock_metrics_controller, server_uri='opc.tcp://localhost:4840', cert_path='', private_key_path='', server_cert_path='', - notification_handler=mock_notification_handler, - reconnection_interval=60, - metrics_controller=mock_metrics_controller, ), ] ) @@ -129,13 +121,12 @@ async def test_init_opc(mock_send_notification, mock_opc_repository): server_name='server2', url='http://localhost:8080', logger=mock_logger, + notification_handler=mock_notification_handler, + metrics_controller=mock_metrics_controller, server_uri='opc.tcp://localhost:4840', cert_path='', private_key_path='', server_cert_path='', - notification_handler=mock_notification_handler, - reconnection_interval=60, - metrics_controller=mock_metrics_controller, ) ] ) @@ -162,51 +153,49 @@ async def test_init_opc(mock_send_notification, mock_opc_repository): ) -@pytest_asyncio.fixture -@patch('laborious.activities.opc.OpcRepository') -async def opc(mock_opc_repository): - servers = { - 'server1': { - 'id': 'server1', - 'server_name': 'server1', - 'url': 'http://localhost:8080', - 'server_uri': 'opc.tcp://localhost:4840', - 'cert_path': '', - 'private_key_path': '', - 'server_cert_path': '', - 'reconnection_interval': 60, +@pytest.fixture +def opc(): + with patch('laborious.activities.opc.OpcRepository') as mock_opc_repository: + mock_opc_repository.return_value.write_data = MagicMock(return_value=(True, {})) + mock_opc_repository.return_value.connect = MagicMock(return_value=(True, {})) + servers = { + 'server1': { + 'id': 'server1', + 'server_name': 'server1', + 'url': 'http://localhost:8080', + 'server_uri': 'opc.tcp://localhost:4840', + 'cert_path': '', + 'private_key_path': '', + 'server_cert_path': '', + 'reconnection_interval': 60, + } } - } - mock_opc_repository.return_value.write_data = AsyncMock(return_value=(True, {})) - mock_opc_repository.return_value.connect = AsyncMock(return_value=(True, {})) - opc = OPC( - opc_servers=servers, - logger=MagicMock(), - notification_handler=MagicMock(), - metrics_controller=AsyncMock(), - ) - await opc.init_opc() - opc.send_notification = MagicMock() - opc.send_notification_async = AsyncMock() - opc.emit_metric = AsyncMock() - return opc + opc_instance = OPC( + opc_servers=servers, + logger=MagicMock(), + notification_handler=MagicMock(), + metrics_controller=MagicMock(), + ) + opc_instance.init_opc() + opc_instance.send_notification = MagicMock() + opc_instance.emit_metric_sync = MagicMock() + yield opc_instance WRITE_DATA_CASES = [ ('tag1', 'int', 50), ('tag2', 'float', 50.5), ('tag3', 'bool', True), - ('tag4', 'string', 'test'), + ('tag4', 'str', 'test'), ] @mark.parametrize('tag,data_type,data', WRITE_DATA_CASES) -@mark.asyncio -async def test_write_data_success(opc, tag, data_type, data): +def test_write_data_success(opc, tag, data_type, data): opc.opc_repository['server1'].write_data.return_value = (True, {'response_time': 0.1}) - result = await opc.write_data( + result = opc.write_data( server_id='server1', tag=tag, data=data, @@ -220,8 +209,7 @@ async def test_write_data_success(opc, tag, data_type, data): ) -@mark.asyncio -async def test_write_data_failed(opc): +def test_write_data_failed(opc): opc.opc_repository['server1'].write_data.return_value = ( False, { @@ -233,7 +221,7 @@ async def test_write_data_failed(opc): }, ) - result = await opc.write_data( + result = opc.write_data( server_id='server1', tag='tag1', data=50, @@ -243,7 +231,7 @@ async def test_write_data_failed(opc): ) assert result is None - opc.send_notification_async.assert_called_once_with( + opc.send_notification.assert_called_once_with( metadata=metadata, notification_id='OPC_WRITE_DATA_ERROR_server1', message='Failed to write data to OPC server: Test error', @@ -253,12 +241,11 @@ async def test_write_data_failed(opc): ) -@mark.asyncio -async def test_write_data_exception(opc): +def test_write_data_exception(opc): opc.opc_repository['server1'].write_data.side_effect = Exception('Test error') try: - await opc.write_data( + opc.write_data( server_id='server1', tag='tag1', data=50, @@ -268,7 +255,7 @@ async def test_write_data_exception(opc): ) except Exception: - opc.send_notification_async.assert_called_once_with( + opc.send_notification.assert_called_once_with( metadata=metadata, notification_id='WRITE_OPC_PREDICTION_ERROR', message='Error writing data to OPC server: Test error', @@ -281,9 +268,8 @@ async def test_write_data_exception(opc): raise AssertionError('Expected an exception to be raised') -@mark.asyncio -async def test_manage_output_tags_success(opc): - opc.write_data = AsyncMock(return_value=0.1) +def test_manage_output_tags_success(opc): + opc.write_data = MagicMock(return_value=0.1) data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]}) config = { @@ -291,7 +277,7 @@ async def test_manage_output_tags_success(opc): 'confidence_tags': {'tag2': {'data_type': 'float'}}, } - output_data, opc_metrics = await opc.manage_output_tags( + output_data, opc_metrics = opc.manage_output_tags( server_id='server1', config=config, data=data, @@ -322,16 +308,15 @@ async def test_manage_output_tags_success(opc): ) -@mark.asyncio @mark.parametrize('side_effect', [[0.1, None], [None, 0.2]]) -async def test_manage_output_tags_failed(opc, side_effect): - opc.write_data = AsyncMock(side_effect=side_effect) +def test_manage_output_tags_failed(opc, side_effect): + opc.write_data = MagicMock(side_effect=side_effect) data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]}) config = { 'prediction_tags': {'tag1': {'data_type': 'float'}}, 'confidence_tags': {'tag2': {'data_type': 'float'}}, } - output_data, opc_metrics = await opc.manage_output_tags( + output_data, opc_metrics = opc.manage_output_tags( server_id='server1', config=config, data=data, @@ -361,14 +346,13 @@ async def test_manage_output_tags_failed(opc, side_effect): ) -@mark.asyncio -async def test_manage_output_tags_do_nothing(opc): - opc.write_data = AsyncMock(return_value=0.1) +def test_manage_output_tags_do_nothing(opc): + opc.write_data = MagicMock(return_value=0.1) data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]}) config = { '_invalid_key': {'tag1': {'data_type': 'float'}}, } - output_data, opc_metrics = await opc.manage_output_tags( + output_data, opc_metrics = opc.manage_output_tags( server_id='server1', config=config, data=data, @@ -379,11 +363,9 @@ async def test_manage_output_tags_do_nothing(opc): opc.write_data.assert_not_called() -@mark.asyncio @patch('laborious.activities.opc.DataFrame') -async def test_write_opc_data_success(mock_dataframe, opc): - # Arrange - input_data = { +def test_write_opc_data_success(mock_dataframe, opc): + input_data: dict[str, object] = { **metadata, 'data': {'prediction': [0.75], 'prediction_confidence': [0.95]}, 'opc_output_config': { @@ -393,19 +375,19 @@ async def test_write_opc_data_success(mock_dataframe, opc): } }, } + opc_output_config = input_data['opc_output_config'] + assert isinstance(opc_output_config, dict) - # Act - opc.manage_output_tags = AsyncMock(return_value=(True, {'tag1': 0.1, 'tag2': 0.2})) + opc.manage_output_tags = MagicMock(return_value=(True, {'tag1': 0.1, 'tag2': 0.2})) opc.process_confidence = MagicMock(return_value={'data': 'data'}) - output_data, opc_metrics = await opc.write_opc_data(input_data) + output_data, opc_metrics = opc.write_opc_data(input_data) - # Assert assert output_data == {'data': 'data'} assert opc_metrics == {'server1': {'tag1': 0.1, 'tag2': 0.2}} opc.manage_output_tags.assert_called_once_with( 'server1', - input_data['opc_output_config']['server1'], + opc_output_config['server1'], mock_dataframe.return_value, metadata['metadata'], ) @@ -416,9 +398,7 @@ async def test_write_opc_data_success(mock_dataframe, opc): ) -@mark.asyncio -async def test_write_opc_data_empty_config(opc): - # Arrange +def test_write_opc_data_empty_config(opc): input_data = { **metadata, 'data': {'prediction': [0.75], 'prediction_confidence': [0.95]}, @@ -426,16 +406,13 @@ async def test_write_opc_data_empty_config(opc): 'opc_output_config': {'server1': {'prediction_tags': {}, 'confidence_tags': {}}}, } - # Act - await opc.write_opc_data(input_data) + opc.write_opc_data(input_data) - # Assert opc.opc_repository['server1'].write_data.assert_not_called() -@mark.asyncio -async def test_write_opc_data_no_validate_server(opc): - opc.validate_server = AsyncMock(return_value=False) +def test_write_opc_data_no_validate_server(opc): + opc.validate_server = MagicMock(return_value=False) input_data = { **metadata, 'data': {'prediction': [0.75], 'prediction_confidence': [0.95]}, @@ -447,10 +424,8 @@ async def test_write_opc_data_no_validate_server(opc): }, } - # Act - await opc.write_opc_data(input_data) + opc.write_opc_data(input_data) - # Assert opc.opc_repository['server1'].write_data.assert_not_called() @@ -462,21 +437,18 @@ async def test_write_opc_data_no_validate_server(opc): ], ) def test_process_confidence(opc, data, success, expected): - # Act result = opc.process_confidence(data, success, metadata) - # Assert assert result['prediction_confidence'][0] == expected -@mark.asyncio -async def test_validate_server(opc): - assert await opc.validate_server('server1', metadata) is True - assert await opc.validate_server('server2', metadata) is False +def test_validate_server(opc): + assert opc.validate_server('server1', metadata) is True + assert opc.validate_server('server2', metadata) is False -@mark.asyncio -async def test_close(opc): - opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True) - await opc.close() - opc.opc_repository['server1'].disconnect.assert_called_once() +def test_close(opc): + disconnect_mock = MagicMock(return_value=None) + opc.opc_repository['server1'].disconnect = disconnect_mock + opc.close() + disconnect_mock.assert_called_once() diff --git a/tests/laborious/activities/test_storage.py b/tests/laborious/activities/test_storage.py index ce99ff7..5438228 100644 --- a/tests/laborious/activities/test_storage.py +++ b/tests/laborious/activities/test_storage.py @@ -1,10 +1,11 @@ import datetime import os -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import ANY, MagicMock, patch -from pytest import fixture, mark, raises +from pytest import fixture, raises from sientia_do.notifications.models import NotificationLevel -from sientia_do.temporal.activities.postgres import Postgres +from sientia_do.observability.sientia_monitoring import SientiaMonitoring +from sientia_do.temporal.activities.postgres_sync import Postgres from laborious.activities.storage import Storage @@ -17,6 +18,15 @@ def _passthrough_from_dict(): yield +@fixture(autouse=True) +def _patch_monitoring_shutdown(): + """ + Avoid running real async SientiaMonitoring.shutdown when Storage.close runs inside tests. + """ + with patch.object(SientiaMonitoring, 'shutdown') as mock_shutdown: + yield mock_shutdown + + metadata = { 'metadata': { 'model_id': 'test_model_id', @@ -42,7 +52,7 @@ def storage(mock_minio_repository): minio_repository=mock_minio_repository.return_value, logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) @@ -50,7 +60,7 @@ def storage(mock_minio_repository): def test___init___not_hasattr(mock_minio_repository): logger = MagicMock() notification_handler = MagicMock() - metrics_controller = AsyncMock() + metrics_controller = MagicMock() minio_repo = mock_minio_repository.return_value storage = Storage( host='localhost', @@ -77,7 +87,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage): storage.minio_repository = None logger = MagicMock() notification_handler = MagicMock() - metrics_controller = AsyncMock() + metrics_controller = MagicMock() storage.__init__( host='localhost', port=5432, @@ -111,54 +121,55 @@ def test___init___done_repository(mock_minio_repository, storage): minio_repository=mock_minio_repository.return_value, logger=MagicMock(), notification_handler=MagicMock(), - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) mock_minio_repository.assert_not_called() assert storage.minio_repository is not None -def test_close(storage): +def test_close(storage, _patch_monitoring_shutdown): storage.minio_repository = MagicMock() storage.close() assert storage.minio_repository is None + _patch_monitoring_shutdown.assert_called_once_with(storage) -def test___del__(storage): - storage.close = MagicMock() +def test_close_when_minio_repository_already_none(storage, _patch_monitoring_shutdown): + """Closing without an initialized MinIO repository skips MinIO teardown.""" + storage.minio_repository = None - storage.__del__() + storage.close() - storage.close.assert_called_once() + assert storage.minio_repository is None + _patch_monitoring_shutdown.assert_called_once_with(storage) -@mark.asyncio -async def test_load_query_with_minio_offload_no_rows(storage): - storage.load_custom_query = AsyncMock(return_value=None) +def test_load_query_with_minio_offload_no_rows(storage): + storage.load_custom_query = MagicMock(return_value=None) storage_result = {'success': False} with patch( 'laborious.activities.storage.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, return_value=storage_result, ) as mock_from_dataframe: - result = await storage.load_query_with_minio_offload( + result = storage.load_query_with_minio_offload( {**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'} ) assert result == storage_result - mock_from_dataframe.assert_awaited_once() + mock_from_dataframe.assert_called_once() -@mark.asyncio -async def test_load_query_with_minio_offload_inline(storage): - storage.load_custom_query = AsyncMock(return_value=[{'a': 1}]) +def test_load_query_with_minio_offload_inline(storage): + storage.load_custom_query = MagicMock(return_value=[{'a': 1}]) storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None} with patch( 'laborious.activities.storage.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, return_value=storage_result, ) as mock_from_dataframe: - result = await storage.load_query_with_minio_offload( + result = storage.load_query_with_minio_offload( { **metadata, 'query': 'SELECT 1', @@ -167,44 +178,42 @@ async def test_load_query_with_minio_offload_inline(storage): } ) assert result == storage_result - mock_from_dataframe.assert_awaited_once() + mock_from_dataframe.assert_called_once() -@mark.asyncio -async def test_load_query_with_minio_offload_minio(storage): - storage.load_custom_query = AsyncMock(return_value=[{'a': 1}]) +def test_load_query_with_minio_offload_minio(storage): + storage.load_custom_query = MagicMock(return_value=[{'a': 1}]) storage_result = {'success': True, 'data': None, 'object_key': 'object-key'} with patch( 'laborious.activities.storage.MinioDataFramePayload.from_dataframe', - new_callable=AsyncMock, + new_callable=MagicMock, return_value=storage_result, ) as mock_from_dataframe: - result = await storage.load_query_with_minio_offload( + result = storage.load_query_with_minio_offload( {**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'} ) assert result == storage_result - mock_from_dataframe.assert_awaited_once() + mock_from_dataframe.assert_called_once() -@mark.asyncio @patch.dict(os.environ, {'SIENTIA_MINIO_RETENTION_HOURS': '1'}) @patch('laborious.activities.storage.now') -async def test_cleanup_minio_objects_expired(mock_now, storage): +def test_cleanup_minio_objects_expired(mock_now, storage): mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0) - storage.minio_repository.list_objects = AsyncMock( + storage.minio_repository.list_objects = MagicMock( return_value=[ 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet', 'sientia/streamlit-connectors/training_datasets/m/m-initial-2025-01-10_12-00-00.parquet', ] ) - storage.minio_repository.delete_file = AsyncMock() - storage.send_notification_async = AsyncMock() + storage.minio_repository.delete_file = MagicMock() + storage.send_notification = MagicMock() data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) + result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 1 assert result['failed_count'] == 0 @@ -224,72 +233,65 @@ async def test_cleanup_minio_objects_expired(mock_now, storage): ) -@mark.asyncio -async def test_load_query_with_minio_offload_minio_not_initialized(storage): +def test_load_query_with_minio_offload_minio_not_initialized(storage): storage.minio_repository = None with raises(ValueError, match='Minio repository not initialized'): - await storage.load_query_with_minio_offload( - {**metadata, 'query': 'SELECT 1', 'model_name': 'm'} - ) + storage.load_query_with_minio_offload({**metadata, 'query': 'SELECT 1', 'model_name': 'm'}) -@mark.asyncio -async def test_export_payload_to_postgres(storage): - payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=MagicMock()) - storage.export_data_to_postgres = AsyncMock(return_value={'success': True}) +def test_export_payload_to_postgres(storage): + payload = MagicMock() + payload.retrieve = MagicMock(return_value=MagicMock()) + storage.export_data_to_postgres = MagicMock(return_value={'success': True}) - result = await storage.export_payload_to_postgres( + result = storage.export_payload_to_postgres( {**metadata, 'data': payload, 'schema': 'public', 'table': 't'} ) - payload.retrieve.assert_awaited_once_with(storage.minio_repository, metadata['metadata']) - storage.export_data_to_postgres.assert_awaited_once() + payload.retrieve.assert_called_once_with(storage.minio_repository, metadata['metadata']) + storage.export_data_to_postgres.assert_called_once() assert result == {'success': True} -@mark.asyncio -async def test_cleanup_minio_objects_expired_minio_not_initialized(storage): +def test_cleanup_minio_objects_expired_minio_not_initialized(storage): storage.minio_repository = None data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'test' with raises(ValueError, match='Minio repository not initialized'): - await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) + storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) -@mark.asyncio @patch('laborious.activities.storage.now') -async def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage): +def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage): mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0) - storage.minio_repository.list_objects = AsyncMock( + storage.minio_repository.list_objects = MagicMock( return_value=['some/random/key-without-timestamp.parquet'] ) - storage.minio_repository.delete_file = AsyncMock() - storage.send_notification_async = AsyncMock() + storage.minio_repository.delete_file = MagicMock() + storage.send_notification = MagicMock() data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'test' - result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) + result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 0 assert result['failed_count'] == 0 storage.minio_repository.delete_file.assert_not_called() -@mark.asyncio @patch('laborious.activities.storage.now') -async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage): +def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage): mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0) old_key = 'training_datasets/m/m-initial-2024-12-01_00-00-00.parquet' - storage.minio_repository.list_objects = AsyncMock(return_value=[old_key]) - storage.minio_repository.delete_file = AsyncMock(side_effect=Exception('delete error')) - storage.send_notification_async = AsyncMock() + storage.minio_repository.list_objects = MagicMock(return_value=[old_key]) + storage.minio_repository.delete_file = MagicMock(side_effect=Exception('delete error')) + storage.send_notification = MagicMock() data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) + result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 0 assert result['failed_count'] == 1 @@ -298,21 +300,20 @@ async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage): assert result['failed'][old_key]['message'] == 'delete error' -@mark.asyncio @patch('laborious.activities.storage.now') -async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage): +def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage): mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0) - storage.minio_repository.list_objects = AsyncMock(side_effect=Exception('list error')) - storage.send_notification_async = AsyncMock() + storage.minio_repository.list_objects = MagicMock(side_effect=Exception('list error')) + storage.send_notification = MagicMock() storage.error = MagicMock() data_mock = MagicMock() data_mock.cleanup_prefix.return_value = 'training_datasets/m' - result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) + result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock}) assert result['deleted_count'] == 0 assert result['failed_count'] == 0 - storage.send_notification_async.assert_called_once_with( + storage.send_notification.assert_called_once_with( metadata=metadata['metadata'], notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED', message='Error cleaning up MinIO objects: list error', diff --git a/tests/laborious/utils/models/test_minio_dataframe_payload.py b/tests/laborious/utils/models/test_minio_dataframe_payload.py index 3414056..b358611 100644 --- a/tests/laborious/utils/models/test_minio_dataframe_payload.py +++ b/tests/laborious/utils/models/test_minio_dataframe_payload.py @@ -1,8 +1,7 @@ from datetime import datetime from io import BytesIO -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock, patch -import pytest from pandas import DataFrame from laborious.utils.models.minio_dataframe_payload import ( @@ -54,17 +53,15 @@ def test_has_data_true_when_object_key_set(): assert payload.has_data() is True -@pytest.mark.asyncio -async def test_retrieve_inline_dict_as_dataframe(): +def test_retrieve_inline_dict_as_dataframe(): payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]}) - minio = AsyncMock() - out = await payload.retrieve(minio, {'metadata': {}}) + minio = MagicMock() + out = payload.retrieve(minio, {'metadata': {}}) assert list(out.columns) == ['a'] minio.download_file.assert_not_called() -@pytest.mark.asyncio -async def test_retrieve_downloads_parquet_when_offloaded(): +def test_retrieve_downloads_parquet_when_offloaded(): source = DataFrame({'a': [1, 2]}) buf = BytesIO() source.to_parquet(buf, engine='pyarrow', index=True) @@ -76,12 +73,12 @@ async def test_retrieve_downloads_parquet_when_offloaded(): object_key='training_datasets/m/f.parquet', object_prefix='training_datasets/m', ) - minio = AsyncMock() - minio.download_file = AsyncMock(return_value=file_bytes) + minio = MagicMock() + minio.download_file = MagicMock(return_value=file_bytes) - out = await payload.retrieve(minio, {'metadata': {}}) + out = payload.retrieve(minio, {'metadata': {}}) - minio.download_file.assert_awaited_once_with( + minio.download_file.assert_called_once_with( object_name='training_datasets/m/f.parquet', metadata={'metadata': {}}, ) @@ -113,21 +110,19 @@ def test_parse_object_timestamp_bad_datetime(): assert MinioDataFramePayload.parse_object_timestamp(key) is None -@pytest.mark.asyncio -async def test_retrieve_empty_when_no_data(): +def test_retrieve_empty_when_no_data(): payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key=None) - minio = AsyncMock() - out = await payload.retrieve(minio, {}) + minio = MagicMock() + out = payload.retrieve(minio, {}) assert out.empty minio.download_file.assert_not_called() -@pytest.mark.asyncio @patch('laborious.utils.models.minio_dataframe_payload.now') -async def test_from_dataframe_none(mock_now): +def test_from_dataframe_none(mock_now): mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0) - minio = AsyncMock() - result = await MinioDataFramePayload.from_dataframe( + minio = MagicMock() + result = MinioDataFramePayload.from_dataframe( dataframe=None, minio_repo=minio, model_name='m', @@ -139,15 +134,14 @@ async def test_from_dataframe_none(mock_now): assert result.object_key is None -@pytest.mark.asyncio @patch('laborious.utils.models.minio_dataframe_payload.now') -async def test_from_dataframe_empty(mock_now): +def test_from_dataframe_empty(mock_now): mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0) - minio = AsyncMock() + minio = MagicMock() mock_df = MagicMock() mock_df.__bool__ = MagicMock(return_value=True) mock_df.empty = True - result = await MinioDataFramePayload.from_dataframe( + result = MinioDataFramePayload.from_dataframe( dataframe=mock_df, minio_repo=minio, model_name='m', @@ -174,12 +168,11 @@ def _mock_dataframe(data_dict, timestamp_values=None): return mock_df -@pytest.mark.asyncio @patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9) -async def test_from_dataframe_inline(): - minio = AsyncMock() +def test_from_dataframe_inline(): + minio = MagicMock() df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]}) - result = await MinioDataFramePayload.from_dataframe( + result = MinioDataFramePayload.from_dataframe( dataframe=df, minio_repo=minio, model_name='m', @@ -190,12 +183,11 @@ 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() +def test_from_dataframe_inline_uses_provided_last_timestamp(): + minio = MagicMock() df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]}) - result = await MinioDataFramePayload.from_dataframe( + result = MinioDataFramePayload.from_dataframe( dataframe=df, minio_repo=minio, model_name='m', @@ -205,17 +197,16 @@ async def test_from_dataframe_inline_uses_provided_last_timestamp(): 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) -async def test_from_dataframe_offloaded(mock_now): +def test_from_dataframe_offloaded(mock_now): mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0) - minio = AsyncMock() - minio.upload_file = AsyncMock(return_value={'minio_object_name': 'full/key.parquet'}) + minio = MagicMock() + minio.upload_file = MagicMock(return_value={'minio_object_name': 'full/key.parquet'}) minio.bucket = 'test-bucket' df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]}) - result = await MinioDataFramePayload.from_dataframe( + result = MinioDataFramePayload.from_dataframe( dataframe=df, minio_repo=minio, model_name='m', @@ -226,7 +217,7 @@ async def test_from_dataframe_offloaded(mock_now): assert result.object_key == 'full/key.parquet' assert result.bucket == 'test-bucket' assert result.uri == 's3://test-bucket/full/key.parquet' - minio.upload_file.assert_awaited_once() + minio.upload_file.assert_called_once() def test_from_dict_inline(): diff --git a/tests/laborious/utils/repository/test_opc_repository.py b/tests/laborious/utils/repository/test_opc_repository.py index 8be96a7..88390fa 100644 --- a/tests/laborious/utils/repository/test_opc_repository.py +++ b/tests/laborious/utils/repository/test_opc_repository.py @@ -1,9 +1,9 @@ import json from datetime import datetime -from unittest.mock import ANY, AsyncMock, MagicMock, Mock, patch +from unittest.mock import ANY, MagicMock, Mock, patch import pytest -from asyncua.crypto.security_policies import SecurityPolicyBasic256 +from opcua.crypto import security_policies from sientia_do.notifications.models import NotificationLevel from laborious.utils.repository.opc_repository import OpcRepository @@ -27,19 +27,18 @@ def opc_repository(mock_logger): cert_path='/path/to/cert.pem', private_key_path='/path/to/key.pem', server_cert_path='/path/to/server_cert.pem', - metrics_controller=AsyncMock(), + metrics_controller=MagicMock(), ) repository.disconnection_interval = 0.1 repository.send_notification = MagicMock() - repository.send_notification_async = AsyncMock() - repository.emit_metric = AsyncMock() + repository.emit_metric_sync = MagicMock() return repository @pytest.fixture def mock_client(): with patch('laborious.utils.repository.opc_repository.Client') as mock: - client_instance = AsyncMock() + client_instance = MagicMock() mock.return_value = client_instance yield client_instance @@ -68,58 +67,48 @@ def test_init(opc_repository): assert opc_repository.error_count == 0 -@pytest.mark.asyncio -async def test_set_security(opc_repository, mock_client): +def test_set_security(opc_repository, mock_client): opc_repository.client = mock_client - await opc_repository.set_security() + opc_repository.set_security() - mock_client.application_uri = 'urn:test:server' mock_client.set_security.assert_called_once_with( - SecurityPolicyBasic256, - certificate='/path/to/cert.pem', - private_key='/path/to/key.pem', - server_certificate='/path/to/server_cert.pem', + security_policies.SecurityPolicyBasic256, + '/path/to/cert.pem', + '/path/to/key.pem', + '/path/to/server_cert.pem', ) assert mock_client.secure_channel_timeout == 10000000 assert mock_client.session_timeout == 10000000 -@pytest.mark.asyncio -async def test_set_security_missing_certificates(opc_repository): +def test_set_security_missing_certificates(opc_repository): opc_repository.cert_path = None opc_repository.private_key_path = None - try: - await opc_repository.set_security() - except ValueError as e: - assert str(e) == 'Certificate and private key paths must be provided for secure connection.' + with pytest.raises(ValueError, match='Certificate and private key paths'): + opc_repository.set_security() -@pytest.mark.asyncio -async def test_set_security_missing_client(opc_repository): +def test_set_security_missing_client(opc_repository): opc_repository.client = None - try: - await opc_repository.set_security() - except ValueError as e: - assert str(e) == 'Client must be initialized before setting security' + with pytest.raises(ValueError, match='Client must be initialized'): + opc_repository.set_security() -@pytest.mark.asyncio -async def test_connect_with_security(opc_repository, mock_client): - opc_repository.try_connect = AsyncMock(return_value=(True, {})) - result = await opc_repository.connect() +def test_connect_with_security(opc_repository, mock_client): + opc_repository.try_connect = MagicMock(return_value=(True, {})) + result = opc_repository.connect() opc_repository.try_connect.assert_called_once() assert opc_repository.client == mock_client assert result == (True, {}) -@pytest.mark.asyncio -async def test_connect_without_security(opc_repository, mock_client): +def test_connect_without_security(opc_repository, mock_client): opc_repository.cert_path = None - opc_repository.try_connect = AsyncMock(return_value=(True, {})) - opc_repository.set_security = AsyncMock() - result = await opc_repository.connect() + opc_repository.try_connect = MagicMock(return_value=(True, {})) + opc_repository.set_security = MagicMock() + result = opc_repository.connect() opc_repository.try_connect.assert_called_once() opc_repository.set_security.assert_not_called() @@ -127,25 +116,23 @@ async def test_connect_without_security(opc_repository, mock_client): assert result == (True, {}) -@pytest.mark.asyncio -async def test_try_connect_success(opc_repository): +def test_try_connect_success(opc_repository): opc_repository.last_reconnection_time = None - opc_repository.client = AsyncMock() - result = await opc_repository.try_connect() + opc_repository.client = MagicMock() + result = opc_repository.try_connect() opc_repository.client.connect.assert_called_once() assert opc_repository.last_reconnection_time is not None assert result == (True, {}) -@pytest.mark.asyncio -async def test_try_connect_fail(opc_repository): +def test_try_connect_fail(opc_repository): opc_repository.last_reconnection_time = None - opc_repository.disconnect = AsyncMock() + opc_repository.disconnect = MagicMock() opc_repository.client = MagicMock() opc_repository.client.connect.side_effect = Exception('Test error') - is_connected, error_data = await opc_repository.try_connect() + is_connected, error_data = opc_repository.try_connect() opc_repository.disconnect.assert_called_once() opc_repository.client.connect.assert_called_once() @@ -157,10 +144,9 @@ async def test_try_connect_fail(opc_repository): assert error_data['attachment_content'] is not None -@pytest.mark.asyncio -async def test_try_connect_no_client(opc_repository): +def test_try_connect_no_client(opc_repository): opc_repository.client = None - result = await opc_repository.try_connect() + result = opc_repository.try_connect() assert result == ( False, { @@ -172,21 +158,24 @@ async def test_try_connect_no_client(opc_repository): ) -@pytest.mark.asyncio -async def test_disconnection_fallback_success(opc_repository, mock_client): +def test_session_alive_returns_false_when_client_none(opc_repository): + opc_repository.client = None + assert opc_repository._session_alive() is False + + +def test_disconnection_fallback_success(opc_repository, mock_client): opc_repository.client = mock_client - mock_client.disconnect.return_value = True - result = await opc_repository.disconnection_fallback() + mock_client.disconnect.return_value = None + result = opc_repository.disconnection_fallback() mock_client.disconnect.assert_called_once() assert result == [] -@pytest.mark.asyncio -async def test_disconnection_fallback_fail(opc_repository, mock_client): +def test_disconnection_fallback_fail(opc_repository, mock_client): opc_repository.client = mock_client mock_client.disconnect.side_effect = Exception('Test error') - result = await opc_repository.disconnection_fallback() + result = opc_repository.disconnection_fallback() assert result == [ {'attempt': 1, 'error': 'Test error', 'traceback': ANY}, {'attempt': 2, 'error': 'Test error', 'traceback': ANY}, @@ -197,32 +186,29 @@ async def test_disconnection_fallback_fail(opc_repository, mock_client): assert mock_client.disconnect.call_count == 5 -@pytest.mark.asyncio -async def test_disconnect(opc_repository, mock_client): +def test_disconnect(opc_repository, mock_client): opc_repository.client = mock_client - opc_repository.disconnection_fallback = AsyncMock(return_value=[]) - await opc_repository.disconnect() + opc_repository.disconnection_fallback = MagicMock(return_value=[]) + opc_repository.disconnect() opc_repository.disconnection_fallback.assert_called_once() assert opc_repository.client is None -@pytest.mark.asyncio -async def test_disconnect_no_client(opc_repository): +def test_disconnect_no_client(opc_repository): opc_repository.client = None - assert await opc_repository.disconnect() is None + assert opc_repository.disconnect() is None -@pytest.mark.asyncio -async def test_disconnect_error(opc_repository, mock_client): +def test_disconnect_error(opc_repository, mock_client): opc_repository.client = mock_client - opc_repository.disconnection_fallback = AsyncMock( + opc_repository.disconnection_fallback = MagicMock( return_value=[{'attempt': 1, 'error': 'Test error', 'traceback': 'text'}] ) - await opc_repository.disconnect() + opc_repository.disconnect() opc_repository.disconnection_fallback.assert_called_once() - opc_repository.send_notification_async.assert_called_once_with( + opc_repository.send_notification.assert_called_once_with( metadata=opc_repository.metadata, notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}', message='Failed to disconnect from OPC server in 5 attempts.', @@ -235,63 +221,39 @@ async def test_disconnect_error(opc_repository, mock_client): assert opc_repository.client is None -@pytest.mark.asyncio -async def test_validate_connection_none_client(opc_repository): +def test_validate_connection_none_client(opc_repository): opc_repository.client = None - opc_repository.connect = AsyncMock(return_value=(True, {})) - response = await opc_repository.validate_connection() + opc_repository.connect = MagicMock(return_value=(True, {})) + response = opc_repository.validate_connection() assert response == (True, {}) opc_repository.connect.assert_called_once() -# @pytest.mark.asyncio -# async def test_validate_connection_error_count_disconnect_error(opc_repository): -# opc_repository.error_count = 6 -# opc_repository.client = AsyncMock() -# opc_repository.disconnect = AsyncMock(side_effect=Exception('Test error')) -# opc_repository.connect = AsyncMock(return_value=(True, {})) +def test_validate_connection_disconnect_raises(opc_repository): + """Outer except path when reconnect cleanup fails mid-validation.""" -# response = await opc_repository.validate_connection() -# assert response == opc_repository.connect.return_value -# opc_repository.disconnect.assert_called_once() -# opc_repository.connect.assert_called_once() -# opc_repository.logger.custom_error.assert_has_calls( -# [ -# call('Failed to disconnect from OPC server: Test error', ANY), -# ] -# ) + opc_repository.client = MagicMock() + opc_repository._session_alive = MagicMock(return_value=False) + opc_repository.last_reconnection_time = datetime(2020, 1, 1, 0, 0, 0) + opc_repository.disconnect = MagicMock(side_effect=RuntimeError('disconnect failed')) + + response = opc_repository.validate_connection() + + assert response[0] is False + assert response[1]['notification_id'] == f'OPC_CONNECTION_CHECK_ERROR_{opc_repository.id}' + assert 'disconnect failed' in response[1]['message'] -@pytest.mark.asyncio -async def test_validate_connection_error_validate_connection_error(opc_repository): - opc_repository.client = MagicMock(uaclient=Exception('Test error')) - opc_repository.error_count = 0 - - response = await opc_repository.validate_connection() - - assert response == ( - False, - { - 'notification_id': f'OPC_CONNECTION_CHECK_ERROR_{opc_repository.id}', - 'message': "Failed to validate connection to OPC server: 'Exception' object has no attribute 'protocol'", - 'block': 'opc_repository', - 'level': NotificationLevel.ERROR, - 'attachment_content': ANY, - }, - ) - - -@pytest.mark.asyncio @patch('laborious.utils.repository.opc_repository.datetime') -async def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, opc_repository): +def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, opc_repository): _mock_datetime.now = MagicMock(return_value=datetime(2025, 1, 1, 0, 0, 0)) opc_repository.error_count = 0 opc_repository.client = MagicMock() - opc_repository.client.uaclient.protocol = None + opc_repository.client.get_root_node.side_effect = RuntimeError('down') opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0) opc_repository.connect = MagicMock(return_value=(True, {})) - response = await opc_repository.validate_connection() + response = opc_repository.validate_connection() opc_repository.connect.assert_not_called() assert response == ( False, @@ -304,55 +266,51 @@ async def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, op ) -@pytest.mark.asyncio @patch('laborious.utils.repository.opc_repository.datetime') -async def test_validate_connection_lost_time_to_reconnect(mock_datetime, opc_repository): +def test_validate_connection_lost_time_to_reconnect(mock_datetime, opc_repository): mock_datetime.now = MagicMock(return_value=datetime(2025, 1, 1, 1, 0, 0)) opc_repository.error_count = 0 - opc_repository.client = AsyncMock() - opc_repository.client.uaclient.protocol = None + opc_repository.client = MagicMock() + opc_repository.client.get_root_node.side_effect = RuntimeError('down') opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0) - opc_repository.connect = AsyncMock(return_value=(True, {})) + opc_repository.connect = MagicMock(return_value=(True, {})) - response = await opc_repository.validate_connection() + response = opc_repository.validate_connection() opc_repository.connect.assert_called_once() assert response == opc_repository.connect.return_value -@pytest.mark.asyncio -async def test_validate_connection_success(opc_repository): +def test_validate_connection_success(opc_repository): opc_repository.client = MagicMock() opc_repository.error_count = 0 - opc_repository.client.uaclient.protocol = MagicMock() - opc_repository.client.uaclient.protocol.state = 'open' + opc_repository.client.get_root_node.return_value = MagicMock() - output = await opc_repository.validate_connection() + output = opc_repository.validate_connection() assert output == (True, {}) -@pytest.mark.asyncio -async def test_write_data_validate_connection_do_nothing(opc_repository): - opc_repository.validate_connection = AsyncMock(return_value=(True, {})) - opc_repository.client = AsyncMock(get_node=MagicMock()) - mock_node = AsyncMock() +def test_write_data_validate_connection_do_nothing(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=(True, {})) + opc_repository.client = MagicMock() + mock_node = MagicMock() opc_repository.client.get_node.return_value = mock_node - result = await opc_repository.write_data( + result = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata'] ) opc_repository.validate_connection.assert_called_once() opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode') + mock_node.set_value.assert_called_once() assert result == (True, {'response_time': ANY}) -@pytest.mark.asyncio -async def test_write_data_validate_connection_failed(opc_repository): - opc_repository.validate_connection = AsyncMock(return_value=(False, {})) - opc_repository.client = AsyncMock() +def test_write_data_validate_connection_failed(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=(False, {})) + opc_repository.client = MagicMock() opc_repository.error_count = 0 - result = await opc_repository.write_data( + result = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata'] ) @@ -361,14 +319,13 @@ async def test_write_data_validate_connection_failed(opc_repository): assert result == (False, {}) -@pytest.mark.asyncio -async def test_write_data_get_node_failed(opc_repository): - opc_repository.validate_connection = AsyncMock(return_value=(True, {})) - opc_repository.client = AsyncMock() +def test_write_data_get_node_failed(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=(True, {})) + opc_repository.client = MagicMock() opc_repository.error_count = 0 opc_repository.client.get_node = MagicMock(side_effect=Exception('Test error')) - is_success, error_data = await opc_repository.write_data( + is_success, error_data = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata'] ) @@ -376,23 +333,15 @@ async def test_write_data_get_node_failed(opc_repository): opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode') assert is_success is False assert error_data['notification_id'] == f'OPC_WRITE_GET_NODE_ERROR_{opc_repository.id}' - assert ( - error_data['message'] - == "Failed to get node from OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}" - ) - assert error_data['block'] == 'opc_repository' - assert error_data['level'] == NotificationLevel.ERROR - assert error_data['attachment_content'] is not None -@pytest.mark.asyncio -async def test_write_data_invalid_data_type(opc_repository, mock_client): - opc_repository.validate_connection = AsyncMock(return_value=(True, {})) +def test_write_data_invalid_data_type(opc_repository, mock_client): + opc_repository.validate_connection = MagicMock(return_value=(True, {})) opc_repository.client = mock_client - mock_node = AsyncMock() + mock_node = MagicMock() mock_client.get_node = MagicMock(return_value=mock_node) - is_success, error_data = await opc_repository.write_data( + is_success, error_data = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'invalid_type', opc_repository.logger, metadata['metadata'] ) @@ -401,53 +350,37 @@ async def test_write_data_invalid_data_type(opc_repository, mock_client): assert is_success is False assert error_data['notification_id'] == f'OPC_WRITE_DATA_TYPE_ERROR_{opc_repository.id}' - assert ( - error_data['message'] - == "Unsupported data type: invalid_type | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}" - ) - assert error_data['block'] == 'opc_repository' - assert error_data['level'] == NotificationLevel.ERROR - assert error_data.get('attachment_content') is None -@pytest.mark.asyncio -async def test_write_data(opc_repository, mock_client): - opc_repository.validate_connection = AsyncMock(return_value=(True, {})) +def test_write_data(opc_repository, mock_client): + opc_repository.validate_connection = MagicMock(return_value=(True, {})) opc_repository.client = mock_client - mock_node = AsyncMock() + mock_node = MagicMock() mock_client.get_node = MagicMock(return_value=mock_node) - result = await opc_repository.write_data( + result = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata'] ) mock_client.get_node.assert_called_once_with('ns=2;s=TestNode') - mock_node.write_value.assert_called_once() + mock_node.set_value.assert_called_once() assert result == (True, {'response_time': ANY}) -@pytest.mark.asyncio -async def test_write_data_write_value_failed(opc_repository, mock_client): - opc_repository.validate_connection = AsyncMock(return_value=(True, {})) +def test_write_data_write_value_failed(opc_repository, mock_client): + opc_repository.validate_connection = MagicMock(return_value=(True, {})) opc_repository.client = mock_client - mock_node = AsyncMock() + mock_node = MagicMock() opc_repository.error_count = 0 mock_client.get_node = MagicMock(return_value=mock_node) - mock_node.write_value.side_effect = Exception('Test error') + mock_node.set_value.side_effect = Exception('Test error') - is_success, error_data = await opc_repository.write_data( + is_success, error_data = opc_repository.write_data( 'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata'] ) opc_repository.validate_connection.assert_called_once() mock_client.get_node.assert_called_once_with('ns=2;s=TestNode') - mock_node.write_value.assert_called_once() + mock_node.set_value.assert_called_once() assert is_success is False assert error_data['notification_id'] == f'OPC_WRITE_DATA_ERROR_{opc_repository.id}' - assert ( - error_data['message'] - == "Failed to write data to OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}" - ) - assert error_data['block'] == 'opc_repository' - assert error_data['level'] == NotificationLevel.ERROR - assert error_data['attachment_content'] is not None diff --git a/tests/laborious/worker/test_worker.py b/tests/laborious/worker/test_worker.py index 8812a1f..32f8276 100644 --- a/tests/laborious/worker/test_worker.py +++ b/tests/laborious/worker/test_worker.py @@ -7,8 +7,8 @@ from laborious.worker import worker def _build_fake_activities(): inst = MagicMock() - inst.init_opc = AsyncMock() - inst.shutdown = AsyncMock() + inst.init_opc = MagicMock() + inst.shutdown = MagicMock() inst.load_query_with_minio_offload = MagicMock() inst.retrain_model = MagicMock() inst.update_production_model = MagicMock() @@ -176,7 +176,7 @@ async def test_main_success_exit_zero(monkeypatch): notif = m_notif_cls.return_value notif.shutdown.assert_called_once() - fake_activities.shutdown.assert_awaited_once() + fake_activities.shutdown.assert_called_once() assert m_prepare.call_count == 4 prepare_calls = m_prepare.call_args_list assert prepare_calls[0].kwargs['runtime'] == 'single' diff --git a/values.yaml b/values.yaml index 2b1c49d..607d5ff 100644 --- a/values.yaml +++ b/values.yaml @@ -1,120 +1,265 @@ -# Default values for sientia-module. +# +# Default values for sientia-laborious-worker using the sientia-module chart (0.6.x). # This is a YAML-formatted file. # Declare variables to be passed into your templates. +# -# This will set the replicaset count more information can be found here: https://kubernetes.io/docs/concepts/workloads/controllers/replicaset/ -replicaCount: 1 +projectName: &projectName "sientia-laborious-worker" -# This sets the container image more information can be found here: https://kubernetes.io/docs/concepts/containers/images/ -image: - repository: aignosi.azurecr.io/sientia-module - # This sets the pull policy for images. - pullPolicy: Always - # Overrides the image tag whose default is the chart appVersion. - tag: "1.1.2" +# ----------------------------------------------------------------------------- +# Global configuration shared by all runtimes +# ----------------------------------------------------------------------------- +global: + namespace: sientia -# This is for the secrets for pulling an image from a private repository more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/pull-image-private-registry/ + image: + repository: aignosi.azurecr.io/sientia-module + pullPolicy: Always + tag: "1.2.0" + + commonLabels: {} + + resources: + # Resource limits and requests are important for ResourceBasedTuner to work correctly. + # The tuner monitors system CPU and memory usage, so proper resource limits must be set. + limits: + cpu: 2000m + memory: 20Gi + requests: + cpu: 1000m + memory: 2Gi + + livenessProbe: + exec: + command: + - sh + - -c + - | + curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1' + initialDelaySeconds: 1260 + periodSeconds: 15 + timeoutSeconds: 5 + failureThreshold: 3 + + readinessProbe: + exec: + command: + - sh + - -c + - | + curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1' + initialDelaySeconds: 1200 + periodSeconds: 10 + timeoutSeconds: 3 + failureThreshold: 2 + + autoscaling: + enabled: false + minReplicas: 1 + maxReplicas: 100 + targetCPUUtilizationPercentage: 80 + # targetMemoryUtilizationPercentage: 80 + + # Environment variables shared by all runtimes. + env: + # Entrypoint variables + - name: GITHUB_REPO_URL + value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git" + - name: GITHUB_BRANCH + value: "feature/SIENTIAPDE-1646" + - name: PYTHON_APP + value: "laborious.worker.worker" + - name: PYPI_SERVER + value: "http://library-distribution-server.library.svc.cluster.local:5000" + + # Application variables + - name: POSTGRES_HOST + value: "paradedb-rw.paradedb.svc.cluster.local" + - name: POSTGRES_PORT + value: "5432" + - name: POSTGRES_USER + value: "postgres" + - name: POSTGRES_PASSWORD + value: "nFqc81y6kwmr2zuAIx43DhiOosFCVPpeEfTtTWZflkNjB2j1KtEeIANkhFR9mAX3" + - name: POSTGRES_DBNAME + value: "sientia" + - name: POSTGRES_MIN_CONNECTIONS + value: "20" + # max_connections = number_of_workers * max_concurrent_activities * safety_factor + # Example: 4 workers * 50 activities * 0.5 = 100 connections + - name: POSTGRES_MAX_CONNECTIONS + value: "100" + + - name: MLFLOW_URL + value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local:80" + - name: MLFLOW_USERNAME + value: "aignosi" + - name: MLFLOW_PASSWORD + value: "1L0FP50j3ncp123" + + # Plugin store (model-library-store Git + runtime packages). + - name: STORE_BASE_URL + value: "http://gitea-http.gitea.svc.cluster.local:3000" + - name: STORE_OWNER + value: "aignosi" + - name: STORE_REPO + value: "suse-model-store" + - name: STORE_USERNAME + valueFrom: + secretKeyRef: + name: sientia-plugin-store-credentials + key: username + - name: STORE_PASSWORD + valueFrom: + secretKeyRef: + name: sientia-plugin-store-credentials + key: password + - name: STORE_CACHE_TTL_SECONDS + value: "3600" + + - name: OPC_ID + value: "1" + - name: OPC_SERVER_NAME + value: "default_server" + - name: OPC_URL + value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840" + + - name: LOG_LEVEL + value: "DEBUG" + - name: HTTP_METRICS_PORT + value: "9090" + - name: HTTP_SDK_METRICS_PORT + value: "9091" + - name: PROJECT_NAME + value: "sientia-laborious" + + - name: TEMPORAL_HOST + value: "temporal-frontend.temporal.svc.cluster.local:7233" + - name: TEMPORAL_NAMESPACE + value: "laborious" + + - name: MONGODB_USERNAME + value: "root" + - name: MONGODB_PASSWORD + value: "wKZDbMNU1c" + - name: MONGODB_URL + value: "my-release-mongodb.mongodb.svc.cluster.local:27017" + - name: MONGODB_DATABASE + value: "sientia" + - name: MONGODB_TTL_INDEX_HOURS + value: "1" + + - name: MINIO_ENDPOINT_URL + value: "minio.minio.svc.cluster.local:9000" + - name: MINIO_ACCESS_KEY + value: "admin" + - name: MINIO_SECRET_KEY + value: "LiArt4eNmJ" + - name: MINIO_DEFAULT_BUCKET + value: "sientia" + - name: MINIO_RETENTION_HOURS + value: "24" + - name: SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES + value: "0.5" + + # Temporal worker tuning for PredictionsBatch. + # IMPORTANT: prefix must be PREDICTIONSBATCH_ (from class name PredictionsBatch). + - name: PREDICTIONSBATCH_MAX_CONCURRENT_WORKFLOW_TASKS + value: "20" + - name: PREDICTIONSBATCH_MAX_CONCURRENT_ACTIVITIES + value: "60" + - name: PREDICTIONSBATCH_ACTIVITY_EXECUTOR_MAX_WORKERS + value: "10" + - name: PREDICTIONSBATCH_MAX_CONCURRENT_LOCAL_ACTIVITIES + value: "20" + - name: PREDICTIONSBATCH_MAX_CACHED_WORKFLOWS + value: "200" + - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM + value: "3" + - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_INITIAL + value: "5" + - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM + value: "15" + - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM + value: "3" + - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_INITIAL + value: "10" + - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM + value: "30" + + - name: MINIMALRETRAIN_MAX_CONCURRENT_ACTIVITIES + value: "5" + - name: MINIMALRETRAIN_ACTIVITY_EXECUTOR_MAX_WORKERS + value: "5" + - name: MINIMALRETRAIN_MAX_CONCURRENT_LOCAL_ACTIVITIES + value: "5" + - name: MINIMALRETRAIN_MAX_CACHED_WORKFLOWS + value: "5" + - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM + value: "5" + - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_INITIAL + value: "5" + - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM + value: "5" + - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM + value: "5" + - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_INITIAL + value: "5" + - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM + value: "5" + + - name: PI_WEB_API_BASE_URL + value: "https://pivision.votorantimcimentos.com/piwebapi" + - name: PI_WEB_API_AUTH_TYPE + value: "basic" + - name: PI_WEB_API_AUTH_TOKEN + valueFrom: + secretKeyRef: + name: pi-web-api-auth-token + key: token + + # Thread-pool size for non-runtime workers that also use prepare_worker. + - name: SIMPLEMETRICS_ACTIVITY_EXECUTOR_MAX_WORKERS + value: "20" + - name: DRIFT_ACTIVITY_EXECUTOR_MAX_WORKERS + value: "20" + +# ----------------------------------------------------------------------------- +# Runtimes configuration +# ----------------------------------------------------------------------------- +# IMPORTANT: +# - The runtime name is used by the worker bootstrap to resolve plugins and task queues. +# - Keep runtime names in sync with the plugin-store runtime names. +runtimes: + - name: "single" + replicas: 1 + +# ----------------------------------------------------------------------------- +# Chart-level configuration (applies to all runtimes) +# ----------------------------------------------------------------------------- imagePullSecrets: -- name: docker-hub-secret -# This is to override the chart name. -nameOverride: "sientia-laborious-worker" -fullnameOverride: "sientia-laborious-worker" -namespace: sientia + - name: docker-hub-secret + +nameOverride: *projectName +fullnameOverride: *projectName -# This section builds out the service account more information can be found here: https://kubernetes.io/docs/concepts/security/service-accounts/ serviceAccount: - # Specifies whether a service account should be created create: true - # Automatically mount a ServiceAccount's API credentials? automount: true - # Annotations to add to the service account annotations: {} - # The name of the service account to use. - # If not set and create is true, a name is generated using the fullname template - name: "sientia-laborious-worker" + name: *projectName -# This is for setting Kubernetes Annotations to a Pod. -# For more information checkout: https://kubernetes.io/docs/concepts/overview/working-with-objects/annotations/ podAnnotations: {} -# This is for setting Kubernetes Labels to a Pod. -# For more information checkout: https://kubernetes.io/docs/concepts/overview/working-with-objects/labels/ podLabels: {} podSecurityContext: {} - # fsGroup: 2000 - securityContext: {} - # capabilities: - # drop: - # - ALL - # readOnlyRootFilesystem: true - # runAsNonRoot: true - # runAsUser: 1000 - -resources: - # Resource limits and requests are important for ResourceBasedTuner to work correctly. - # The tuner monitors system CPU and memory usage, so proper resource limits must be set. - limits: - cpu: 2000m # 2 CPU cores - memory: 20Gi # 20 GB memory - requests: - cpu: 1000m # 1 CPU core - memory: 2Gi # 2 GB memory - -# This is to setup the liveness and readiness probes more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/configure-liveness-readiness-startup-probes/ -# This is to setup the liveness and readiness probes more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/configure-liveness-readiness-startup-probes/ -livenessProbe: - exec: - command: - - sh - - -c - - | - curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1' - initialDelaySeconds: 1260 - periodSeconds: 15 - timeoutSeconds: 5 - failureThreshold: 3 - -readinessProbe: - exec: - command: - - sh - - -c - - | - curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1' - initialDelaySeconds: 1200 - periodSeconds: 10 - timeoutSeconds: 3 - failureThreshold: 2 - - - -# This section is for setting up autoscaling more information can be found here: https://kubernetes.io/docs/concepts/workloads/autoscaling/ -autoscaling: - enabled: false - minReplicas: 1 - maxReplicas: 100 - targetCPUUtilizationPercentage: 80 - # targetMemoryUtilizationPercentage: 80 - -# Additional volumes on the output Deployment definition. volumes: [] -# - name: foo -# secret: -# secretName: mysecret -# optional: false - -# Additional volumeMounts on the output Deployment definition. volumeMounts: [] -# - name: foo -# mountPath: "/etc/foo" -# readOnly: true nodeSelector: {} - tolerations: [] - affinity: {} services: @@ -124,7 +269,6 @@ services: port: 9091 targetPort: 9091 name: sdk-metrics - metrics: enabled: true type: ClusterIP @@ -135,195 +279,18 @@ services: # Configuração do ServiceMonitor para o Prometheus Operator # ref: https://github.com/prometheus-operator/prometheus-operator serviceMonitor: - # Se true, um recurso ServiceMonitor será criado. enabled: true - # O intervalo no qual as métricas devem ser coletadas (ex: 30s, 1m). endpoints: - - port: metrics - path: /metrics - interval: 30s - relabelings: [] - - port: sdk-metrics - path: /metrics - interval: 30s - relabelings: [] - + - port: metrics + path: /metrics + interval: 30s + relabelings: [] + - port: sdk-metrics + path: /metrics + interval: 30s + relabelings: [] additionalLabels: release: kube-prometheus-stack - - -env: - # Entrypoint variables - - name: GITHUB_REPO_URL - value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git" - - name: GITHUB_BRANCH - value: "feature/SIENTIAPDE-1712" - - name: PYTHON_APP - value: "laborious.worker.worker" - - # Application variables - - name: POSTGRES_HOST - value: "paradedb-rw.paradedb.svc.cluster.local" - - name: POSTGRES_PORT - value: "5432" - - name: POSTGRES_USER - value: "postgres" - - name: POSTGRES_PASSWORD - value: "nFqc81y6kwmr2zuAIx43DhiOosFCVPpeEfTtTWZflkNjB2j1KtEeIANkhFR9mAX3" - - name: POSTGRES_DBNAME - value: "sientia" - - name: POSTGRES_MIN_CONNECTIONS - value: "20" - # max_connections = number_of_workers * max_concurrent_activities * safety_factor - # Example: 4 workers * 50 activities * 0.5 = 100 connections - - name: POSTGRES_MAX_CONNECTIONS - value: "100" - - - name: MLFLOW_HOST - value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local" - - name: MLFLOW_PORT - value: "80" - - name: MLFLOW_USERNAME - value: "aignosi" - - name: MLFLOW_PASSWORD - value: "1L0FP50j3ncp123" - - # Worker runtime (PluginStore): required for PredictionsBatch / MinimalRetrain workers. - - name: RUNTIME - value: "single" - - # Plugin store (model-library-store Git + runtime packages). - - name: STORE_BASE_URL - value: "http://gitea.sientia.svc.cluster.local:3000" - - name: STORE_OWNER - value: "sientia" - - name: STORE_REPO - value: "model-library-store" - - name: STORE_BRANCH - value: "main" - - name: STORE_USERNAME - valueFrom: - secretKeyRef: - name: store-credentials - key: username - - name: STORE_PASSWORD - valueFrom: - secretKeyRef: - name: store-credentials - key: password - - name: STORE_CACHE_TTL_SECONDS - value: "" - - - name: OPC_ID - value: "1" - - name: OPC_SERVER_NAME - value: "default_server" - - name: OPC_URL - value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840" - - - - name: LOG_LEVEL - value: "DEBUG" - - name: HTTP_METRICS_PORT - value: "9090" - - name: HTTP_SDK_METRICS_PORT - value: "9091" - - name: PROJECT_NAME - value: "sientia-laborious" - - - name: TEMPORAL_HOST - value: "temporal-frontend.temporal.svc.cluster.local:7233" - - name: TEMPORAL_NAMESPACE - value: "laborious" - - - name: MONGODB_USERNAME - value: "root" - - name: MONGODB_PASSWORD - value: "wKZDbMNU1c" - - name: MONGODB_URL - value: "my-release-mongodb.mongodb.svc.cluster.local:27017" - - name: MONGODB_DATABASE - value: "sientia" - - name: MONGODB_TTL_INDEX_HOURS - value: "1" - - - name: MINIO_ENDPOINT_URL - value: "minio.minio.svc.cluster.local:9000" - - name: MINIO_ACCESS_KEY - value: "admin" - - name: MINIO_SECRET_KEY - value: "LiArt4eNmJ" - - name: MINIO_DEFAULT_BUCKET - value: "sientia" - - name: MINIO_RETENTION_HOURS - value: "24" - - name: SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES - value: "0.5" - - # Temporal worker tuning for PredictionsBatch. - # IMPORTANT: prefix must be PREDICTIONSBATCH_ (from class name PredictionsBatch). - # Keep workflow-task concurrency moderate to reduce task completion races under load. - - name: PREDICTIONSBATCH_MAX_CONCURRENT_WORKFLOW_TASKS - value: "20" - # Allow higher activity parallelism because most activities are I/O-bound, but keep headroom. - - name: PREDICTIONSBATCH_MAX_CONCURRENT_ACTIVITIES - value: "60" - # Keep local activities controlled so they do not monopolize the event loop. - - name: PREDICTIONSBATCH_MAX_CONCURRENT_LOCAL_ACTIVITIES - value: "20" - # Cache enough workflows for reuse without excessive memory growth. - - name: PREDICTIONSBATCH_MAX_CACHED_WORKFLOWS - value: "200" - # Start with one workflow poller to avoid burst contention at startup. - - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM - value: "3" - # Small initial poller count warms up gradually instead of spiking task fetches. - - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_INITIAL - value: "5" - # Cap workflow pollers to limit scheduling pressure and avoid over-polling. - - name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM - value: "15" - # Keep at least two activity pollers so activity queues do not starve during spikes. - - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM - value: "3" - # Moderate initial activity pollers for faster ramp-up with controlled pressure. - - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_INITIAL - value: "10" - # Limit max activity pollers to preserve CPU for workflow-task completion. - - name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM - value: "30" - - - name: MINIMALRETRAIN_MAX_CONCURRENT_ACTIVITIES - value: "1" - - name: MINIMALRETRAIN_MAX_CONCURRENT_LOCAL_ACTIVITIES - value: "1" - - name: MINIMALRETRAIN_MAX_CACHED_WORKFLOWS - value: "1" - - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM - value: "1" - - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_INITIAL - value: "1" - - name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM - value: "1" - - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM - value: "1" - - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_INITIAL - value: "1" - - name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM - value: "1" - - - name: PI_WEB_API_BASE_URL - value: "https://pivision.votorantimcimentos.com/piwebapi" - - name: PI_WEB_API_AUTH_TYPE - value: "basic" - - name: PI_WEB_API_AUTH_TOKEN - valueFrom: - secretKeyRef: - name: pi-web-api-auth-token - key: token - - - name: PYPI_SERVER - value: "http://library-distribution-server.library.svc.cluster.local:5000" ssh: enabled: true @@ -331,10 +298,10 @@ ssh: sshPath: /mnt/.ssh knownHostsPath: /mnt/known_hosts -# kubectl create secret docker-registry docker-hub-secret --namespace sientia --docker-server=http://aignosi.azurecr.io --docker-username=aignosi --docker-password=5I5zpQ6sRaHqX1hD3dr+2mo647yO3FRc359/wu6gsP+ACRDRz5mp - -# helm upgrade --install sientia-laborious-worker sientia/sientia-module -n sientia --create-namespace -f ./values.yaml --version 0.6.0 - +# kubectl create secret docker-registry docker-hub-secret --namespace sientia --docker-server=http://aignosi.azurecr.io --docker-username=aignosi --docker-password= +# +# helm upgrade --install sientia-laborious-worker sientia/sientia-module -n sientia --create-namespace -f ./values.yaml --version 0.6.1 +# # kubectl create secret generic git-ssh-key-sientia-laborious-worker \ # --namespace sientia \ # --from-file=ssh-privatekey=git_key \