From 1ce8b9d3a7ad741137ccf58aca3c9325f84b25b8 Mon Sep 17 00:00:00 2001 From: vitor-aignosi Date: Tue, 5 May 2026 16:52:51 -0300 Subject: [PATCH] SIENTIAPDE-1773 Enhance environment configuration and update dependencies - Added new environment variables for PluginStore and MLflow configuration in `.env.example`, including `RUNTIME`, `STORE_BASE_URL`, `STORE_OWNER`, `STORE_REPO`, `STORE_BRANCH`, `STORE_USERNAME`, `STORE_PASSWORD`, `STORE_CACHE_TTL_SECONDS`, `PYPI_SERVER`, `PYPI_USERNAME`, and `PYPI_PASSWORD`. - Updated `git-requirements-mapping.txt` to reflect changes in repository names. - Modified `requirements-light.txt` and `requirements.txt` to upgrade `sientia-dataops-library` to version 1.12.0 and `sientia-mlops-library` to version 0.8.1. - Updated `values.yaml` to include new environment variables for worker runtime and PluginStore configuration. - Refactored E2E tests to utilize new MLflow repository stubs and PluginStore mocks for improved testing accuracy. --- .env.example | 16 + README.md | 27 +- e2e/conftest.py | 141 +- e2e/test_minio_offload.py | 3 +- e2e/test_predictions_batch_format_export.py | 3 +- e2e/test_predictions_batch_main_workflow.py | 9 +- ...st_predictions_batch_prediction_process.py | 66 +- git-requirements-mapping.txt | 4 +- ...us-temporal-plugin-store-migration-plan.md | 305 +++ laborious/activities/activities.py | 117 +- laborious/activities/mlflow.py | 554 ++++-- laborious/utils/connectors_config.py | 94 +- .../utils/repository/model_repository.py | 1482 --------------- laborious/worker/prepare_worker.py | 73 - laborious/worker/worker.py | 118 +- requirements-light.txt | 2 +- requirements.txt | 4 +- tests/laborious/activities/test_activities.py | 18 +- tests/laborious/activities/test_mlflow.py | 462 ++--- .../utils/repository/test_model_repository.py | 1680 ----------------- .../laborious/utils/test_connectors_config.py | 37 +- .../workflows/test_minimal_retrain.py | 9 +- values.yaml | 26 + 23 files changed, 1280 insertions(+), 3970 deletions(-) create mode 100644 laborious-temporal-plugin-store-migration-plan.md delete mode 100644 laborious/utils/repository/model_repository.py delete mode 100644 laborious/worker/prepare_worker.py delete mode 100644 tests/laborious/utils/repository/test_model_repository.py diff --git a/.env.example b/.env.example index 6e621e9..bc0f933 100644 --- a/.env.example +++ b/.env.example @@ -11,6 +11,22 @@ MLFLOW_PORT="80" MLFLOW_USERNAME="aignosi" MLFLOW_PASSWORD="mlflow_password" +# Worker runtime name for PluginStore.install_runtime (PredictionsBatch / MinimalRetrain workers). +RUNTIME="single" + +# Plugin store (Git-backed catalog + runtime install). +STORE_BASE_URL="http://gitea.sientia.svc.cluster.local:3000" +STORE_OWNER="sientia" +STORE_REPO="model-library-store" +STORE_BRANCH="main" +STORE_USERNAME="" +STORE_PASSWORD="" +STORE_CACHE_TTL_SECONDS="" + +PYPI_SERVER="http://library-distribution-server.library.svc.cluster.local:5000" +PYPI_USERNAME="" +PYPI_PASSWORD="" + OPC_ID="1" OPC_URL="opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840" diff --git a/README.md b/README.md index b63be57..bf037dc 100644 --- a/README.md +++ b/README.md @@ -163,7 +163,7 @@ Laborious uses a Temporal-based architecture with strong separation of concerns #### **Data Services (`laborious/utils/`)** - `connectors_config.py`: Env-driven configuration builders - `models/minio_dataframe_payload.py`: MinIO-offloaded DataFrame payload model -- `repository/model_repository.py`: MLFlow operations and retraining +- ML models are loaded via `SientiaMLflowRepository` (wrapper-based, `@production` alias) constructed in `Activities` from `build_mlflow_config()`. - `repository/opc_repository.py`: OPC communication and writes - `repository/minio_manager.py`: MinIO object storage operations - `filters/conditional_filters.py` and `filters/mlflow_filters.py` @@ -730,7 +730,6 @@ tests/ │ │ ├── test_conditional_filters.py │ │ └── test_mlflow_filters.py │ └── repository/ -│ ├── test_model_repository.py │ └── test_opc_repository.py ``` @@ -803,7 +802,20 @@ The Laborious system exposes comprehensive Prometheus metrics for operational vi | `MLFLOW_PORT` | MLFlow server port | `5080` | Yes | | `MLFLOW_USERNAME` | MLFlow username | `aignosi` | Yes | | `MLFLOW_PASSWORD` | MLFlow password | `aignosi` | Yes | +| `RUNTIME` | Plugin store runtime name installed at worker boot (required for model-loading workers) | `single` | Yes* | +| `STORE_BASE_URL` | Plugin store Git server base URL | `http://localhost:3000` | Yes* | +| `STORE_OWNER` | Plugin store repository owner | `sientia` | Yes* | +| `STORE_REPO` | Plugin store repository name | `model-library-store` | Yes* | +| `STORE_BRANCH` | Optional branch for the store repository | `main` | No | +| `STORE_USERNAME` | HTTP username for the Git store | `None` | No | +| `STORE_PASSWORD` | HTTP password/token for the Git store | `None` | No | +| `STORE_CACHE_TTL_SECONDS` | Optional cache TTL for store metadata | `None` | No | +| `PYPI_SERVER` | Private PyPI index URL for runtime wheels | `http://localhost:5000` | Yes* | +| `PYPI_USERNAME` | Optional PyPI basic-auth username | `None` | No | +| `PYPI_PASSWORD` | Optional PyPI basic-auth password | `None` | No | | `OPC_CONFIG` | OPC server configuration (JSON) | `{}` | No | + +\* `RUNTIME`, PluginStore (`STORE_*`), and `PYPI_SERVER` are required for workers that install a runtime and load `SientiaModel` wrappers (`PredictionsBatch`, `MinimalRetrain`). Workers that only run drift/simple-metrics style jobs may omit them when those workflows are deployed separately. | `OPC_ID` | OPC server identifier | `1` | No | | `OPC_URL` | OPC server URL | `opc.tcp://localhost:4840` | No | | `OPC_SERVER_URI` | OPC server URI | `opc.tcp://localhost:4840` | No | @@ -1038,11 +1050,8 @@ This is the configuration created by the Orchestrator in Temporal. "EMPTY_DATA":{"config":{},"policy":"STOP"} }, "model_config":{ - "is_compressed":true, - "predict_flavor":"pyfunc", - "retention_minutes":60, - "retention_target":"artifact", - "transform_function_keyword":"transform" + "target":"sensor_or_label_column", + "retention_minutes":60 }, "model_id":"352", "model_name":"courier", @@ -1081,8 +1090,7 @@ laborious/ │ ├── prediction_process.py # Core prediction workflow │ └── format_and_export_prediction.py # Export workflow ├── worker/ # Worker implementation -│ ├── worker.py # Main worker orchestrator -│ └── prepare_worker.py # Worker factory with autoscaling config +│ └── worker.py # Entrypoint; workers built via `sientia_do.temporal.worker.prepare_worker` ├── utils/ # Utility functions │ ├── connectors_config.py # Environment-driven config builders │ ├── models/ # Data models @@ -1091,7 +1099,6 @@ laborious/ │ │ ├── conditional_filters.py # Conditional data filters │ │ └── mlflow_filters.py # MLFlow response filters │ └── repository/ # Data access layer -│ ├── model_repository.py # MLFlow model operations │ ├── opc_repository.py # OPC server operations │ └── minio_manager.py # MinIO object storage operations ├── metrics.py # Prometheus metrics definitions diff --git a/e2e/conftest.py b/e2e/conftest.py index 7825335..0a55f00 100644 --- a/e2e/conftest.py +++ b/e2e/conftest.py @@ -313,88 +313,55 @@ def patch_minio_repository(mock_minio_repository): @pytest_asyncio.fixture def patch_pi_web_api_repository(mock_pi_web_api_repository): - """Patch MLflowRepository to return mock.""" + """Patch PI Web API client to return mock.""" with patch('laborious.activities.api.PIWebAPIClient', return_value=mock_pi_web_api_repository): yield + @pytest_asyncio.fixture -def mock_mlflow_models(): - """Create mock models for MLflow load_model methods.""" - # Mock transform model - returns DataFrame with same index as input - mock_transform_model = MagicMock() - def mock_transform_predict(data): - num_rows = max(len(data), 1) if hasattr(data, '__len__') else 1 - print(data.to_csv()) - print(data.index) - result = pd.DataFrame({ - 'feature_1': [0.234] * num_rows, - 'feature_2': [0.783] * num_rows, - }) +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): + result = pd.DataFrame( + { + 'feature_1': [0.234] * len(data), + 'feature_2': [0.783] * len(data), + } + ) result.index = data.index - return result - mock_transform_model.predict = MagicMock(side_effect=mock_transform_predict) - - # Mock predict model - returns array/list of predictions - mock_predict_model = MagicMock() - def mock_predict_predict(data): - num_rows = max(len(data), 1) if hasattr(data, '__len__') else 1 - return [0.5] * num_rows - mock_predict_model.predict = MagicMock(side_effect=mock_predict_predict) - - # Mock PyFuncModel for compressed models - mock_pyfunc_model = MagicMock() - mock_pyfunc_model._model_impl = MagicMock() - mock_pyfunc_model._model_impl.python_model = mock_transform_model - - return { - 'transform_model': mock_transform_model, - 'predict_model': mock_predict_model, - 'pyfunc_model': mock_pyfunc_model, - } + return result, {} + def _predict_side_effect(_params: dict, data: pd.DataFrame): + pred = pd.DataFrame([0.5] * len(data), columns=['placeholder']) + pred.index = data.index + return pred, {} -@pytest_asyncio.fixture -def patch_mlflow(mock_mlflow_models): - """Patch mlflow module in repository with load_model mocks.""" - mock_mlflow = MagicMock() - - # Mock sklearn.load_model - def mock_sklearn_load_model(model_uri): - if 'data_model' in model_uri or 'transform' in model_uri.lower(): - return mock_mlflow_models['transform_model'] - return mock_mlflow_models['predict_model'] - mock_mlflow.sklearn = MagicMock() - mock_mlflow.sklearn.load_model = MagicMock(side_effect=mock_sklearn_load_model) - - # Mock pyfunc.load_model - def mock_pyfunc_load_model(model_uri): - if 'artifacts' in model_uri or 'tmp' in model_uri: - return mock_mlflow_models['pyfunc_model'] - if 'data_model' in model_uri or 'transform' in model_uri.lower(): - return mock_mlflow_models['transform_model'] - return mock_mlflow_models['predict_model'] - mock_mlflow.pyfunc = MagicMock() - mock_mlflow.pyfunc.load_model = MagicMock(side_effect=mock_pyfunc_load_model) - - # Mock pytorch.load_model - mock_mlflow.pytorch = MagicMock() - mock_mlflow.pytorch.load_model = MagicMock(return_value=mock_mlflow_models['predict_model']) - - # Mock other mlflow methods that might be called - mock_mlflow.set_tracking_uri = MagicMock() - mock_mlflow.get_run = MagicMock(return_value=MagicMock(info=MagicMock(artifact_uri='mlflow-artifacts:/test_run_id'))) - mock_mlflow.tracking = MagicMock() - mock_mlflow.tracking.MlflowClient = MagicMock(return_value=MagicMock( - search_registered_models=MagicMock(return_value=[MagicMock(name='test_model')]), - search_model_versions=MagicMock(return_value=[MagicMock( - current_stage='Production', - version='1', - source='runs:/artifacts/test_run_id' - )]) - )) - - with patch('laborious.utils.repository.model_repository.mlflow', new=mock_mlflow): - yield mock_mlflow + wrapper = MagicMock() + wrapper.transform.side_effect = _transform_side_effect + wrapper.predict.side_effect = _predict_side_effect + + repo.get_cached_model = MagicMock(return_value=wrapper) + repo.stub_wrapper = wrapper + repo._client = MagicMock() + return repo @pytest_asyncio.fixture(scope='function') @@ -407,7 +374,8 @@ async def test_activities( mock_minio_repository, patch_create_engine, patch_minio_repository, - patch_mlflow, + mlflow_repository_stub, + plugin_store_stub, patch_pi_web_api_repository, mock_opc_repository ): @@ -429,12 +397,7 @@ async def test_activities( 'min_connections': 1, 'max_connections': 5, }, - mlflow_config={ - 'host': 'http://localhost', - 'port': '5000', - 'username': 'test', - 'password': 'test', - }, + plugin_store=plugin_store_stub, minio_config={ # Host:port only; Minio() prepends http(s):// from the secure flag. 'endpoint_url': 'localhost:9000', @@ -452,6 +415,8 @@ async def test_activities( }, logger=mock_logger, notification_handler=notification_handler, + metrics_controller=metrics_controller, + mlflow_repository=mlflow_repository_stub, ) activities.opc_repository = { @@ -474,7 +439,8 @@ async def test_activities_real_minio( notification_handler, metrics_controller, patch_create_engine, - patch_mlflow, + mlflow_repository_stub, + plugin_store_stub, patch_pi_web_api_repository, mock_opc_repository, ): @@ -495,12 +461,7 @@ async def test_activities_real_minio( 'min_connections': 1, 'max_connections': 5, }, - mlflow_config={ - 'host': 'http://localhost', - 'port': '5000', - 'username': 'test', - 'password': 'test', - }, + plugin_store=plugin_store_stub, minio_config={ 'endpoint_url': f'localhost:{minio_port}', 'access_key': 'minioadmin', @@ -517,6 +478,8 @@ async def test_activities_real_minio( }, logger=mock_logger, notification_handler=notification_handler, + metrics_controller=metrics_controller, + mlflow_repository=mlflow_repository_stub, ) activities.opc_repository = {'1': mock_opc_repository} try: diff --git a/e2e/test_minio_offload.py b/e2e/test_minio_offload.py index f8c0122..694e208 100644 --- a/e2e/test_minio_offload.py +++ b/e2e/test_minio_offload.py @@ -103,8 +103,7 @@ async def test_predictions_batch_with_minio_offload_path( 'prediction_store_policy': 'lts:1', 'model_config': { 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', + 'target': 'sensor_1', }, 'datetime_columns': ['timestamp', 'created_at'], } diff --git a/e2e/test_predictions_batch_format_export.py b/e2e/test_predictions_batch_format_export.py index 35543f8..91c795c 100644 --- a/e2e/test_predictions_batch_format_export.py +++ b/e2e/test_predictions_batch_format_export.py @@ -40,8 +40,7 @@ base_input_data = { 'prediction_store_policy': 'lts:1', 'model_config': { 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', + 'target': 'sensor_1', }, 'datetime_columns': ['timestamp', 'created_at'], } diff --git a/e2e/test_predictions_batch_main_workflow.py b/e2e/test_predictions_batch_main_workflow.py index 3f8d37f..1408031 100644 --- a/e2e/test_predictions_batch_main_workflow.py +++ b/e2e/test_predictions_batch_main_workflow.py @@ -68,9 +68,8 @@ async def test_scenario_1_1_1_happy_path_complete_success( 'save_transform': True, 'prediction_store_policy': 'lts:1', 'model_config': { + 'target': 'sensor_1', 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', }, 'datetime_columns': ['timestamp', 'created_at'], } @@ -157,9 +156,8 @@ async def test_scenario_1_2_1_sql_query_execution_error( 'save_transform': True, 'prediction_store_policy': 'lts:1', 'model_config': { + 'target': 'sensor_1', 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', }, } @@ -277,9 +275,8 @@ async def test_scenario_1_2_3_invalid_datetime_column_specification( 'save_transform': True, 'prediction_store_policy': 'lts:1', 'model_config': { + 'target': 'sensor_1', 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', }, 'datetime_columns': ['nonexistent_column'], } diff --git a/e2e/test_predictions_batch_prediction_process.py b/e2e/test_predictions_batch_prediction_process.py index af097cd..80249f0 100644 --- a/e2e/test_predictions_batch_prediction_process.py +++ b/e2e/test_predictions_batch_prediction_process.py @@ -50,8 +50,7 @@ base_input_data = { 'prediction_store_policy': 'lts:1', 'model_config': { 'retention_minutes': 0, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'sklearn', + 'target': 'sensor_1', }, 'datetime_columns': ['timestamp', 'created_at'], } @@ -80,24 +79,30 @@ def insert_sample_prediction(postgres_engine, model_id): @pytest.fixture -def bad_data_model(patch_mlflow): - model = MagicMock(predict=MagicMock(side_effect=Exception('Bad data model'))) - patch_mlflow.sklearn.load_model = MagicMock(return_value=model) - return model +def bad_data_model(mlflow_repository_stub): + mlflow_repository_stub.stub_wrapper.transform = MagicMock( + side_effect=Exception('Bad data model') + ) + return mlflow_repository_stub.stub_wrapper @pytest.fixture -def bad_predict_model(patch_mlflow, mock_mlflow_models): - model = MagicMock(predict=MagicMock(side_effect=Exception('Bad predict model'))) +def bad_predict_model(mlflow_repository_stub): + wrapper = mlflow_repository_stub.stub_wrapper - def mock_sklearn_load_model(model_uri): - if 'data_model' in model_uri or 'transform' in model_uri.lower(): - return mock_mlflow_models['transform_model'] - return model + def _good_transform(data): + result = pd.DataFrame( + { + 'feature_1': [0.234] * len(data), + 'feature_2': [0.783] * len(data), + } + ) + result.index = data.index + return result, {} - patch_mlflow.sklearn = MagicMock() - patch_mlflow.sklearn.load_model = MagicMock(side_effect=mock_sklearn_load_model) - return model + wrapper.transform.side_effect = _good_transform + wrapper.predict = MagicMock(side_effect=Exception('Bad predict model')) + return wrapper @pytest.mark.asyncio @@ -107,7 +112,7 @@ async def test_scenario_2_1_1_input_gate_triggers_continue( temporal_worker: Worker, test_activities: Activities, postgres_engine, - mock_mlflow_models, + mlflow_repository_stub, ): """Input gate CONTINUE: export default prediction; MLflow transform/predict not used.""" client = temporal_test_env.client @@ -118,8 +123,8 @@ async def test_scenario_2_1_1_input_gate_triggers_continue( client, PredictionsBatch.run, input_data, make_workflow_id('test-continue-policy') ) assert_continue(postgres_engine, model_id) - mock_mlflow_models['transform_model'].predict.assert_not_called() - mock_mlflow_models['predict_model'].predict.assert_not_called() + mlflow_repository_stub.stub_wrapper.transform.assert_not_called() + mlflow_repository_stub.stub_wrapper.predict.assert_not_called() @pytest.mark.asyncio @@ -129,7 +134,7 @@ async def test_scenario_2_1_2_input_gate_triggers_stop( temporal_worker: Worker, test_activities: Activities, postgres_engine, - mock_mlflow_models, + mlflow_repository_stub, ): """Input gate STOP: no export, no MLflow.""" client = temporal_test_env.client @@ -141,7 +146,7 @@ async def test_scenario_2_1_2_input_gate_triggers_stop( client, PredictionsBatch.run, input_data, make_workflow_id('test-input-stop') ) assert_stop(postgres_engine, model_id) - mock_mlflow_models['transform_model'].predict.assert_not_called() + mlflow_repository_stub.stub_wrapper.transform.assert_not_called() @pytest.mark.asyncio @@ -151,7 +156,7 @@ async def test_scenario_2_1_3_input_gate_triggers_repeat( temporal_worker: Worker, test_activities: Activities, postgres_engine, - mock_mlflow_models, + mlflow_repository_stub, ): """Input gate REPEAT with existing history.""" client = temporal_test_env.client @@ -164,7 +169,7 @@ async def test_scenario_2_1_3_input_gate_triggers_repeat( client, PredictionsBatch.run, input_data, make_workflow_id('test-input-repeat') ) assert_repeat(postgres_engine, model_id, data) - mock_mlflow_models['transform_model'].predict.assert_not_called() + mlflow_repository_stub.stub_wrapper.transform.assert_not_called() @pytest.mark.asyncio @@ -222,7 +227,7 @@ async def test_scenario_2_2_2_transform_gate_triggers_stop( test_activities: Activities, postgres_engine, bad_data_model, - mock_mlflow_models, + mlflow_repository_stub, ): client = temporal_test_env.client model_id = 222 @@ -233,7 +238,7 @@ async def test_scenario_2_2_2_transform_gate_triggers_stop( client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-stop') ) assert_stop(postgres_engine, model_id) - mock_mlflow_models['predict_model'].predict.assert_not_called() + mlflow_repository_stub.stub_wrapper.predict.assert_not_called() @pytest.mark.asyncio @@ -264,19 +269,20 @@ async def test_scenario_2_2_4_transform_content_gate_nan_values_stop( temporal_worker: Worker, test_activities: Activities, postgres_engine, - mock_mlflow_models, + mlflow_repository_stub, ): """mlflow_content_gate triggers STOP when transform output is all NaN (NAN_VALUES filter).""" client = temporal_test_env.client model_id = 224 def all_nan_transform(data): - num_rows = max(len(data), 1) if hasattr(data, '__len__') else 1 - result = pd.DataFrame({'feature_1': [np.nan] * num_rows, 'feature_2': [np.nan] * num_rows}) + result = pd.DataFrame( + {'feature_1': [np.nan] * len(data), 'feature_2': [np.nan] * len(data)} + ) result.index = data.index - return result + return result, {} - mock_mlflow_models['transform_model'].predict = MagicMock(side_effect=all_nan_transform) + mlflow_repository_stub.stub_wrapper.transform = MagicMock(side_effect=all_nan_transform) insert_sample_data(postgres_engine, model_id, [60.0, 78.2]) input_data = get_base_input_data(model_id) @@ -288,7 +294,7 @@ async def test_scenario_2_2_4_transform_content_gate_nan_values_stop( client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-content-stop') ) assert_stop(postgres_engine, model_id) - mock_mlflow_models['predict_model'].predict.assert_not_called() + mlflow_repository_stub.stub_wrapper.predict.assert_not_called() @pytest.mark.asyncio diff --git a/git-requirements-mapping.txt b/git-requirements-mapping.txt index 24095da..02df29a 100644 --- a/git-requirements-mapping.txt +++ b/git-requirements-mapping.txt @@ -1,2 +1,2 @@ -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git:sientia-do -git+ssh://git@github.com/Aignosi/sientia-mlops-library.git:sientia \ No newline at end of file +git+ssh://git@github.com/Aignosi/sientia-dataops-library.git:sientia_do +git+ssh://git@github.com/Aignosi/sientia-model-library.git:sientia_model \ No newline at end of file diff --git a/laborious-temporal-plugin-store-migration-plan.md b/laborious-temporal-plugin-store-migration-plan.md new file mode 100644 index 0000000..c4da423 --- /dev/null +++ b/laborious-temporal-plugin-store-migration-plan.md @@ -0,0 +1,305 @@ +--- +tags: + - engineering + - sientia + - runtime-system + - laborious-temporal + - plugin-store + - migration-plan +created: 2026-03-02 +modified: 2026-03-02 +created_by: Vitor Pimentel +modified_by: Vitor Pimentel +status: draft +--- + +# Sientia Laborious Temporal — PluginStore & Wrapper Migration Plan + +> Migration plan for evolving `sientia-dataops-laborious_temporal` from direct MLflow model loading to a runtime-aware architecture that uses Sientia model wrappers (`SientiaModel`) via their public methods, aligned with the runtime strategy. + +## Summary + +1. [[#Objectives and Scope|Objectives and Scope]] — What this migration must achieve +2. [[#Existing State Overview (laborious_temporal)|Existing State Overview]] — Current responsibilities and coupling points +3. [[#Requirements Mapping|Requirements Mapping]] — Functional and non-functional requirements +4. [[#Target Architecture|Target Architecture]] — Desired runtime and model interaction architecture +5. [[#Implementation Plan|Implementation Plan]] — Phased, detailed changes to apply +6. [[#Testing Strategy|Testing Strategy]] — How to validate the new behavior +7. [[#Rollout and Migration Strategy|Rollout and Migration Strategy]] — How to safely roll out the changes +8. [[#Related Documents|Related Documents]] — Cross-links to supporting documents + +--- + +## Objectives and Scope + +This migration focuses on the `sientia-dataops-laborious_temporal` application and aims to: + +- Keep the **runtime-aware deployment model** consistent with the rest of the runtime system (Helm + `RUNTIME` env var, runtime installation via PluginStore). +- Ensure that **all interactions with models use the public methods of the Sientia wrapper** (`SientiaModel`): + - Use `SientiaModel.train(...)` and `retrain(...)` for training and retraining flows. + - Use `SientiaModel.predict(...)` and `SientiaModel.transform(...)` for inference and preprocessing. +- **Use the shared MLflow repository** (`SientiaMLflowRepository`) for all MLflow operations (load, runs, artifacts, promotion, production lookup, metadata logging); do not implement these in Laborious. + +Out of scope: + +- Replacing MLflow as the tracking and registry backend. +- Redesigning Temporal workflows (queues, retry policies) beyond what is required for the new model interaction style. + +--- + +## Existing State Overview (laborious_temporal) + +Key components in `sientia-dataops-laborious_temporal`: + +- **MLflow activities** (`laborious/activities/mlflow.py`) + - `MLFlow` class exposes Temporal activities for: + - `request_transform` — loads transformation models from MLflow and applies them to input data. + - `request_predict` — loads predictive models from MLflow and generates predictions. + - `retrain_model` — orchestrates retraining using historical data stored in MinIO and MLflow registry. + - `update_production_model` — promotes new versions to production. + - `get_reference_data` — fetches evaluation/reference datasets from model artifacts. + - These activities delegate ML-specific work to `MLFlowRepository`. + +- **MLflow repository** (`laborious/utils/repository/model_repository.py`) + - `MLFlowRepository` encapsulates the interaction with MLflow: + - Model discovery and run resolution (`get_model_run_id`, `get_model_uri`, `get_experiment`, etc.). + - Artifact download and loading for both transformer and prediction models. + - Model caching and retention (`get_model`, `get_cached_operation`). + - Transformation and prediction entry points: + - `transform(...)` wraps `get_cached_operation(..., operation='transform')`. + - `predict(...)` wraps `get_cached_operation(..., operation='predict')`. + - Retraining orchestration (`fit_models`, `create_new_experiment`, `retrain_model`, `update_production_model`). + - Today: + - Models are loaded via MLflow flavors: sklearn, pyfunc, pytorch. + - When `flavor == 'pyfunc'` and `load_wrapper=True`, the repository loads a wrapper via: + - `raw_model = mlflow.pyfunc.load_model(artifact_path)` + - `model = raw_model._model_impl.python_model` + - Production models are resolved using **stages** in the Model Registry (for example, selecting the latest version in stage `Production`); **aliases such as `@production` are not used yet**, and models are registered explicitly as part of the current retrain/promotion flows. + - The wrapper’s `_model_impl` class does not extend `SientiaModel`. + - All this MLflow-specific logic is local to `laborious_temporal` and partially duplicated in `sientia-dataops-model-manager`, which motivates the extraction of a shared MLflow repository in `sientia-dataops-library` (see `mlflow-shared-repository-migration-plan`). + +**MLflow:** All MLflow ops → [[mlflow-shared-repository-migration-plan|shared repository]]. Laborious uses the interface; `SientiaModel` lifecycle is in sientia-model-library. + +### Current vs Target — High-level Flow + +```mermaid +flowchart LR + subgraph current [Current State — Laborious Temporal] + direction TB + TemporalWorker["Temporal Worker"] + MlflowActivities["MLFlow Activities\nrequest_transform / request_predict / retrain_model"] + MLFlowRepositoryNode["MLFlowRepository"] + MLflowRegistry["MLflow Tracking + Registry"] + RawModel["Loaded Model\n(sklearn / pyfunc / pytorch)"] + + TemporalWorker -->|"start workflow\n(Temporal)"| MlflowActivities + MlflowActivities -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode + MLFlowRepositoryNode -->|"search_model_versions()\ncurrent_stage == 'Production'"| MLflowRegistry + MLFlowRepositoryNode -->|"mlflow.*.load_model(model_uri)"| RawModel + MLFlowRepositoryNode -->|"raw_model.predict(data)\nraw_model.fit(data)"| RawModel + end + + subgraph target [Target State — Laborious Temporal] + direction TB + TemporalWorker2["Temporal Worker"] + MlflowActivities2["MLFlow Activities\n(same APIs)"] + MLFlowRepositoryNode2["MLFlowRepository\n(wrapper-aware)"] + MLflowRegistry2["MLflow Tracking + Registry\n(aliases enabled)"] + SientiaWrapperNode["Wrapper Instance\n(extends SientiaModel)"] + + TemporalWorker2 -->|"start workflow\n(Temporal)"| MlflowActivities2 + MlflowActivities2 -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode2 + MLFlowRepositoryNode2 -->|"get_model_version_by_alias('production')\n& models:/name@production"| MLflowRegistry2 + MLFlowRepositoryNode2 -->|"mlflow.pyfunc.load_model(...)"| SientiaWrapperNode + MLFlowRepositoryNode2 -->|"wrapper.transform(...)\nwrapper.predict(...)\nwrapper.train()/retrain()"| SientiaWrapperNode + SientiaWrapperNode -->|"store_model(...)\n(auto-register + update alias)"| MLflowRegistry2 + end +``` + +### Current vs Target — Retrain Hot Path (Code Sketch) + +Current retrain flow inside `MLFlowRepository.fit_models` / `retrain_model` (simplified): + +```python +data_model, _ = await self.download_model( + model_name=model_name, + metadata=metadata, + model_type="transform", + flavor=transform_flavor, + load_wrapper=(transform_flavor == "pyfunc"), +) + +prediction_model, _ = await self.download_model( + model_name=model_name, + metadata=metadata, + model_type="predict", + flavor=predict_flavor, + load_wrapper=(predict_flavor == "pyfunc"), +) + +treated_data_candidate = data_model.fit(data) # or data_model.predict(data) +... +prediction_model.fit(retrain_dataset) # direct fit on underlying model +``` + +Target retrain flow when using wrappers that extend `SientiaModel`: + +```python +wrapper, _ = await self.download_model( + model_name=model_name, + metadata=metadata, + # model_type="predict", # wrapper owns both transformer + model stack + # flavor="pyfunc", flavor will always be pyfunc, this parameter will be removed + # load_wrapper=True, wrapper will always be loaded from _model_impl +) + +# First-time training or full retrain using public API +wrapper.train( + train_data=train_df, # features + target + val_data=val_df, + target=target_name, +) + +# Incremental retrain (when appropriate) +wrapper.retrain(full_retrain_df) + +transformed_df, trans_meta = wrapper.transform(raw_df) +pred_df, pred_meta = wrapper.predict({}, transformed_df, params={}) +``` + +--- + +## Requirements Mapping + +### Functional Requirements + +| ID | Requirement | Description | Impacted Areas | +| ----- | -------------------------------------------------- | ------------------------------------------------------------------------------------------------ | ------------------------------------------------------ | +| FR-01 | Runtime detection and installation | Align worker startup with `RUNTIME` and runtime installation via PluginStore | Worker | +| FR-02 | Wrapper‑based training using public API | Retraining must call the public `retrain(...)` method of `SientiaModel` | `fit_models`, `retrain_model` paths | +| FR-03 | Wrapper‑based inference using public API | Inference must call `predict(...)` and `transform(...)` on the wrapper; obtain wrappers via `SientiaMLflowRepository` | Activities, shared repository | +| FR-04 | Use shared MLflow repository | Delegate all MLflow operations (load, runs, artifacts, promotion, production lookup) to `SientiaMLflowRepository`; do not implement in Laborious | [[mlflow-shared-repository-migration-plan]] | +| FR-05 | Model configuration for training/retraining | `model_config` must define `target` and `retention_minutes`; all models are pyfunc + wrapper (no flavor selection) | `model_config` structures | + +### Non-Functional Requirements + +| ID | Requirement | Description | Impacted Areas | +| ------ | -------------------------------- | ------------------------------------------------------------------------------------------------------------------------ | --------------------------------------- | +| NFR-01 | Consistent public API usage | All wrapper interactions must go through public `SientiaModel` methods; metadata logging is handled by the shared repository | Activities | +| NFR-02 | Fail fast when wrappers unavailable | Where wrappers are not yet available, raise an explicit error requesting model update | Shared repository, config | +| NFR-03 | Testability | Enable unit tests to validate wrapper-based flows and integration with shared repository | `tests/laborious` | +| NFR-04 | Operational safety | Retraining and promotion behavior must remain auditable and robust | Retrain & promotion flows | + +--- + +## Target Architecture + +### Wrapper-centric model interactions + +The target state for `laborious_temporal` is: + +- All training and retraining logic goes through the wrapper’s public methods: `train(...)` and `retrain(...)`. +- Inference and transformation use `transform(df)` and `predict(context, df, params)`. +- All MLflow operations (load, runs, artifacts, promotion, production lookup) go through `SientiaMLflowRepository` (see [[mlflow-shared-repository-migration-plan]]). + +### Use of Shared MLflow Repository + +Laborious obtains wrappers and performs all MLflow operations via `SientiaMLflowRepository`. The shared repository (see [[mlflow-shared-repository-migration-plan]]) owns: alias-based URIs, pyfunc loading, wrapper extraction, promotion, runs, artifacts, and metadata logging. Laborious activities call `repo.load_wrapper(...)` and then the wrapper’s public methods (`transform`, `predict`, `train`, `retrain`); they do not implement MLflow logic. + +--- + +## Implementation Plan + +### Phase 0 — Design and Configuration Alignment + +- **P0-01**: Catalog model types and flavors used by `laborious_temporal`: + - For each active model: + - Flavor will be always pyfunc + - Whether a `_model_impl` wrapper is already present and extends `SientiaModel`. +- **P0-02**: Define configuration fields in `model_config` for wrapper usage: + - Example: + - `target` field for training/retraining (already partially present). + - `retention_minutes` field for model retention in minutes (unchanged). + +### Phase 1 — Inference via SientiaModel Public API (using shared repository) + +- **P1-01**: Replace `download_model` with `SientiaMLflowRepository.load_wrapper(...)`; remove local MLflow loading logic. +- **P1-02**: Update `get_cached_operation` to call `wrapper.transform(data)` and `wrapper.predict({}, data)`; unpack returned metadata; metadata logging is handled by the shared repository. +- **P1-03**: Ensure activities (`request_transform`, `request_predict`) remain unchanged externally (inputs/outputs unchanged). + +### Phase 2 — Retraining via SientiaModel.retrain + +- **P2-01**: Refactor `fit_models` to call `wrapper.retrain(data)` and `wrapper.train(...)`; use `SientiaMLflowRepository` for runs, metrics, artifacts, and promotion (no local MLflow logic). + +### Phase 3 — MLflow Logging and Promotion (via shared repository) + +- **P3-01**: In `create_new_experiment`, use the wrapper’s `store_model(...)` for model artifacts; use `SientiaMLflowRepository` for runs, metrics, and any additional MLflow operations. +- **P3-02**: Use `SientiaMLflowRepository.promote_to_alias(...)` for production promotion; model registration is auto-handled by wrappers. + +### Phase 4 — Runtime Alignment + +- **P4-01**:`laborious_temporal` is also deployed via the runtime-aware Helm chart: + - Read `RUNTIME` env var. + - Install runtime via PluginStore before starting Temporal workers. +- **P4-02**: Standardize worker queue name as `{project_name}-{runtime}-queue`. +- **P4-03**: Fix quality pipelines to a single runtime (to be decided). + +--- + +## Testing Strategy + +- **T1 — Unit tests** + - Add tests in `tests/laborious/utils/repository/test_model_repository.py` to cover: + - Wrapper-based `get_cached_operation` for both `transform` and `predict` (using shared repository). + - Wrapper-based `fit_models` and `retrain_model` paths calling `train` and `retrain` respectively. + - Code coverage must be 100%. +- **T2 — Integration tests** + - Use existing end-to-end tests under `e2e/`: + - Configure a model with `SientiaModel` wrapper (loaded via shared repository). + - Run full prediction and retrain workflows; compare predictions, retrain outcomes, and MLflow artifacts. + +--- + +## Rollout and Migration Strategy + +- **R1 — Create models with the new architecture** + - Update or create models in the modeling pipeline so they: + - Use wrappers that extend `SientiaModel`. + - Correctly implement `train`, `retrain`, `predict`, `transform`, and `store_model`. + - Produce structured metadata in `transform_meta` and `pred_meta`. + - Publish these models to a test store (or dedicated branch/experiment) for initial validation. + +- **R2 — Provision runtimes for the new models** + - Configure and install dedicated runtimes for the new models: + - Ensure runtime dependencies (Python and system libraries) are available via PluginStore/runtime installer. + - Validate that each runtime can: + - Load the wrapper through MLflow. + - Execute `transform` and `predict` end-to-end on sample data. + +- **R3 — Run new models in real workflows** + - Integrate the new wrapper-based models into real Laborious workflows, initially in non-critical environments: + - Route only a subset of flows or entities to the new models. + - Monitor logs (including metadata), business metrics, and retraining/promotion behavior. + - Promote these models to production using aliases (`@production`) in MLflow 3+. + +- **R4 — Migrate remaining models progressively** + - Define migration waves by model family: + - For each existing model: + - Create or adapt a `SientiaModel` wrapper. + - Provision the corresponding runtime. + - Execute the test cycle (T1–T2) from the previous section. + - Update aliases so production traffic uses the new wrapper. + - After all models are migrated: + - Remove legacy stage-based paths (`Production`) and non-wrapper models. + - Simplify the codebase to assume wrappers + aliases only. + +--- + +## Related Documents + +- [[mlflow-shared-repository-migration-plan|MLflow Shared Repository Migration Plan]] — Concepts implemented in the common interface (production lookup, wrapper loading, promotion, metadata logging, etc.) +- [[model-manager-plugin-store-migration-plan|Model Manager PluginStore Migration Plan]] +- [[analytics-implementation-plan|Runtime Analytics Helm Implementation Plan]] +- [[analytics|Runtime Analytics Architecture and Analysis]] +- [[../model-plugin-system/06-end-to-end-flow|Model Plugin System — End-to-End Flow]] + diff --git a/laborious/activities/activities.py b/laborious/activities/activities.py index 9a03f54..2bb3d2a 100644 --- a/laborious/activities/activities.py +++ b/laborious/activities/activities.py @@ -7,6 +7,10 @@ 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_model.model_repository.mlflow_repository import SientiaMLflowRepository + from sientia_model.model_repository.plugin_store import PluginStore + + from laborious.utils.connectors_config import build_mlflow_config from laborious.activities.api import API from laborious.activities.gates import Gates @@ -18,61 +22,73 @@ with workflow.unsafe.imports_passed_through(): class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): """ - Main activities orchestrator for the Laborious system. + Central orchestrator for all Temporal activities used by Laborious workflows. - This class combines functionality from multiple activity classes to provide - a unified interface for all workflow operations. It manages database connections, - MLFlow model interactions, data quality validation, and OPC server communications. + Composes Storage (Postgres + MinIO offload), MLFlow (wrapper-based inference and retrain + via ``SientiaMLflowRepository``), Gates (data quality and ML response filters), OPC exports, + drift/simple metrics, and PI Web API writes. The worker constructs one ``Activities`` instance + per process and registers its callables on multiple workers bound to different task queues. - The class implements multiple inheritance to combine specialized functionality: - - Storage: Database operations and data persistence - - MLFlow: Model inference and transformation operations - - Gates: Data quality validation and filtering mechanisms - - OPC: Real-time data export to OPC servers - - ModelMetrics: Model performance metrics and drift detection - - API: PI Web API export operations for industrial systems + MLflow connectivity: unless ``mlflow_repository`` is injected (tests only), this class builds + ``SientiaMLflowRepository`` from ``build_mlflow_config()`` so tracking credentials and URL + stay aligned with the rest of Laborious env-based configuration. Attributes: - postgres_config (dict): PostgreSQL connection configuration - mlflow_config (dict): MLFlow server configuration - opc_config (dict): OPC server configuration - pi_web_api_config (dict): PI Web API server configuration - logger (Logger): Logging and observability instance - notification_handler (NotificationHandler): Notification management instance + Inherits and exposes behaviour from mixins; the MLFlow mixin holds ``mlflow_repository`` + and ``plugin_store`` after ``__init__``. """ def __init__( self, postgres_config: dict[str, Any], - mlflow_config: dict[str, Any], + plugin_store: PluginStore, minio_config: dict[str, Any], opc_config: dict[str, Any], pi_web_api_config: dict[str, Any], logger: Logger, notification_handler: NotificationHandler, + metrics_controller: MetricsController | None = None, + mlflow_repository: SientiaMLflowRepository | None = None, ): """ - Initialize the Activities orchestrator with all required configurations. + Wire Postgres, MinIO, MLflow, OPC, gates, metrics, and PI Web API into a single object. - This constructor initializes all parent classes with their respective - configurations and sets up the foundation for all activity operations. + A single ``MetricsController`` instance is created (or reused) and passed to MinIO, + MLflow repository, and all mixins so Prometheus and SDK metrics stay consistent. Args: - postgres_config: PostgreSQL connection configuration dictionary - Required keys: host, port, user, password, dbname, min_connections, max_connections - mlflow_config: MLFlow server configuration dictionary - Required keys: host, port, username, password - opc_config: OPC server configuration dictionary - Can contain multiple server configurations - pi_web_api_config: PI Web API server configuration dictionary - Required keys: base_url, auth_type, auth_token - logger: Logger instance for observability and debugging - notification_handler: Notification handler for alerts and monitoring + - postgres_config: Host, port, credentials, db name, and pool bounds for Storage. + - plugin_store: ``PluginStore`` instance; the worker must call ``install_runtime`` before + activities run so wrapper code is importable. + - minio_config: Endpoint, keys, bucket, retention, and TLS flag for object storage payloads. + - opc_config: Map of OPC server id to connection settings for ``OPC`` mixin. + - pi_web_api_config: Base URL and auth for ``API`` mixin. + - logger: Structured logger used across all activities. + - notification_handler: Handler for alerts and persisted notifications. + - metrics_controller: Optional shared controller; if ``None``, a new one is created. + - mlflow_repository: Optional ``SientiaMLflowRepository`` for unit/e2e tests; in production + leave unset so the repository is built from environment via ``build_mlflow_config()``. Raises: - Exception: If any parent class initialization fails + Exception: If any parent ``__init__`` fails (e.g. invalid config keys). + + Return: + None """ - metrics_controller = MetricsController(logger=logger) + + mc = metrics_controller or MetricsController(logger=logger) + + # Production path: one shared MLflow client for all model registry / tracking calls. + if mlflow_repository is None: + mlflow_cfg = build_mlflow_config() + mlflow_repository = SientiaMLflowRepository( + host=mlflow_cfg['url'], + username=mlflow_cfg['username'], + password=mlflow_cfg['password'], + logger=logger, + notification_handler=notification_handler, + metrics_controller=mc, + ) minio_repository = MinioRepository( endpoint=minio_config['endpoint_url'], @@ -81,11 +97,10 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): bucket=minio_config['default_bucket'], logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, secure=minio_config['secure'], ) - # Initialize parent classes Storage.__init__( self, host=postgres_config['host'], @@ -99,19 +114,17 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): minio_repository=minio_repository, logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) MLFlow.__init__( self, - mlflow_host=mlflow_config['host'], - mlflow_port=mlflow_config['port'], - mlflow_username=mlflow_config['username'], - mlflow_password=mlflow_config['password'], + mlflow_repository=mlflow_repository, + plugin_store=plugin_store, minio_repository=minio_repository, logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) Gates.__init__( @@ -119,7 +132,7 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): minio_repository=minio_repository, logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) OPC.__init__( @@ -127,14 +140,14 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): opc_servers=opc_config, logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) ModelMetrics.__init__( self, logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) API.__init__( @@ -144,22 +157,18 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): auth_token=pi_web_api_config['auth_token'], logger=logger, notification_handler=notification_handler, - metrics_controller=metrics_controller, + metrics_controller=mc, ) async def shutdown(self): """ - Gracefully shutdown all activities and clean up resources. + Close database pools, sync clients, and OPC sessions in a defined order. - This method ensures proper cleanup of all resources including: - - PostgreSQL connection pools - - OPC server connections - - PI Web API client connections - - MLFlow model repositories - - Any other resources that need explicit cleanup + Should be invoked on worker exit so connection pools and OPC sessions are released + cleanly before process termination. - The method should be called before the application terminates to ensure - proper resource cleanup and prevent resource leaks. + Return: + None """ Storage.close(self) MLFlow.close(self) diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index f4d42a3..7f210dc 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -1,11 +1,18 @@ from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): + import tempfile import traceback + from datetime import datetime + from pathlib import Path + from shutil import rmtree from typing import Any + import mlflow import numpy as np - from pandas import to_datetime + import pandas as pd + from pandas import DataFrame, to_datetime + from sklearn.model_selection import train_test_split from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger @@ -18,79 +25,71 @@ with workflow.unsafe.imports_passed_through(): now, ) from sientia_do.utils.formatters import create_sample_dict + from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository + from sientia_model.model_repository.plugin_store import PluginStore from 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 - from laborious.utils.repository.model_repository import MLFlowRepository class MLFlow(MinioManager): """ - MLFlow integration activities for model inference operations. + Temporal activities that talk to MLflow through ``SientiaMLflowRepository`` and ``SientiaModel`` wrappers. - This class provides activities for interacting with MLFlow models, including - data transformation and prediction operations. It handles authentication, - data preprocessing, and model management with configurable retention policies. + Models are resolved by registered name and the ``production`` alias (not by legacy stages or + separate transform/predict flavors). ``get_cached_model`` loads or reuses a wrapper; inference + uses ``wrapper.transform`` / ``wrapper.predict``; retrain uses ``wrapper.retrain`` or + ``wrapper.train`` plus ``store_model`` and registry promotion via ``promote_to_alias``. - The class implements comprehensive error handling and logging for all - MLFlow operations, ensuring reliable model inference in production environments. + Large inputs and outputs flow through ``MinioDataFramePayload`` when workflows offload parquet + to MinIO. On failure, transform/predict still return a payload with ``success: False`` and + error details for downstream gates. Attributes: - mlflow_host (str): MLFlow server hostname - mlflow_port (int): MLFlow server port - mlflow_username (str): MLFlow authentication username - mlflow_password (str): MLFlow authentication password - model_monitoring_repository (MLFlowRepository): Repository for MLFlow operations + mlflow_repository: Client for tracking, registry, artifact download, and run lifecycle. + plugin_store: Reference to the store (runtime is installed on the worker; reserved for + future store-backed helpers). """ _MAX_DEBUG_DATAFRAME_ROWS = 100 def __init__( self, - mlflow_host: str, - mlflow_port: int, - mlflow_username: str, - mlflow_password: str, + mlflow_repository: SientiaMLflowRepository, + plugin_store: PluginStore, minio_repository: MinioRepository | None = None, logger: Logger | None = None, notification_handler: NotificationHandler | None = None, metrics_controller: MetricsController | None = None, ): """ - Initialize MLFlow activities with server configuration. + Attach shared MLflow and MinIO clients used by all ML activities in this mixin. Args: - mlflow_host: MLFlow server hostname or IP address - mlflow_port: MLFlow server port number - mlflow_username: Username for MLFlow authentication - mlflow_password: Password for MLFlow authentication - logger: Logger instance for observability and debugging - notification_handler: Notification handler for alerts and monitoring + - mlflow_repository: Repository built by ``Activities`` (or injected in tests). + - plugin_store: Plugin store instance from worker bootstrap. + - minio_repository: MinIO client for ``MinioDataFramePayload`` upload/download. + - logger: Structured logger. + - notification_handler: Notifications on hard failures where applicable. + - metrics_controller: Shared metrics controller. - Raises: - Exception: If MLFlowRepository initialization fails + Return: + None """ + MinioManager.__init__( self, minio_repository, logger, notification_handler, metrics_controller ) - self.mlflow_host = mlflow_host - self.mlflow_port = mlflow_port - self.mlflow_username = mlflow_username - self.mlflow_password = mlflow_password - - self.model_monitoring_repository = MLFlowRepository( - f'{mlflow_host}:{mlflow_port}', - mlflow_username, - mlflow_password, - logger, - notification_handler, - metrics_controller, - ) + self.mlflow_repository = mlflow_repository + self.plugin_store = plugin_store def close(self) -> None: """ - Close the MLFlow activity and clean up resources. + Release MinIO manager resources held by the mixin. + + Return: + None """ MinioManager.close(self) @@ -115,35 +114,100 @@ class MLFlow(MinioManager): metadata, ) + def _detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame: + """ + Ensure the transform output index is homogeneous and encoded as ``DATETIME_FORMAT_WITH_TZ`` strings. + + Accepts an all-string index (validated against the format), or all-``datetime`` / + ``Timestamp`` (naive timestamps are localized to UTC before formatting). Mixed element types + or unsupported types raise ``ValueError`` with a message logged at info level. + + Args: + - data: DataFrame whose index carries the time dimension after transform. + - metadata: Workflow metadata for log correlation. + + Return: + ``pd.DataFrame``: Same frame with a normalized string index; empty frames are returned as-is. + """ + + if data.empty: + self.info('Data is empty, skipping datetime index detection and parsing', metadata) + return data + + index = data.index + index_type = type(index[0]) + + self.info(f'Index type: {index_type}', metadata) + + message = ( + f'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, ' + f'string in format {DATETIME_FORMAT_WITH_TZ}.' + ) + + if not all(isinstance(i, index_type) for i in index): + types = map(str, map(type, index)) + raise ValueError(f'{message}. Elements are {",".join(types)}') + + if index_type is str: + try: + pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ) + except ValueError as e: + raise ValueError(f'{message}. Unable to parse given date format: {e}') from e + + elif index_type is datetime or index_type is pd.Timestamp: + idx = data.index + if hasattr(idx, 'tz') and idx.tz is None: + data.index = idx.tz_localize('UTC') # type: ignore[attr-defined] + + data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ) # type: ignore[attr-defined] + else: + raise ValueError(f'{message}. Got {index_type}.') + + return data + + def _resolve_model_version_for_run(self, run_id: str) -> str: + """ + Map an MLflow ``run_id`` to the latest registered model version that produced that run. + + ``search_model_versions`` may return multiple versions if the model was registered more than + once for the same run; the highest numeric ``version`` wins so promotion targets the newest + artifact set. + + Args: + - run_id: Run UUID from ``retrain_model`` / experiment payload. + + Return: + str: Registry version string acceptable by ``promote_to_alias``. + + Raises: + ValueError: If the filter returns no versions (model not registered for this run). + """ + + versions = self.mlflow_repository._client.search_model_versions( + filter_string=f"run_id='{run_id}'" + ) + if not versions: + raise ValueError(f'No registered model version found for run_id={run_id}') + latest = max(versions, key=lambda v: int(v.version)) + return str(latest.version) + @activity.defn(name='request_transform') async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ - Transform input data using MLFlow models. + Pivot long-format sensor rows, load the production wrapper, and run ``wrapper.transform``. - This activity processes input data through MLFlow model transformation, - including data preprocessing, format conversion, and validation. It handles - data deduplication, pivoting, and cleanup to ensure optimal model performance. - - The transformation process includes: - 1. Data deduplication based on variable and timestamp - 2. Data pivoting for model input format - 3. Null value handling and cleanup - 4. MLFlow model transformation request - 5. Response validation and logging + Expected tabular shape after load: columns including ``variable``, ``timestamp``, ``value``, + and ``created_at`` for deduplication. Data are sorted by ``created_at``, de-duplicated per + ``(variable, timestamp)``, pivoted wide, then passed to the model. ``model_config`` may + include ``retention_minutes`` for wrapper cache TTL. Args: - input_data: Configuration and data for transformation - Required keys: - - metadata (dict): Workflow execution metadata - - data (dict): Input data for transformation - - model_name (str): Name of the MLFlow model to use - - model_retention (int): Model retention period in minutes + - input_data: Dict with ``metadata``, ``model_name``, ``data`` (``MinioDataFramePayload`` + dict or inline dataframe dict), and optional ``model_config``. - Returns: - dict: Transformed data from MLFlow model - - Raises: - Exception: If transformation fails or MLFlow model is unavailable + Return: + ``MinioDataFramePayload`` with transformed frame and ``success: True``, or a payload + with ``success: False`` and exception details in ``status`` if transform fails. """ metadata = input_data['metadata'] self.info('Transforming data...', metadata) @@ -156,12 +220,11 @@ class MLFlow(MinioManager): self._debug_dataframe('Raw input data:', data, metadata) - # Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair + # Long → wide: keep newest row per (variable, timestamp), then pivot for the wrapper API. data = data.sort_values('created_at', ascending=False).drop_duplicates( subset=['variable', 'timestamp'], keep='first' ) - # Pivot data for model input format data = data.pivot(index='timestamp', columns='variable', values='value') data.fillna(np.nan, inplace=True) @@ -172,10 +235,25 @@ class MLFlow(MinioManager): self._debug_dataframe('Processed input data:', data, metadata) - # Request transformation from MLFlow model - response_data = await self.model_monitoring_repository.transform( - model_name, data, model_config, metadata - ) + try: + wrapper = self.mlflow_repository.get_cached_model( + model_name=model_name, + alias='production', + retention_minutes=model_config.get('retention_minutes', 0), + metadata=metadata, + ) + transformed_df, transform_meta = wrapper.transform(data) + if transform_meta: + self.info(f'Wrapper transform metadata: {transform_meta}', metadata) + + transformed_df = self._detect_and_parse_datetime_index(transformed_df, metadata) + + response_data: dict[str, Any] = {'success': True, 'content': transformed_df} + except Exception as e: + response_data = { + 'success': False, + 'content': {'message': str(e), 'traceback': traceback.format_exc()}, + } self.debug( f'Transform raw response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}', @@ -217,32 +295,19 @@ class MLFlow(MinioManager): @activity.defn(name='request_predict') async def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload: """ - Execute predictions using MLFlow models. + Load the production wrapper and call ``wrapper.predict`` on the prepared feature frame. - This activity performs ML model inference using MLFlow models with the - transformed data. It handles data format conversion, null value processing, - and model prediction requests with comprehensive error handling. - - The prediction process includes: - 1. Data format validation and cleanup - 2. Null value handling for model compatibility - 3. MLFlow model prediction request - 4. Response validation and logging - 5. Performance monitoring and metrics + The activity normalizes ``NaN`` to ``None`` for JSON-friendly columns, rebuilds a + ``timestamp`` column in the internal string format, preserves the original index for + alignment, and records ``response_time`` seconds on the output frame. Non-DataFrame + predictions are coerced to a single ``prediction`` column. Args: - input_data: Configuration and data for prediction - Required keys: - - metadata (dict): Workflow execution metadata - - data (dict): Transformed data for prediction - - model_name (str): Name of the MLFlow model to use - - model_retention (int): Model retention period in minutes + - input_data: Same envelope as ``request_transform`` (``metadata``, ``model_name``, + ``data``, optional ``model_config`` with ``retention_minutes``). - Returns: - dict: Prediction results from MLFlow model - - Raises: - Exception: If prediction fails or MLFlow model is unavailable + Return: + ``MinioDataFramePayload`` with predictions or error status mirroring transform behaviour. """ metadata = input_data['metadata'] self.info('Predicting data...', metadata) @@ -255,7 +320,8 @@ class MLFlow(MinioManager): self._debug_dataframe('Input data for prediction:', data, metadata) - # Convert numpy.nan to None for model compatibility + input_index = data.index + data.replace(np.nan, None, inplace=True) data['timestamp'] = data.index @@ -263,10 +329,41 @@ class MLFlow(MinioManager): data['timestamp'], format=DATETIME_FORMAT_WITH_TZ ).dt.strftime(DATETIME_FORMAT) - # Request prediction from MLFlow model - response_data = await self.model_monitoring_repository.predict( - model_name, data, model_config, metadata - ) + try: + wrapper = self.mlflow_repository.get_cached_model( + model_name=model_name, + alias='production', + retention_minutes=model_config.get('retention_minutes', 0), + metadata=metadata, + ) + start_time = datetime.now() + predict_data, pred_meta = wrapper.predict({}, data) + end_time = datetime.now() + + if pred_meta: + self.info(f'Wrapper predict metadata: {pred_meta}', metadata) + + if isinstance(predict_data, DataFrame): + self._debug_dataframe( + 'Data received from model prediction:', predict_data, metadata + ) + predict_data.columns = pd.Index(['prediction']) + else: + self.debug( + f'Data received from model prediction (not a DataFrame): {predict_data}', + metadata, + ) + predict_data = pd.DataFrame(predict_data, columns=['prediction']) + + predict_data.index = input_index + predict_data['response_time'] = (end_time - start_time).total_seconds() + + response_data: dict[str, Any] = {'success': True, 'content': predict_data} + except Exception as e: + response_data = { + 'success': False, + 'content': {'message': str(e), 'traceback': traceback.format_exc()}, + } self.debug( f'Prediction response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}', @@ -303,34 +400,22 @@ class MLFlow(MinioManager): @activity.defn(name='retrain_model') async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]: """ - Retrain MLFlow models with updated training data. + Fit an updated wrapper from historical data, log a new run, and register a model version. - This activity orchestrates the complete model retraining process, - including data preparation, model retraining execution, and result - validation. It handles data preprocessing, column cleanup, and - comprehensive error handling for production model management. - - The retraining process includes: - 1. Data timestamp extraction and validation - 2. Column cleanup and data preparation - 3. Data pivoting for model input format - 4. MLFlow model retraining execution - 5. Result validation and error handling + Flow: load long-format data from MinIO → dedupe/pivot like inference prep → require + ``model_config['target']`` → read current ``production`` version for ``source_run_id`` tag → + ``start_run`` with retrain tags → ``wrapper.retrain`` or ``wrapper.train`` when + ``full_retrain`` is set (optional ``validation_fraction``) → log input CSV artifact → + ``store_model`` and ``log_params``. Does not promote; the workflow calls + ``update_production_model`` after validation. Args: - input_data (dict): Input data containing: - - metadata (dict): Workflow execution metadata - - data (dict[str, Any]): Training data for model retraining - - model_name (str): Name of the MLFlow model to retrain + - input_data: Must include ``metadata``, ``model_name``, ``data`` (payload), and + ``model_config`` with at least ``target``; optional ``full_retrain``, ``validation_fraction``. - Returns: - dict: Retraining results containing: - - status (str): Retraining operation status - - timestamp (str): Timestamp of the retraining operation - - experiment (str): MLFlow experiment identifier - - Raises: - Exception: If retraining fails or encounters critical errors + Return: + On success: ``success``, ``experiment`` (``run_id``, ``experiment_id``, ``experiment_name``), + ``message``, ``timestamp``. On failure: ``success: False``, error fields, and optional trace. """ if self.minio_repository is None: @@ -339,7 +424,6 @@ class MLFlow(MinioManager): metadata = input_data['metadata'] try: - # Payload-based retrain input (inline dict or MinIO offloaded). payload = MinioDataFramePayload.from_dict(input_data['data']) data = await payload.retrieve(self.minio_repository, metadata) @@ -371,7 +455,6 @@ class MLFlow(MinioManager): timestamp = data['timestamp'].max() self.debug(f'Timestamp: {timestamp}', metadata) - # Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair if 'created_at' in data.columns: data = data.sort_values('created_at', ascending=False).drop_duplicates( subset=['variable', 'timestamp'], keep='first' @@ -382,10 +465,8 @@ class MLFlow(MinioManager): data.drop(columns=['model_id'], inplace=True, errors='ignore') data.drop(columns=['created_at'], inplace=True, errors='ignore') - # Pivot data for model input format data = data.pivot(index='timestamp', columns='variable', values='value') data.fillna(np.nan, inplace=True) - # data.reset_index(inplace=True) data.columns.name = None data['timestamp'] = data.index @@ -396,60 +477,110 @@ class MLFlow(MinioManager): data.columns.name = None - retrain_output = await self.model_monitoring_repository.retrain_model( - data=data, model_name=model_name, model_config=model_config, metadata=metadata - ) + target = model_config.get('target') + if target is None: + msg = 'model_config must include "target" for retraining' + self.info(msg, metadata) + return { + 'success': False, + 'experiment': None, + 'message': msg, + 'traceback': '', + 'timestamp': str(timestamp), + } - if not retrain_output['success']: - trace = retrain_output['traceback'] - await self.send_notification_async( - metadata=metadata, - notification_id='RETRAIN_MODEL_ERROR', - message=f'Error retraining model {model_name}: {retrain_output["message"]}', - block='retrain_model', - level=NotificationLevel.ERROR, - attachment_content=trace, + try: + mv_src = self.mlflow_repository._client.get_model_version_by_alias( + name=model_name, + alias='production', ) - self.error(trace, metadata=metadata) + source_run_id = mv_src.run_id - return {**retrain_output, 'timestamp': timestamp} + wrapper = self.mlflow_repository.get_cached_model( + model_name=model_name, + alias='production', + retention_minutes=0, + metadata=metadata, + ) + + run_name = f'{model_name}-retrain-{datetime.now().strftime("%Y%m%d%H%M%S")}' + + with self.mlflow_repository.start_run( + model_name=model_name, + run_name=run_name, + experiment_name=model_name, + tags={'retrain': 'true', 'source_run_id': source_run_id}, + metadata=metadata, + ) as run_info: + if model_config.get('full_retrain'): + val_frac = float(model_config.get('validation_fraction', 0.2)) + train_df, val_df = train_test_split(data, test_size=val_frac, random_state=42) + wrapper.train( + train_data=train_df, + val_data=val_df, + target=target, + ) + else: + wrapper.retrain(data) + + tmp_dir = tempfile.mkdtemp(prefix='laborious_retrain_') + try: + raw_csv = Path(tmp_dir) / 'retrain_input.csv' + data.to_csv(raw_csv, index=False) + mlflow.log_artifact(str(raw_csv)) + finally: + rmtree(tmp_dir, ignore_errors=True) + + wrapper.store_model(name=model_name) + + self.mlflow_repository.log_params( + { + 'retrain': 'true', + 'retrain_date': datetime.now().isoformat(), + 'source_run_id': source_run_id, + 'retrain_samples': str(data.shape), + } + ) + + experiment_payload = { + 'run_id': run_info.run_id, + 'experiment_id': run_info.experiment_id, + 'experiment_name': model_name, + } + + return { + 'success': True, + 'experiment': experiment_payload, + 'message': 'Model retrained successfully.', + 'timestamp': str(timestamp), + } + + except Exception as e: + error_msg = f'Error retraining model {model_name}: {e}' + self.info(error_msg, metadata) + return { + 'success': False, + 'experiment': None, + 'message': error_msg, + 'traceback': traceback.format_exc(), + 'timestamp': str(timestamp), + } @activity.defn(name='update_production_model') async def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ - Update production model with newly trained model version. + Point the ``production`` alias at the model version registered for the retrain run. - This activity manages the critical process of updating production - models with newly trained versions. It handles model deployment, - status tracking, and comprehensive reporting for operational - visibility and audit trails. - - The update process includes: - 1. Production model update execution - 2. Status and metadata tracking - 3. Comprehensive reporting and logging - 4. Error handling and notification - 5. Audit trail maintenance + Resolves the highest numeric registry version whose ``run_id`` matches + ``experiment['run_id']``, then calls ``promote_to_alias``. On failure, sends a notification + and re-raises so the workflow can surface the error. Args: - input_data (dict): Input data containing: - - metadata (dict): Workflow execution metadata - - model_name (str): Name of the MLFlow model to update - - experiment (str): MLFlow experiment identifier - - model_id (str): Unique identifier for the model version - - timestamp (str): Timestamp of the update operation - - status (str): Current status of the model update + - input_data: ``metadata``, ``model_name``, and ``experiment`` with ``run_id`` and + ``experiment_id`` (as returned from ``retrain_model``). - Returns: - dict[Any, Any]: Comprehensive update report containing: - - model_id (str): Model version identifier - - model_name (str): Name of the updated model - - timestamp (str): Update operation timestamp - - status (str): Update operation status - - Additional MLFlow response metadata - - Raises: - Exception: If production model update fails + Return: + Dict with ``model_name``, promoted ``version``, ``mlflow_run_id``, ``mlflow_experiment_id``. """ metadata = input_data['metadata'] model_name = input_data['model_name'] @@ -459,12 +590,25 @@ class MLFlow(MinioManager): ) try: - response = await self.model_monitoring_repository.update_production_model( - experiment=experiment, model_name=model_name, metadata=metadata + run_id = experiment['run_id'] + experiment_id = experiment['experiment_id'] + + version = self._resolve_model_version_for_run(run_id) + + self.mlflow_repository.promote_to_alias( + model_name=model_name, + version=version, + alias='production', + metadata=metadata, ) self.info(f'Production model {model_name} updated successfully', metadata) - return response + return { + 'model_name': model_name, + 'version': version, + 'mlflow_run_id': run_id, + 'mlflow_experiment_id': experiment_id, + } except Exception as e: trace = traceback.format_exc() @@ -482,47 +626,51 @@ class MLFlow(MinioManager): @activity.defn(name='get_reference_data') async def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None: """ - Get reference data from the MLflow Model Registry. + Download ``evaluation_data.csv`` from the MLflow run linked to ``production`` and parse it. - This method retrieves evaluation reference data stored as artifacts in the - MLflow Model Registry. The reference data is typically used for model - drift detection, performance comparison, and quality validation. The method - loads the data from a CSV artifact file and formats timestamps for - consistent processing. - - The method handles: - 1. Loading evaluation data artifact from MLflow Model Registry - 2. Timestamp parsing and formatting for consistency - 3. Data conversion to dictionary format for workflow consumption - 4. Graceful handling of missing reference data + Used by drift workflows to compare live data against the reference distribution logged with + the model. Artifacts are downloaded to a temp directory, discovered via ``rglob`` (nested + layout-safe), then timestamps are normalized to ``DATETIME_FORMAT`` string columns before + returning record-oriented dicts. Args: - input_data (dict): Input data containing: - - metadata (dict): Workflow execution metadata - - model_name (str): Name of the MLFlow model to get reference data from + - input_data: ``metadata`` and ``model_name`` for registry lookup. - Returns: - list[dict[Hashable, Any]] | None: Reference data from the MLflow Model Registry - as a list of dictionaries. Returns None if reference data is not found - or if the artifact does not exist. - - Raises: - Exception: If artifact loading fails or encounters errors during processing + Return: + List of row dicts, or ``None`` if the artifact path is missing or any step fails. """ - metadata = input_data['metadata'] model_name = input_data['model_name'] - artifact = 'evaluation_data.csv' - reference_data = await self.model_monitoring_repository.load_artifact_dataframe( - model_name=model_name, artifact_path=artifact, metadata=metadata - ) + try: + mv = self.mlflow_repository._client.get_model_version_by_alias( + name=model_name, + alias='production', + ) + run_id = mv.run_id - if reference_data is None: - self.warning(f'Reference data not found for model {model_name}', metadata) + tmpdir = tempfile.mkdtemp(prefix='laborious_eval_') + try: + self.mlflow_repository.download_artifacts( + run_id=run_id, + artifact_path='evaluation_data.csv', + dst_path=tmpdir, + metadata=metadata, + ) + csv_candidates = list(Path(tmpdir).rglob('evaluation_data.csv')) + if not csv_candidates: + self.warning(f'Reference data not found for model {model_name}', metadata) + return None + + reference_data = pd.read_csv(csv_candidates[0]) + finally: + rmtree(tmpdir, ignore_errors=True) + + reference_data['timestamp'] = to_datetime(reference_data['timestamp']) + reference_data['timestamp'] = reference_data['timestamp'].dt.strftime(DATETIME_FORMAT) + + return reference_data.to_dict(orient='records') + + except Exception as e: + self.warning(f'Reference data not found for model {model_name}: {e}', metadata) return None - - reference_data['timestamp'] = to_datetime(reference_data['timestamp']) - reference_data['timestamp'] = reference_data['timestamp'].dt.strftime(DATETIME_FORMAT) - - return reference_data.to_dict(orient='records') diff --git a/laborious/utils/connectors_config.py b/laborious/utils/connectors_config.py index ed3477a..467c8e2 100644 --- a/laborious/utils/connectors_config.py +++ b/laborious/utils/connectors_config.py @@ -3,31 +3,88 @@ from os import getenv from typing import Any +def _build_mlflow_tracking_url() -> str: + """ + Compose a single tracking URI for ``SientiaMLflowRepository`` from host and port env vars. + + If ``MLFLOW_HOST`` already contains a port in the authority (e.g. ``http://tracker:80``), + it is returned unchanged so operators can override port logic explicitly. + + Return: + str: Full tracking URL (scheme + host [+ port]). + """ + + mlflow_host = getenv('MLFLOW_HOST', 'http://localhost').rstrip('/') + mlflow_port = getenv('MLFLOW_PORT', '5080') + + # Host already includes an explicit port (e.g. http://tracker:80) + host_after_scheme = mlflow_host.split('://', 1)[-1] + if ':' in host_after_scheme: + return mlflow_host + + return f'{mlflow_host}:{mlflow_port}' + + def build_mlflow_config() -> dict[str, Any]: """ - Build MLFlow server configuration from environment variables. + Read MLflow tracking and registry credentials from the environment. - This function constructs an MLFlow configuration dictionary from - environment variables with sensible defaults for local development. - It handles server connection and authentication parameters. + Used by ``Activities`` when constructing ``SientiaMLflowRepository``. The ``url`` value is the + same string workers and notebooks should use for ``MLFLOW_TRACKING_URI``-style clients. Environment Variables: - MLFLOW_HOST: MLFlow server hostname (default: http://localhost) - MLFLOW_PORT: MLFlow server port (default: 5080) - MLFLOW_USERNAME: MLFlow username (default: aignosi) - MLFLOW_PASSWORD: MLFlow password (default: aignosi) + MLFLOW_HOST: Host with scheme; port optional if MLFLOW_PORT is set (default: http://localhost) + MLFLOW_PORT: Appended when host has no explicit port (default: 5080) + MLFLOW_USERNAME: Basic-auth or service user (default: aignosi) + MLFLOW_PASSWORD: Password or token (default: aignosi) - Returns: - dict: MLFlow configuration dictionary with all required parameters + Return: + dict[str, Any]: ``url``, ``username``, ``password``. """ return { - 'host': getenv('MLFLOW_HOST', 'http://localhost'), - 'port': int(getenv('MLFLOW_PORT', '5080')), + 'url': _build_mlflow_tracking_url(), 'username': getenv('MLFLOW_USERNAME', 'aignosi'), 'password': getenv('MLFLOW_PASSWORD', 'aignosi'), } +def build_plugin_store_config() -> dict[str, Any]: + """ + Collect settings for ``PluginStore`` (Git-backed catalog + runtime install via pip). + + Mirrors the model-manager service: the worker passes these kwargs into ``PluginStore`` after + ``install_runtime`` resolves wheels from the configured PyPI index. Missing optional env vars + become ``None`` so the store can run without auth in local dev. + + Environment Variables: + STORE_BASE_URL: Git HTTP(S) server (e.g. Gitea) base URL (default: http://localhost:3000) + STORE_OWNER: Namespace or org owning the store repo (default: sientia) + STORE_REPO: Repository name (default: model-library-store) + STORE_BRANCH: Checkout branch; unset lets the client use default + STORE_USERNAME / STORE_PASSWORD: HTTP basic credentials for Git fetch + STORE_CACHE_TTL_SECONDS: Optional integer seconds for metadata cache TTL + PYPI_SERVER: Index URL for ``pip install`` during runtime install (default: http://localhost:5000) + PYPI_USERNAME / PYPI_PASSWORD: Optional index authentication + + Return: + dict[str, Any]: Keys aligned with ``PluginStore`` constructor parameter names. + """ + + cache_ttl_seconds = getenv('STORE_CACHE_TTL_SECONDS') + return { + 'base_url': getenv('STORE_BASE_URL', 'http://localhost:3000'), + 'owner': getenv('STORE_OWNER', 'sientia'), + 'repo': getenv('STORE_REPO', 'model-library-store'), + 'username': getenv('STORE_USERNAME'), + 'password': getenv('STORE_PASSWORD'), + 'branch': getenv('STORE_BRANCH'), + 'cache_ttl_seconds': int(cache_ttl_seconds) if cache_ttl_seconds else None, + 'pypi_index_url': getenv('PYPI_SERVER', 'http://localhost:5000'), + 'pypi_username': getenv('PYPI_USERNAME'), + 'pypi_password': getenv('PYPI_PASSWORD'), + } + + def build_opc_config() -> dict[str, Any]: """ Build OPC server configuration from environment variables. @@ -73,14 +130,15 @@ def build_minio_config() -> dict[str, Any]: Build MinIO (S3-compatible) configuration from environment variables. Environment Variables: - MINIO_ENDPOINT: MinIO endpoint including scheme (default: http://localhost:9000) + MINIO_ENDPOINT_URL: Host:port or URL for the S3 API (default: http://localhost:9000) MINIO_ACCESS_KEY: Access key (default: minioadmin) MINIO_SECRET_KEY: Secret key (default: minioadmin) - MINIO_REGION: Region name for S3 client (default: us-east-1) - MINIO_BUCKET_DEFAULT: Default bucket for uploads (default: laborious) - MINIO_SECURE: Whether to use HTTPS (default: false) - Returns: - dict: MinIO configuration dictionary + MINIO_DEFAULT_BUCKET: Default bucket for Laborious payloads (default: laborious) + MINIO_RETENTION_HOURS: Offloaded object retention window (default: 24) + MINIO_SECURE: If ``true``, use HTTPS (default: false) + + Return: + dict[str, Any]: Keys consumed by ``Activities`` / ``MinioRepository``. """ return { 'endpoint_url': getenv('MINIO_ENDPOINT_URL', 'http://localhost:9000'), diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py deleted file mode 100644 index b4fdad5..0000000 --- a/laborious/utils/repository/model_repository.py +++ /dev/null @@ -1,1482 +0,0 @@ -""" -MLflow repository utilities - -This module provides the `MLFlowRepository` class and helpers to interact with -an MLflow tracking server and model registry. It covers model discovery, -downloading/loading with multiple flavors, cached operations with retention -policies, transformation/prediction interfaces, retraining workflows, and -production model promotion. - -Capabilities: -- Model loading and caching with retention policies -- Data transformation and prediction operations -- Model retraining workflows -- Production model updates and versioning -""" - -import ctypes -import gc -import threading -import time -import traceback -from datetime import datetime, timedelta -from io import StringIO -from os import environ, makedirs, path -from shutil import rmtree -from typing import Any, Literal, overload - -import mlflow -import pandas as pd -from mlflow.entities import Experiment -from numpy import ndarray -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.temporal.constants import DATETIME_FORMAT_WITH_TZ - -from laborious import metrics -from laborious.utils.dataframe_debug import build_dataframe_debug_message - -ARTIFACTS_PATH = './tmp/artifacts' -TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl' -PREDICTION_COMPRESSED_PATH = 'artifacts/stacking_model.pkl' - -INVALID_FLAVOR_MESSAGE = "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'." - - -def force_memory_release(logger: Logger): - """Attempt to release memory from the Python process. - - Executes a garbage collection cycle and calls `malloc_trim(0)` on glibc - where available to return free memory to the OS. This may be a no-op on - non-glibc systems. - - Args: - logger (Logger): Logger for observability. - """ - gc.collect() - - try: - ctypes.CDLL('libc.so.6').malloc_trim(0) - logger.info('Memory released') - except Exception as e: - logger.info(f'Memory release failed: {e}') - - -class MLFlowRepository(SientiaMonitoring): - _MAX_DEBUG_DATAFRAME_ROWS = 100 - - def __init__( - self, - host: str, - username: str, - password: str, - logger: Logger, - notification_handler: NotificationHandler, - metrics_controller: MetricsController, - ): - """Initialize MLflow client and base state. - - Args: - host (str): MLflow tracking URI. - username (str): MLflow username. - password (str): MLflow password. - logger (Logger): Logger instance. - """ - # set tracking uri - SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) - mlflow.set_tracking_uri(host) - - environ['MLFLOW_TRACKING_USERNAME'] = username - environ['MLFLOW_TRACKING_PASSWORD'] = password - # Create an MLflow client - self.client = mlflow.tracking.MlflowClient() - - self.model_cache: dict[str, Any] = {} - self._cache_lock = threading.RLock() - self.logger = logger - - def _debug_dataframe(self, message: str, data: Any, metadata: dict[str, Any]) -> None: - """ - Log dataframe content only when row count is below the configured threshold - - Args: - - message (str): Base log message to identify the dataframe in logs - - data (Any): Dataframe-like object expected to expose shape and to_csv - - metadata (dict[str, Any]): Metadata for contextual logging - """ - self.debug( - build_dataframe_debug_message( - message=message, - data=data, - max_rows=self._MAX_DEBUG_DATAFRAME_ROWS, - ), - metadata, - ) - - """ - Functions related to get model registry parameters - """ - - def get_model_uri(self, run_id: str, prediction: bool = True): - """Build the artifact URI for a run's model. - - Args: - run_id (str): MLflow run identifier. - prediction (bool): If True, return `prediction_model` URI, - otherwise return `data_model` URI. - - Returns: - str: Artifact URI to the selected model within the run. - """ - run_info = mlflow.get_run(run_id) - if prediction: - model_uri = run_info.info.artifact_uri + '/prediction_model' - else: - model_uri = run_info.info.artifact_uri + '/data_model' - return model_uri - - def get_model_run_id(self, model_name: str, stage: str = 'Production') -> str: - """Resolve the run_id for a registered model at a given stage. - - Args: - model_name (str): Registered model name. - stage (str): Desired stage (e.g., 'Production'). - - Returns: - str: Run ID for the latest version at the given stage. - """ - # Use search_registered_models instead of deprecated get_latest_versions - registered_models = self.client.search_registered_models( - filter_string=f"name='{model_name}'" - ) - - if not registered_models: - raise mlflow.exceptions.MlflowException( - f"Model '{model_name}' not found in the Model Registry." - ) - - # Get all versions of the model and filter by stage - model_versions = self.client.search_model_versions(filter_string=f"name='{model_name}'") - - # Filter versions by the desired stage using current_stage attribute - stage_versions = [mv for mv in model_versions if mv.current_stage == stage] - - if not stage_versions: - raise mlflow.exceptions.MlflowException( - f"Model '{model_name}' in stage '{stage}' not found in the Model Registry." - ) - - # Sort by version number to get the latest - latest_version = max(stage_versions, key=lambda v: int(v.version)) - run_id = latest_version.source.split('/') - return run_id[2] - - def get_next_run_name(self, model_name: str) -> str: - """ - Generate the next run name for a specific MLflow model. - - This method calculates the next sequential run number for a model - by searching existing runs and incrementing the count. It ensures - unique run names for model training and retraining operations. - - Args: - model_name (str): The name of the MLflow model - - Returns: - str: The next run name in format 'model_name-run_number' - """ - runs = mlflow.search_runs(experiment_names=[model_name], order_by=['start_time desc']) - next_run_number = len(runs) + 1 - return f'{model_name}-{next_run_number}' - - def get_experiment( - self, experiment_name: str, create_if_not_exists: bool = False - ) -> Experiment: - """ - Retrieve MLflow experiment by name, optionally creating it. - - This method searches for an MLFlow experiment by name and - returns its unique identifier. It provides error handling - for non-existent experiments. - - Args: - experiment_name (str): Name of the MLflow experiment - - Returns: - Experiment: MLflow experiment object - - Raises: - ValueError: If the experiment name is not found and creation is disabled - """ - experiment = mlflow.get_experiment_by_name(experiment_name) - - if experiment is None: - if create_if_not_exists: - experiment_id = mlflow.create_experiment(experiment_name) - experiment = mlflow.get_experiment(experiment_id) - if experiment is None: - raise ValueError( - f'Experiment {experiment_name} not found after creation, unknown reason' - ) - else: - raise ValueError(f'Experiment {experiment_name} not found') - - return experiment - - def get_model_params(self, run_id: str): - """Fetch parameters associated with a given MLflow run. - - Args: - run_id (str): Run identifier to inspect. - - Returns: - dict: Mapping of parameter names to values. - """ - run_info = mlflow.get_run(run_id) - return run_info.data.params - - def check_artifact_exists( - self, run_id: str, artifact_path: str, metadata: dict[str, Any] - ) -> bool: - """ - Check if an artifact exists in the MLflow Model Registry. - - Args: - run_id (str): Run identifier to inspect. - artifact_path (str): Path to the artifact to check. - - Returns: - bool: True if the artifact exists, False otherwise. - """ - artifacts = self.client.list_artifacts(run_id) - - self.debug(f'Artifacts of {run_id}: \n{artifacts}', metadata) - self.debug(f'Looking for artifact {artifact_path} in {run_id}', metadata) - - return any(artifact.path == artifact_path for artifact in artifacts) - - """ - Functions related to download and load models - """ - - async def dowload_artifacts( - self, model_name: str, metadata: dict[str, Any], artifact_path: str = 'data_model' - ) -> str: - """ - Download artifacts from the latest production run of a model. - - Args: - model_name (str): Registered model name. - metadata (dict[str, Any]): Metadata used for structured logging. - artifact_path (str): Relative path to artifacts within the run. - - Returns: - str: Local filesystem path where artifacts are saved. - """ - run_id = self.get_model_run_id(model_name=model_name, stage='Production') - output_dir = f'{ARTIFACTS_PATH}/{model_name}' - - full_path = path.join(output_dir, artifact_path) - - if path.exists(full_path): - # Remove the directory and create a new one - rmtree(full_path) - makedirs(output_dir, exist_ok=True) - - self.info(f'Downloading artifacts from {run_id} to {output_dir}') - - core_labels = self.get_core_labels(metadata, operation_type='download_artifacts') - - start_time = time.time() - try: - artifacts = self.client.download_artifacts(run_id, artifact_path, output_dir) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels) - - return artifacts - - async def load_artifact_dataframe( - self, model_name: str, artifact_path: str, metadata: dict[str, Any] - ) -> pd.DataFrame | None: - """ - Load the dataframe content of an artifact from the MLflow Model Registry. - - Args: - model_name (str): The name of the model to download from the registry. - artifact_path (str): The path to the artifact to load. - metadata (dict[str, Any]): Metadata used for structured logging. - - Returns: - pd.DataFrame: The dataframe content of the artifact. - """ - run_id = self.get_model_run_id(model_name=model_name, stage='Production') - core_labels = self.get_core_labels(metadata, operation_type='load_text') - - if not self.check_artifact_exists(run_id, artifact_path, metadata): - return None - - artifact_path = path.join('runs:/', run_id, artifact_path) - - start_time = time.time() - try: - content = mlflow.artifacts.load_text(artifact_path) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels) - - self.debug(f'Content of {run_id}/{artifact_path}: \n{content}', metadata) - - dataframe = pd.read_csv(StringIO(content)) - - self.info(f'Loaded dataframe from {model_name}:{artifact_path}', metadata) - return dataframe - - async def load_predict_model( - self, model_name: str, metadata: dict[str, Any], flavor: str = 'sklearn' - ) -> Any: - """ - Load a predictive model from the MLflow Model Registry. - - Args: - model_name (str): The name of the model to download from the registry. - metadata (dict[str, Any]): Metadata used for structured logging. - flavor (str): Model flavor ('pyfunc', 'sklearn', 'pytorch') - artifact_path (str | None): Path to compressed artifacts if model is compressed - - Returns: - mlflow.pyfunc.PyFuncModel: The loaded predictive model. - - Notes: - - The model is fetched from the "Production" stage of the MLflow Model Registry. - - Warnings during the model loading process are suppressed. - """ - model_uri = f'models:/{model_name}/production' - self.info(f'Loading prediction model {model_name} from {model_uri}') - - core_labels = self.get_core_labels(metadata, operation_type='load_predict_model') - start_time = time.time() - try: - if flavor == 'pyfunc': - model = mlflow.pyfunc.load_model(model_uri) - elif flavor == 'sklearn': - model = mlflow.sklearn.load_model(model_uri) - elif flavor == 'pytorch': - model = mlflow.pytorch.load_model(model_uri) - else: - raise ValueError(INVALID_FLAVOR_MESSAGE) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels) - return model - - async def load_transform_model( - self, model_name: str, metadata: dict[str, Any], flavor: str = 'sklearn' - ) -> Any: - """ - Load the latest Production version of a transformation model. - - This method retrieves the latest production model run ID for the given - model name, constructs the model URI, and loads the model using MLflow. - - Args: - model_name (str): The name of the model to download. - metadata (dict[str, Any]): Metadata used for structured logging. - flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch') - artifact_path (str | None): Path to compressed artifacts if model is compressed - - Returns: - Any: The loaded model object, depending on the flavor used. - - Raises: - Exception: If the model run ID or URI cannot be retrieved, or if the - model cannot be loaded. - """ - - latest_production_id = self.get_model_run_id(model_name=model_name, stage='Production') - model_uri = self.get_model_uri(latest_production_id, prediction=False) - - self.info(f'Loading data model {model_name} from {model_uri}') - - core_labels = self.get_core_labels(metadata, operation_type='load_transform_model') - start_time = time.time() - try: - if flavor == 'sklearn': - model = mlflow.sklearn.load_model(model_uri) - elif flavor == 'pyfunc': - model = mlflow.pyfunc.load_model(model_uri) - elif flavor == 'pytorch': - model = mlflow.pytorch.load_model(model_uri) - else: - raise ValueError(INVALID_FLAVOR_MESSAGE) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels) - - return model - - async def download_model( - self, - model_name: str, - metadata: dict[str, Any], - model_type: str, - flavor: str, - load_wrapper: bool = False, - ) -> tuple[Any, str | None]: - """ - Download model based on type ("predict" or "transform"). - - Args: - model_name (str): Name of the model to download - metadata (dict[str, Any]): Metadata used for structured logging. - model_type (str): Type of model ('predict' or 'transform') - flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch') - load_wrapper (bool): Whether to load wrapper - - Returns: - tuple[Any, str | None]: Model object and optional artifact path. - """ - - self.info( - f'Downloading {model_type} model {model_name} with flavor {flavor} and load_wrapper {load_wrapper}' - ) - - if model_type not in ['predict', 'transform']: - raise ValueError("Invalid model_type. Use 'predict' or 'transform'.") - - artifact_path = None - - if load_wrapper: - self.info(f'Loading wrapper for {model_type} model {model_name} with flavor {flavor}') - - target = 'prediction_model' if model_type == 'predict' else 'data_model' - - artifact_path = await self.dowload_artifacts(model_name, metadata, target) - - self.info( - f'Model with type {model_type} and name {model_name} is compressed, loading from {artifact_path}' - ) - - raw_model = mlflow.pyfunc.load_model(artifact_path) - model = raw_model._model_impl.python_model - - self.debug( - f'Model wrapper loaded: {model.__class__.__name__}:{model.__dict__}', metadata - ) - else: - if model_type == 'predict': - model = await self.load_predict_model(model_name, metadata, flavor) - - else: - model = await self.load_transform_model(model_name, metadata, flavor) - - return model, artifact_path - - """ - Functions related to data format - """ - - def detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame: - """ - Normalize DataFrame index to the expected timestamp string format. - - The index must be timestamp-like. If the index is: - - string: it must match `DATETIME_FORMAT_WITH_TZ` - - datetime or pandas Timestamp: it will be converted to that format - Any other type raises a ValueError. - - Args: - data (pd.DataFrame): DataFrame with timestamp index. - metadata (dict): Metadata for structured logging. - - Returns: - pd.DataFrame: DataFrame with converted datetime index. - """ - if data.empty: - self.info('Data is empty, skipping datetime index detection and parsing', metadata) - return data - - index = data.index - - # Get type of first element of index - index_type = type(index[0]) - - self.info(f'Index type: {index_type}', metadata) - - message = f'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}.' - - # Check if all in index are of the same type - if not all(isinstance(i, index_type) for i in index): - types = map(str, map(type, index)) - raise ValueError(f'{message}. Elements are {",".join(types)}') - - # Check type and converts to DATETIME_FORMAT_WITH_TZ - if index_type is str: - # Validate format of string and return error if not valid - try: - pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ) - except ValueError as e: - raise ValueError(f'{message}. Unable to parse given date format: {e}') from e - - elif index_type == datetime or index_type == pd.Timestamp: - index = data.index - if hasattr(index, 'tz') and index.tz is None: - data.index = index.tz_localize('UTC') # type: ignore[attr-defined] - - data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ) # type: ignore[attr-defined] - else: - raise ValueError(f'{message}. Got {index_type}.') - - return data - - """ - Functions related to cache management of models - """ - - def check_cache_retention(self, cache: dict, retention: int) -> bool: - """ - Check whether cached model data is still valid. - - Args: - cache (dict): Cached model data with a 'timestamp' key. - retention (int): Retention time in minutes. - - Returns: - bool: True if cache is still valid, False if expired. - """ - current_time = datetime.now() - cache_time = cache['timestamp'] - if current_time - cache_time >= timedelta(minutes=retention): - return False - return True - - def handle_valid_model(self, model_name: str, cache: dict) -> dict: - """ - Return the cached model configuration when retention is valid. - - Args: - model_name (str): Name of the model (for logging/consistency). - cache (dict): Cached model data structure. - - Returns: - dict: Model configuration. - """ - self.debug(f'Model {model_name} is still valid, using cached version') - - return cache['target'] - - def handle_outdated_model(self, model_name: str, model_key: str) -> None: - """ - Clean up outdated cached model and its artifacts. - - Args: - model_name (str): Name of the model. - model_key (str): Cache key for the model. - - Returns: - None - """ - self.debug(f'Model {model_name} is outdated, downloading a new one') - - del self.model_cache[model_key]['target'] - del self.model_cache[model_key] - - async def get_model( - self, - model_name: str, - metadata: dict[str, Any], - retention: int, - model_type: str, - flavor: str, - ) -> Any: - """ - Retrieve a model with caching support based on retention policy. - - Args: - model_name (str): Name of the model to retrieve - metadata (dict[str, Any]): Metadata used for structured logging. - retention (int): Cache retention time in minutes (0 = no cache). - model_type (str): Type of model ('predict' or 'transform') - flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch') - - Returns: - Any: Model object. - """ - # Retention is 0, download a new model - if retention <= 0: - model, _artifact_path = await self.download_model( - model_name=model_name, - metadata=metadata, - model_type=model_type, - flavor=flavor, - load_wrapper=False, - ) - return model - - model_key = f'{model_name}_{model_type}' - - # Acquire lock to check cache - with self._cache_lock: - if model_key in self.model_cache: - cache = self.model_cache[model_key] - - # Check if config has changed or is outdated - if self.check_cache_retention(cache, retention): - return self.handle_valid_model(model_name=model_name, cache=cache) - else: - # Model is outdated, delete old model files - self.handle_outdated_model(model_name=model_name, model_key=model_key) - else: - self.debug( - f'Model {model_name} is not in {model_type} cache, downloading a new one' - ) - - # Donwload new model (without lock to avoid blocking other threads) - model, _artifact_path = await self.download_model( - model_name=model_name, - metadata=metadata, - model_type=model_type, - flavor=flavor, - load_wrapper=False, - ) - - # Update cache with lock - with self._cache_lock: - cache = {'target': model, 'timestamp': datetime.now()} - self.model_cache[model_key] = cache - - return model - - @overload - async def get_cached_operation( - self, - model_name: str, - data: pd.DataFrame, - operation: Literal['transform'], - retention: int, - flavor: str, - metadata: dict[str, Any], - ) -> pd.DataFrame: ... - - @overload - async def get_cached_operation( - self, - model_name: str, - data: pd.DataFrame, - operation: Literal['predict'], - retention: int, - flavor: str, - metadata: dict[str, Any], - ) -> pd.DataFrame | ndarray: ... - - async def get_cached_operation( - self, - model_name: str, - data: pd.DataFrame, - operation: str, - retention: int, - flavor: str, - metadata: dict[str, Any], - ) -> pd.DataFrame | ndarray: - """ - Execute a cached operation using the requested model. - - Args: - model_name (str): Registered model name. - data (pd.DataFrame): Input data. - retention (int): Cache retention in minutes. - flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch'). - metadata (dict[str, Any]): Metadata used for structured logging. - Returns: - pd.DataFrame | ndarray: Operation result. - """ - if operation not in ['transform', 'predict']: - raise ValueError("Invalid operation. Use 'transform' or 'predict'.") - - model = await self.get_model( - model_name=model_name, - metadata=metadata, - retention=retention, - model_type=operation, - flavor=flavor, - ) - - prediction = model.predict(data) - - if retention == 0: - self.info(f'Deleting model {model_name}:{operation} from memory') - del model - - force_memory_release(self.logger) - - return prediction - - """ - Functions related to model retraining - """ - - def get_prediction_data( - self, - prediction_model: Any, - retrain_dataset: pd.DataFrame, - target_name: str, - predict_flavor: str, - ) -> pd.DataFrame: - """ - Get prediction data from prediction model. - """ - input_index = retrain_dataset.index - - if predict_flavor == 'pyfunc': - prediction_data = prediction_model.predict({}, retrain_dataset) - else: - prediction_data = prediction_model.predict(retrain_dataset) - - if isinstance(prediction_data, pd.DataFrame): - prediction_data.columns = pd.Index(['prediction']) - - else: - prediction_data = pd.DataFrame(prediction_data, columns=['prediction']) - - prediction_data.index = input_index - - # Merge prediction data with retrain_dataset on index - prediction_data = pd.merge( # NOSONAR - retrain_dataset, prediction_data, left_index=True, right_index=True, how='left' - ) - - # Rename column "target_name" to "target" - prediction_data.rename(columns={target_name: 'target'}, inplace=True) - - prediction_data['timestamp'] = prediction_data.index - - prediction_data.reset_index(drop=True, inplace=True) - - prediction_data = prediction_data.sort_values(by='timestamp', ascending=True) - - return prediction_data - - async def fit_models( - self, - model_name: str, - data: pd.DataFrame, - latest_production_id: str, - metadata: dict, - transform_flavor: str = 'sklearn', - skip_transform: bool = False, - predict_flavor: str = 'sklearn', - target_name: str | None = None, - ) -> dict[str, Any]: - """ - Prepare models and data for a retraining run. - - This method sets up the complete environment for model retraining by: - 1. Loading the current production prediction model - 2. Loading the current production transformation model - 3. Fitting the transformation model with new data - 4. Preparing data for prediction model retraining - 5. Setting up the MLFlow experiment context - - Args: - model_name (str): Name of the MLflow model to retrain. - data (pd.DataFrame): Training data for model retraining. - transform_flavor (str): Flavor for transformation model. - predict_flavor (str): Flavor for prediction model. - target_name (str | None): Optional target column; if None, use model target. - metadata (dict): Metadata for logging. - - Returns: - dict[str, dict[str, Any]]: Mapping with prepared `prediction_model` and - `data_model`, including optional artifact paths. - """ - - self.info(f'Starting model experiment creation for {model_name}', metadata) - self.debug( - f'Model configuration - transform_flavor: {transform_flavor}, predict_flavor: {predict_flavor}, target_name: {target_name}', - metadata, - ) - - # data.to_csv( - # f"tmp/retrain_data_{model_name}.csv", index=True) - - self.info(f'Retrieved latest production run ID: {latest_production_id}', metadata) - self.info(f'Loading transformation model for {model_name}', metadata) - - load_transform_wrapper = transform_flavor == 'pyfunc' - - data_model, data_artifact_path = await self.download_model( - model_name=model_name, - metadata=metadata, - model_type='transform', - flavor=transform_flavor, - load_wrapper=load_transform_wrapper, - ) - - self.info(f'Loading prediction model for {model_name}', metadata) - - load_predict_wrapper = predict_flavor == 'pyfunc' - - prediction_model, prediction_artifact_path = await self.download_model( - model_name=model_name, - metadata=metadata, - model_type='predict', - flavor=predict_flavor, - load_wrapper=load_predict_wrapper, - ) - - if not skip_transform: - treated_data_candidate = data_model.fit(data) - else: - treated_data_candidate = data_model - - if not isinstance(treated_data_candidate, pd.DataFrame): - data_model = treated_data_candidate - treated_data = data_model.predict(data) - else: - treated_data = treated_data_candidate - - # Stores current index as timestamp, Courier model expects a timestamp column - # with specific format - treated_data['timestamp'] = treated_data.index - - # Parses timestamp column to datetime format to align with data - treated_data = self.detect_and_parse_datetime_index(treated_data, metadata) - - treated_data = treated_data.drop_duplicates(subset=['timestamp'], keep='first') - - self.debug(f'Treated data index: {treated_data.index}', metadata) - - # treated_data.to_csv( - # f"tmp/retrain_treated_data_{model_name}.csv", index=True) - - self.debug(f'Transformed data shape: {treated_data.shape}', metadata) - - if target_name is None: - target_name = data_model.target_variable - self.debug(f'Using target variable from data model: {target_name}', metadata) - else: - self.debug(f'Using provided target variable: {target_name}', metadata) - - # Check if treated_data contains target variable - if target_name not in treated_data.columns: - self.debug( - f'Target variable {target_name} not found in treated data, aligning data with treated data indexes', - metadata, - ) - # Aligns data with treated data indexes to get target variable - aligned_data = data.loc[treated_data.index] - aligned_series = aligned_data[target_name] - retrain_dataset = pd.merge( # NOSONAR - treated_data, aligned_series, left_index=True, right_index=True - ) - else: - # Uses target variable from treated data - self.debug(f'Target variable {target_name} found in treated data, using it', metadata) - retrain_dataset = treated_data - - # retrain_dataset.to_csv( - # f"tmp/retrain_retrain_dataset_{model_name}.csv", index=True) - - prediction_model.fit(retrain_dataset) - - # get prediction data - prediction_data = self.get_prediction_data( - prediction_model, retrain_dataset, target_name, predict_flavor - ) - - self.info(f'Model experiment creation completed successfully for {model_name}', metadata) - - retrain_data = { - 'prediction_model': { - 'model': prediction_model, - 'artifact_path': prediction_artifact_path, - }, - 'data_model': {'model': data_model, 'artifact_path': data_artifact_path}, - 'prediction_data': prediction_data, - } - return retrain_data - - async def log_model(self, model_data: dict, flavor: str, model_type: str, metadata: dict): - """Log a model into the active MLflow run. - - Args: - model_data (dict): Model holder with keys 'model' and optional 'artifact_path'. - flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch'). - model_type (str): Artifact name, e.g., 'prediction_model' or 'data_model'. - metadata (dict): Metadata for structured logging. - """ - model = model_data['model'] - - self.debug(f'Logging {model_type} model to {model_type}', metadata) - - core_labels = self.get_core_labels(metadata, operation_type='log_model') - start_time = time.time() - - try: - if flavor == 'sklearn': - mlflow.sklearn.log_model(model, model_type) - elif flavor == 'pyfunc': - code_path = [path.join(model_data['artifact_path'], 'code', 'utils')] - - self.debug(f'Code path: {code_path}', metadata) - - model.store_model(artifact_path=model_type, code_path=code_path, to_disk=False) - - self.debug('Model uploaded successfully', metadata) - elif flavor == 'pytorch': - mlflow.pytorch.log_model(model, model_type) - else: - raise ValueError(INVALID_FLAVOR_MESSAGE) - - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels) - - async def create_new_experiment( - self, - model_name: str, - data: pd.DataFrame, - retrain_data: dict, - latest_production_id: str, - metadata: dict, - transform_flavor: str = 'sklearn', - predict_flavor: str = 'sklearn', - ) -> dict: - """ - Execute the complete model retraining process in MLflow. - - This method performs the actual model retraining by: - 1. Starting a new MLFlow run with descriptive metadata - 2. Logging model parameters and hyperparameters - 3. Retraining both prediction and transformation models - 4. Logging training data as artifacts - 5. Saving retrained models to MLFlow registry - - Args: - prediction_model: MLFlow prediction model to retrain - data_model: MLFlow transformation model to retrain - experiment (str): MLflow experiment name for the retraining. - model_name (str): Name of the model being retrained. - data (pd.DataFrame): Training data used for retraining. - transform_flavor (str): Flavor for transformation model. - predict_flavor (str): Flavor for prediction model. - metadata (dict): Metadata for logging. - - Returns: - dict: Metadata about the created run and experiment. - """ - - prediction_model = retrain_data['prediction_model'] - data_model = retrain_data['data_model'] - prediction_data = retrain_data['prediction_data'] - - model_temp_path = path.join(ARTIFACTS_PATH, model_name) - - self.info(f'Starting model retraining process for {model_name}', metadata) - - original_params = self.get_model_params(latest_production_id) - retrain_params = { - **original_params, - 'retrain': True, - 'retrain_date': datetime.now().isoformat(), - 'source_run_id': latest_production_id, - 'retrain_samples': str(data.shape), - } - experiment_description = f'Retrain model {model_name} with new data' - - experiment = self.get_experiment(model_name, create_if_not_exists=True) - experiment_name = experiment.name - - current_run_name = self.get_next_run_name(experiment_name) - - self.debug(f'Attributes: {retrain_params}', metadata) - - data_path = f'{model_temp_path}/retrain_data.csv' - prediction_data_path = f'{model_temp_path}/evaluation_data.csv' - - makedirs(model_temp_path, exist_ok=True) - - data.to_csv(data_path, index=False) - prediction_data.to_csv(prediction_data_path, index=False) - - self.info( - f'Starting model upload for {experiment_name} with run name {current_run_name}', - metadata, - ) - - core_labels = self.get_core_labels(metadata, operation_type='create_new_experiment') - start_time = time.time() - try: - with mlflow.start_run( - experiment_id=experiment.experiment_id, - run_name=current_run_name, - description=experiment_description, - ) as _run: - run_id = _run.info.run_id - self.info('Logging data model', metadata) - # dynamic parameters, including model itself - await self.log_model(data_model, transform_flavor, 'data_model', metadata) - - # dynamic parameters, including model itself - self.info('Logging prediction model', metadata) - await self.log_model(prediction_model, predict_flavor, 'prediction_model', metadata) - - self.info(f'Model logged successfully for {model_name}', metadata) - - self.info(f'Logging remaining parameters for {model_name}', metadata) - - # update transfomation model - # fixed parameters - mlflow.log_params(retrain_params) - - # log the data raw - mlflow.log_artifact(data_path) - mlflow.log_artifact(prediction_data_path) - - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels) - - self.info('Deleting model from filesystem', metadata) - if path.exists(model_temp_path): - rmtree(model_temp_path) - - self.info('Deleting prediction model from memory', metadata) - del prediction_model['model'] - del prediction_model - - self.info('Deleting data model from memory', metadata) - del data_model['model'] - del data_model - - force_memory_release(self.logger) - - return { - 'run_id': run_id, - 'experiment_id': experiment.experiment_id, - 'experiment_name': experiment.name, - } - - async def update_production_model_by_run_id( - self, run_id: str, model_name: str, metadata: dict - ) -> dict: - """ - Promote a specific run's model to Production. - - This method promotes a model from a specific MLFlow run to - production stage. It handles model registration, versioning, - and stage transitions with proper error handling. - - Args: - run_id (str): MLflow run ID containing the model to promote. - model_name (str): Name of the MLflow model. - metadata (dict): Metadata for logging. - - Returns: - dict: Model update metadata containing: - - model_name (str): Name of the updated model - - version (str): New model version number - - mlflow_run_id (str): Source run ID - - Update Process: - 1. Registers the model from the specified run - 2. Retrieves the latest model version - 3. Transitions the model to 'Production' stage - 4. Archives existing production versions - """ - - self.info( - f'Starting production model update for {model_name} with run ID: {run_id}', metadata - ) - - # Registrar o modelo - # Aqui estamos assumindo que você já tem um modelo salvo, caso contrário você precisará treiná-lo e salvá-lo primeiro. - # Se o modelo já está registrado, você pode usar o método register_model() ou pyfunc.load_model() para isso. - - core_labels = self.get_core_labels(metadata, operation_type='register_model') - start_time = time.time() - try: - mlflow.register_model(f'runs:/{run_id}/prediction_model', model_name) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels) - - # Obter a versão mais recente registrada do modelo - model_versions = self.client.get_registered_model(model_name).latest_versions - - if not isinstance(model_versions, list): - raise ValueError('Model versions is not a list') - - max_version = max(model_versions, key=lambda x: int(x.version)).version - - # Mover a versão mais recente do modelo para o estágio de 'Production' - core_labels = self.get_core_labels( - metadata, operation_type='transition_model_version_stage' - ) - start_time = time.time() - try: - self.client.transition_model_version_stage( - name=model_name, - version=max_version, - stage='Production', - archive_existing_versions=True, - ) - except Exception as e: - await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels) - raise e - - await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels) - await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels) - - return {'model_name': model_name, 'version': max_version, 'mlflow_run_id': run_id} - - """ - Functions that provide the interface to model operations - """ - - async def transform( - self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict - ) -> dict[str, Any]: - """ - Transform data using a cached transformation model. - - This method provides a high-level interface for data transformation operations - using MLFlow models. It handles model caching, error management, and data - format conversion automatically. - - Process Flow: - 1. Retrieves or downloads the transformation model using caching mechanism - 2. Applies the transformation model to the input data - 3. Converts the transformed data to dictionary format for API response - 4. Handles any exceptions and returns structured error information - 5. Manages model lifecycle based on retention policy (cleanup artifacts if needed) - - Parameters: - model_name (str): The name of the MLflow model to use for transformation. - data (pd.DataFrame): The input data to be transformed by the model. - model_config (dict): Model configuration parameters - metadata (dict): Metadata for logging - - Returns: - dict: Response dictionary containing: - - success (bool): Operation success status - - content (dict): Transformed data as dictionary, or error information - if operation failed. Error content includes: - - message (str): Error description - - traceback (str): Full exception traceback - - Raises: - Exception: Any exception during model loading or transformation is caught - and returned in the response structure rather than propagated. - """ - - self._debug_dataframe('Data received for model transformation:', data, metadata) - - # data.to_csv( - # f"tmp/data_{model_name}.csv", index=True) - - model_retention = model_config.get('retention_minutes', 0) - flavor = model_config.get('transform_flavor', 'sklearn') - - try: - transformed_data: pd.DataFrame = await self.get_cached_operation( - model_name=model_name, - data=data, - operation='transform', - retention=model_retention, - flavor=flavor, - metadata=metadata, - ) - - self._debug_dataframe( - 'Data received from model transformation:', transformed_data, metadata - ) - - # transformed_data.to_csv( - # f"tmp/transformed_data_{model_name}.csv", index=True) - - transformed_data = self.detect_and_parse_datetime_index(transformed_data, metadata) - - return {'success': True, 'content': transformed_data} - - except Exception as e: - return { - 'success': False, - 'content': {'message': str(e), 'traceback': traceback.format_exc()}, - } - - async def predict( - self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict - ): - """ - Generate predictions using a cached prediction model. - - This method provides a high-level interface for model prediction operations - using MLFlow models. It handles model caching, performance monitoring, - response formatting, and error management automatically. - - Process Flow: - 1. Preserves input data index for result alignment - 2. Records prediction start time for performance measurement - 3. Retrieves or downloads the prediction model using caching mechanism - 4. Executes model prediction on the input data - 5. Formats predictions into DataFrame with proper column naming - 6. Restores original data index to maintain data alignment - 7. Calculates and adds response time measurement - 8. Converts results to dictionary format for API response - 9. Handles any exceptions and returns structured error information - - Parameters: - model_name (str): The name of the MLflow model to use for prediction. - data (pd.DataFrame): The input data to make predictions on. - model_retention (int): Cache retention time in minutes (0 = no caching). - model_config (dict): Model configuration parameters - metadata (dict): Metadata for logging - - Returns: - dict: Response dictionary containing: - - success (bool): Operation success status - - content (dict): Prediction results as dictionary with: - - prediction: Model predictions array - - response_time: Prediction execution time in seconds - Or error information if operation failed: - - message (str): Error description - - traceback (str): Full exception traceback - - Raises: - Exception: Any exception during model loading or prediction is caught - and returned in the response structure rather than propagated. - """ - - model_retention = model_config.get('retention_minutes', 0) - flavor = model_config.get('predict_flavor', 'sklearn') - - try: - input_index = data.index - start_time = datetime.now() - - self._debug_dataframe('Data received for model prediction:', data, metadata) - - # data.to_csv( - # f"tmp/treated_data_{model_name}.csv", index=True) - - predict_data = await self.get_cached_operation( - model_name=model_name, - data=data, - operation='predict', - retention=model_retention, - flavor=flavor, - metadata=metadata, - ) - - end_time = datetime.now() - - if isinstance(predict_data, pd.DataFrame): - self._debug_dataframe( - 'Data received from model prediction:', predict_data, metadata - ) - - # predict_data.to_csv( - # f"tmp/predicted_data_{model_name}.csv", index=True) - predict_data.columns = pd.Index(['prediction']) - - else: - self.debug( - f'Data received from model prediction (not a DataFrame): {predict_data}', - metadata, - ) - predict_data = pd.DataFrame(predict_data, columns=['prediction']) - # predict_data.to_csv( - # f"tmp/predicted_data_{model_name}.csv", index=True) - - predict_data.index = input_index - predict_data['response_time'] = (end_time - start_time).total_seconds() - - return {'success': True, 'content': predict_data} - - except Exception as e: - return { - 'success': False, - 'content': {'message': str(e), 'traceback': traceback.format_exc()}, - } - - async def retrain_model( - self, data: pd.DataFrame, model_name: str, model_config: dict, metadata: dict - ) -> dict[str, Any]: - """ - Orchestrate the complete model retraining workflow. - - This method coordinates the entire model retraining process by managing - the MLFlow experiment lifecycle, model loading, retraining execution, - and artifact management. It provides a comprehensive retraining solution - that maintains model versioning and experiment tracking. - - Process Flow: - 1. Creates MLFlow experiment environment: - - Loads current production prediction model - - Loads current production transformation model - - Fits transformation model with new training data - - Prepares transformed data for prediction model retraining - - Sets up MLFlow experiment context - 2. Executes model retraining: - - Starts new MLFlow run with descriptive metadata - - Logs model parameters and hyperparameters - - Retrains both prediction and transformation models - - Logs training data as artifacts - - Saves retrained models to MLFlow registry - - Cleans up temporary files - 3. Returns comprehensive retraining results - - Args: - data (pd.DataFrame): Training data for model retraining. Must contain - all features required by both transformation and - prediction models, including target variable. - model_name (str): Name of the MLflow model to retrain. Must exist - in the MLflow Model Registry in Production stage. - model_config (dict): Model configuration parameters - metadata (dict): Metadata for logging - - Returns: - dict: Retraining operation results and experiment details. - - Raises: - mlflow.exceptions.MlflowException: If model not found in registry - ValueError: If experiment cannot be created or models cannot be loaded - Exception: Any other exception during the retraining process - """ - - self.info(f'Starting model retraining workflow for {model_name}', metadata) - self._debug_dataframe('Data received for model retraining:', data, metadata) - - target_name = model_config.get('target', None) - - transform_flavor = model_config.get('transform_flavor', 'sklearn') - predict_flavor = model_config.get('predict_flavor', 'sklearn') - skip_transform = model_config.get('skip_transform', False) - - self.debug( - f'Model configuration - transform_flavor: {transform_flavor}, predict_flavor: {predict_flavor}, target_name: {target_name}', - metadata, - ) - - try: - latest_production_id = self.get_model_run_id(model_name, stage='Production') - self.info('Creating model experiment environment', metadata) - retrain_data = await self.fit_models( - model_name=model_name, - data=data, - transform_flavor=transform_flavor, - skip_transform=skip_transform, - predict_flavor=predict_flavor, - target_name=target_name, - metadata=metadata, - latest_production_id=latest_production_id, - ) - self.info(f'Model experiment created successfully: {retrain_data}', metadata) - - self.info('Saving model retrain', metadata) - experiment = await self.create_new_experiment( - model_name=model_name, - data=data, - retrain_data=retrain_data, - transform_flavor=transform_flavor, - predict_flavor=predict_flavor, - metadata=metadata, - latest_production_id=latest_production_id, - ) - self.info( - f'Model retraining completed successfully for experiment: {experiment}', metadata - ) - - return { - 'success': True, - 'experiment': experiment, - 'message': 'Model retrained successfully.', - } - except Exception as e: - error_msg = f'Error retraining model {model_name}: {e}' - self.info(error_msg, metadata) - return { - 'success': False, - 'experiment': None, - 'message': error_msg, - 'traceback': traceback.format_exc(), - } - - async def update_production_model( - self, experiment: dict[str, Any], model_name: str, metadata: dict - ) -> dict: - """ - Update production model using the latest retraining run. - - This method orchestrates the complete production model update process by - identifying the most recent retraining run and promoting it to production - stage. It handles model registration, versioning, and stage transitions - with comprehensive metadata tracking. - - Process Flow: - 1. Retrieves experiment information: - - Converts experiment name to MLFlow experiment ID - - Searches for the most recent retraining run in the experiment - - Filters runs by 'retrain' parameter and orders by completion time - 2. Promotes model to production: - - Registers the model from the specified run to MLFlow Model Registry - - Retrieves the latest model version number - - Transitions the model to 'Production' stage - - Archives existing production versions automatically - 3. Returns comprehensive update metadata - - Args: - experiment (str): MLflow experiment name containing the retraining runs. - Must be a valid experiment that exists in MLFlow. - model_name (str): Name of the MLFlow model to update. Must exist - in the MLFlow Model Registry. - metadata (dict): Metadata for logging - - Returns: - dict: Complete model update metadata containing: - - model_name (str): Name of the updated model - - version (str): New model version number (incremented automatically) - - mlflow_run_id (str): Source run ID of the promoted model - - mlflow_experiment_id (int): Experiment ID for tracking - - Raises: - ValueError: If experiment not found or model versions are invalid - mlflow.exceptions.MlflowException: If model registration or stage - transition fails - Exception: Any other exception during the update process - - Note: - This operation is irreversible. The previous production model will - be automatically archived when the new version is promoted. - """ - run_id = experiment['run_id'] - experiment_id = experiment['experiment_id'] - metadata_result = await self.update_production_model_by_run_id(run_id, model_name, metadata) - - metadata_result['mlflow_experiment_id'] = experiment_id - - return metadata_result diff --git a/laborious/worker/prepare_worker.py b/laborious/worker/prepare_worker.py deleted file mode 100644 index 9766af2..0000000 --- a/laborious/worker/prepare_worker.py +++ /dev/null @@ -1,73 +0,0 @@ -import os -import re -from collections.abc import Sequence -from typing import Any - -from sientia_do.observability.logger import Logger -from temporalio.client import Client -from temporalio.worker import PollerBehaviorAutoscaling, Worker - -parameters = [ - ('MAX_CONCURRENT_WORKFLOW_TASKS', '200'), - ('MAX_CONCURRENT_ACTIVITIES', '200'), - ('MAX_CONCURRENT_LOCAL_ACTIVITIES', '200'), - ('MAX_CACHED_WORKFLOWS', '200'), - ('WORKFLOW_POLLER_BEHAVIOUR_MINIMUM', '10'), - ('WORKFLOW_POLLER_BEHAVIOUR_INITIAL', '100'), - ('WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM', '200'), - ('ACTIVITY_POLLER_BEHAVIOUR_MINIMUM', '10'), - ('ACTIVITY_POLLER_BEHAVIOUR_INITIAL', '100'), - ('ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM', '200'), -] - - -def camel_to_snake(text: str) -> str: - """Convert camelCase or PascalCase to snake_case.""" - text = re.sub('(.)([A-Z][a-z]+)', r'\1_\2', text) - text = re.sub('([a-z0-9])([A-Z])', r'\1_\2', text) - return text.lower() - - -def prepare_worker( - main_workflow: type, - other_workflows: Sequence[type], - activities: Sequence[Any], - temporal_client: Client, - logger: Logger, -) -> Worker: - main_workflow_name = main_workflow.__name__.upper() - - queue_name = f'{camel_to_snake(main_workflow.__name__)}-queue' - - local_workflow_parameters = {} - - for parameter in parameters: - local_workflow_parameters[parameter[0]] = int( - os.getenv(main_workflow_name + '_' + parameter[0], parameter[1]) - ) - - logger.info(f'Preparing worker for {main_workflow_name} with queue {queue_name}') - logger.info(f'Worker runtime config: {local_workflow_parameters}') - - return Worker( - temporal_client, - task_queue=queue_name, - workflows=[main_workflow, *other_workflows], - activities=[*activities], - max_concurrent_workflow_tasks=local_workflow_parameters['MAX_CONCURRENT_WORKFLOW_TASKS'], - max_concurrent_activities=local_workflow_parameters['MAX_CONCURRENT_ACTIVITIES'], - max_concurrent_local_activities=local_workflow_parameters[ - 'MAX_CONCURRENT_LOCAL_ACTIVITIES' - ], - max_cached_workflows=local_workflow_parameters['MAX_CACHED_WORKFLOWS'], - workflow_task_poller_behavior=PollerBehaviorAutoscaling( - minimum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MINIMUM'], - initial=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_INITIAL'], - maximum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM'], - ), - activity_task_poller_behavior=PollerBehaviorAutoscaling( - minimum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MINIMUM'], - initial=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_INITIAL'], - maximum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM'], - ), - ) diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index fa09259..8280caf 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -1,32 +1,32 @@ """ Laborious Worker Module -This module provides the main worker implementation for the Sientia DataOps Laborious system. -It orchestrates Temporal workers, manages task queues, and handles the lifecycle of -prediction and retraining workflows. +Entry process that connects to Temporal, registers Laborious activities, and runs four workers in +parallel. Each worker shares the same ``Activities`` instance (single Postgres pool, single MLflow +repository, single PluginStore handle) but polls a different task queue. -The worker supports multiple task queues: -- predictions_batch-queue: Handles batch prediction workflows (heavy workload) - Includes activities for MLFlow, data quality gates, OPC export, PI Web API export, and PostgreSQL -- minimal_retrain-queue: Handles model retraining workflows -- drift-queue: Handles drift detection workflows -- simple_metrics-queue: Handles simple metrics calculation workflows +Task queues (see ``sientia_do.temporal.worker.prepare_worker``): +- ``predictions_batch-{runtime}-queue`` + sub-workflows on the same queue (ML-heavy path). +- ``minimal_retrain-{runtime}-queue`` (retrain + promote + export). +- ``drift-queue`` and ``simple_metrics-queue`` without a runtime suffix so existing schedulers + keep stable queue names. -Key Features: -- Resource-based scaling with WorkerTuner (CPU and memory aware) -- Automatic polling scaling with PollerBehaviorAutoscaling -- Prometheus metrics integration -- Comprehensive error handling and logging -- Graceful shutdown with cleanup -- Multiple worker instances for different workflow types +Bootstrap order: +1. Prometheus app metrics and Mongo-backed notification handler. +2. ``RUNTIME`` validation and ``PluginStore.install_runtime`` so ``SientiaModel`` code is importable. +3. ``Activities`` construction (builds ``SientiaMLflowRepository`` internally from env). +4. OPC client initialization inside activities. +5. Temporal ``Runtime`` with SDK Prometheus bind, client connect, then ``prepare_worker`` per workflow. + +Shutdown closes workers, notifications, activities (pools + OPC), and clears ``app_up``. Environment Variables: -- TEMPORAL_HOST: Temporal server address (default: localhost:7233) -- TEMPORAL_NAMESPACE: Temporal namespace (default: laborious) -- POD_ID: Kubernetes pod identifier for metrics -- HTTP_METRICS_PORT: Prometheus metrics server port (default: 9090) -- HTTP_SDK_METRICS_PORT: Temporal SDK metrics port (default: 9091) -- PROJECT_NAME: Project name for notifications (default: laborious) +- RUNTIME: Required non-empty string passed to ``install_runtime``. +- STORE_* / PYPI_*: Plugin store and private index (see ``build_plugin_store_config``). +- TEMPORAL_HOST, TEMPORAL_NAMESPACE: Cluster connection. +- POD_ID, HTTP_METRICS_PORT, HTTP_SDK_METRICS_PORT: Observability. +- PROJECT_NAME, MONGODB_*: Notifications (via ``build_mongodb_config`` in handler). +- POSTGRES_*, MINIO_*, OPC_*, PI_WEB_API_*, MLFLOW_*: Passed through ``Activities`` helpers. """ from temporalio import client, workflow @@ -40,20 +40,22 @@ with workflow.unsafe.imports_passed_through(): from prometheus_client import start_http_server from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.observability.logger import get_logger + from sientia_do.observability.metrics_controller import MetricsController + from sientia_do.temporal.worker.prepare_worker import prepare_worker from sientia_do.utils.connectors_config import ( build_api_config, build_mongodb_config, build_postgres_config, ) + from sientia_model.model_repository.plugin_store import PluginStore from laborious import metrics from laborious.activities.activities import Activities from laborious.utils.connectors_config import ( build_minio_config, - build_mlflow_config, build_opc_config, + build_plugin_store_config, ) - from laborious.worker.prepare_worker import prepare_worker from laborious.workflows.drift import Drift from laborious.workflows.minimal_retrain import MinimalRetrain from laborious.workflows.predictions_batch import PredictionsBatch @@ -69,23 +71,18 @@ SDK_METRICS_PORT = int(os.getenv('HTTP_SDK_METRICS_PORT', '9091')) async def main(): """ - Main entry point for the Laborious worker application. + Run the full worker lifecycle: metrics, notifications, runtime install, workers, gather. - This function initializes and starts all components of the worker: - 1. Sets up logging and metadata - 2. Starts Prometheus metrics server - 3. Initializes notification handler - 4. Creates and configures activities - 5. Initializes OPC connections - 6. Starts Temporal client and workers - 7. Manages worker lifecycle and graceful shutdown - - The function runs indefinitely until interrupted or an error occurs. - On error, it performs cleanup and exits with a non-zero status code. + Exits the process with code 0 on normal completion of all worker tasks, or 1 after logging + if any worker raises. ``finally`` always shuts down notifications and activities and sets + ``app_up`` to 0 before ``sys.exit``. Raises: - Exception: Any unhandled exception during worker execution - SystemExit: On graceful shutdown or error conditions + Exception: Propagated from ``asyncio.gather`` only before ``finally`` handling; typically + workers run until cancelled. + + Return: + None (process terminates via ``sys.exit`` from the ``finally`` block). """ host = os.getenv('TEMPORAL_HOST', 'localhost:7233') logger = get_logger(__name__) @@ -113,16 +110,55 @@ async def main(): project_name=os.getenv('PROJECT_NAME', 'laborious'), ) + metrics_controller = MetricsController(logger=logger) + + runtime = os.getenv('RUNTIME', '').strip() + if not runtime: + logger.custom_critical( + 'RUNTIME environment variable is required and must be non-empty', + metadata, + ) + metrics.APP_UP.labels(pod_id=POD_ID).set(0) + sys.exit(1) + + metadata_runtime = {**metadata, 'runtime': runtime} + logger.custom_info(f'Installing PluginStore runtime: {runtime}', metadata_runtime) + + ps_cfg = build_plugin_store_config() + plugin_store = PluginStore( + base_url=ps_cfg['base_url'], + owner=ps_cfg['owner'], + repo=ps_cfg['repo'], + username=ps_cfg['username'], + password=ps_cfg['password'], + branch=ps_cfg['branch'], + cache_ttl_seconds=ps_cfg['cache_ttl_seconds'], + pypi_index_url=ps_cfg['pypi_index_url'], + pypi_username=ps_cfg['pypi_username'], + pypi_password=ps_cfg['pypi_password'], + logger=logger, + notification_handler=notification_handler, + metrics_controller=metrics_controller, + ) + + try: + await plugin_store.install_runtime(runtime_name=runtime, metadata=metadata_runtime) + except Exception as exc: + logger.custom_critical(f'Failed to install runtime {runtime}: {exc}', metadata_runtime) + metrics.APP_UP.labels(pod_id=POD_ID).set(0) + sys.exit(1) + logger.custom_info('Starting Activities...', metadata) activities = Activities( postgres_config=build_postgres_config(), - mlflow_config=build_mlflow_config(), + plugin_store=plugin_store, minio_config=build_minio_config(), opc_config=build_opc_config(), pi_web_api_config=build_api_config(), logger=logger, notification_handler=notification_handler, + metrics_controller=metrics_controller, ) logger.custom_info('Initializing OPC...', metadata) @@ -159,6 +195,7 @@ async def main(): activities.export_data_to_postgres, ], logger=logger, + runtime=runtime, ), prepare_worker( temporal_client=temporal_client, @@ -210,6 +247,7 @@ async def main(): activities.write_pi_web_api_data, ], logger=logger, + runtime=runtime, ), ] @@ -221,8 +259,6 @@ async def main(): exit_code = 0 try: - # This will run the workers and wait for them to complete. - # If an exception occurs in any of the worker handlers, it will be propagated here. await asyncio.gather(*handlers) except BaseException as e: # NOSONAR logger.custom_error(f'An unhandled exception occurred: {e}', metadata) diff --git a/requirements-light.txt b/requirements-light.txt index 91a9c33..f105052 100644 --- a/requirements-light.txt +++ b/requirements-light.txt @@ -3,7 +3,7 @@ psycopg2-binary sqlalchemy asyncua redis -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.10.4 +git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 prometheus-client botocore boto3 diff --git a/requirements.txt b/requirements.txt index 6c5c242..d05e447 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,8 +3,8 @@ psycopg2-binary sqlalchemy asyncua redis -git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.10.4 -git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.41.0 +git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 +git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.1 prometheus-client botocore boto3 diff --git a/tests/laborious/activities/test_activities.py b/tests/laborious/activities/test_activities.py index 8e6cd78..2b1c03f 100644 --- a/tests/laborious/activities/test_activities.py +++ b/tests/laborious/activities/test_activities.py @@ -48,7 +48,8 @@ def test___init__( 'secure': False, } - mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'} + mlflow_repository = MagicMock() + plugin_store = MagicMock() opc_config = { 'bootstrap_servers': 'localhost:9092', @@ -67,12 +68,13 @@ def test___init__( activities = Activities( postgres_config=postgres_config, - mlflow_config=mlflow_config, + plugin_store=plugin_store, minio_config=minio_config, opc_config=opc_config, pi_web_api_config=pi_web_api_config, logger=logger, notification_handler=notification_handler, + mlflow_repository=mlflow_repository, ) assert isinstance(activities, Activities) @@ -101,10 +103,8 @@ def test___init__( mock_mlflow_init.assert_called_once_with( ANY, - mlflow_host=mlflow_config['host'], - mlflow_port=mlflow_config['port'], - mlflow_username=mlflow_config['username'], - mlflow_password=mlflow_config['password'], + mlflow_repository=mlflow_repository, + plugin_store=plugin_store, minio_repository=mock_minio_repository.return_value, logger=logger, notification_handler=notification_handler, @@ -193,7 +193,8 @@ async def test_shutdown( 'secure': False, } - mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'} + mlflow_repository = MagicMock() + plugin_store = MagicMock() opc_config = { 'bootstrap_servers': 'localhost:9092', @@ -212,12 +213,13 @@ async def test_shutdown( activities = Activities( postgres_config=postgres_config, - mlflow_config=mlflow_config, + plugin_store=plugin_store, minio_config=minio_config, opc_config=opc_config, pi_web_api_config=pi_web_api_config, logger=logger, notification_handler=notification_handler, + mlflow_repository=mlflow_repository, ) await activities.shutdown() diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index 5d242e5..18121b5 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -1,9 +1,10 @@ -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, AsyncMock, MagicMock, patch import numpy as np +import pandas as pd from pytest import fixture, mark, raises from sientia_do.notifications.models import NotificationLevel -from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ +from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ from laborious.activities.mlflow import MLFlow @@ -16,12 +17,14 @@ def _passthrough_from_dict(): yield -@patch('laborious.activities.mlflow.MLFlowRepository') @patch('laborious.activities.mlflow.MinioRepository') -def test___init__(mock_minio_repository, mock_mlflow_repository): +def test___init__(mock_minio_repository): logger = MagicMock() notification_handler = MagicMock() metrics_controller = AsyncMock() + mlflow_repo = MagicMock() + plugin_store = MagicMock() + minio_repo = mock_minio_repository( endpoint='localhost:9000', access_key='minio', @@ -32,24 +35,16 @@ def test___init__(mock_minio_repository, mock_mlflow_repository): bucket='test', ) mlflow = MLFlow( - mlflow_host='http://localhost', - mlflow_port=5000, - mlflow_username='admin', - mlflow_password='admin', + mlflow_repository=mlflow_repo, + plugin_store=plugin_store, minio_repository=minio_repo, logger=logger, notification_handler=notification_handler, metrics_controller=metrics_controller, ) - assert mlflow.mlflow_host == 'http://localhost' - assert mlflow.mlflow_port == 5000 - assert mlflow.mlflow_username == 'admin' - assert mlflow.mlflow_password == 'admin' - - mock_mlflow_repository.assert_called_once_with( - 'http://localhost:5000', 'admin', 'admin', ANY, ANY, ANY - ) + assert mlflow.mlflow_repository is mlflow_repo + assert mlflow.plugin_store is plugin_store mock_minio_repository.assert_called_once_with( endpoint='localhost:9000', @@ -63,12 +58,14 @@ def test___init__(mock_minio_repository, mock_mlflow_repository): @fixture -@patch('laborious.activities.mlflow.MLFlowRepository') @patch('laborious.activities.mlflow.MinioRepository') -def mlflow(mock_minio_repository, mock_mlflow_repository): +def mlflow(mock_minio_repository): logger = MagicMock() notification_handler = MagicMock() metrics_controller = AsyncMock() + mlflow_repo = MagicMock() + plugin_store = MagicMock() + minio_repo = mock_minio_repository( endpoint='localhost:9000', access_key='minio', @@ -79,17 +76,14 @@ def mlflow(mock_minio_repository, mock_mlflow_repository): bucket='test', ) mlflow = MLFlow( - mlflow_host='http://localhost:5000', - mlflow_port=5000, - mlflow_username='admin', - mlflow_password='admin', + mlflow_repository=mlflow_repo, + plugin_store=plugin_store, minio_repository=minio_repo, logger=logger, notification_handler=notification_handler, metrics_controller=metrics_controller, ) - mlflow.model_monitoring_repository = AsyncMock() mlflow.minio_repository = AsyncMock() mlflow.send_notification = MagicMock() @@ -120,9 +114,32 @@ metadata = { new_callable=AsyncMock, ) async def test_request_transform_success(mock_from_dataframe, mlflow): - data_mock = MagicMock() + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw = pd.DataFrame( + { + 'variable': ['v1', 'v1'], + 'timestamp': [ts, ts], + 'value': [1.0, 2.0], + 'created_at': [ts, ts], + } + ) + pivoted = raw.sort_values('created_at', ascending=False).drop_duplicates( + subset=['variable', 'timestamp'], keep='first' + ) + pivoted = pivoted.pivot(index='timestamp', columns='variable', values='value') + pivoted = pivoted.fillna(np.nan) + pivoted.columns.name = None + pivoted.index.name = None + pivoted['timestamp'] = pivoted.index + + out_idx = pd.Index([ts.strftime(DATETIME_FORMAT_WITH_TZ)], name=None) + out_df = pd.DataFrame({'v1': [1.0]}, index=out_idx) + wrapper = MagicMock() + wrapper.transform.return_value = (out_df, {'meta': True}) + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + payload = AsyncMock() - payload.retrieve = AsyncMock(return_value=data_mock) + payload.retrieve = AsyncMock(return_value=raw) input_data = { **metadata, @@ -131,17 +148,13 @@ async def test_request_transform_success(mock_from_dataframe, mlflow): 'model_config': {}, } - transform_response = {'success': True, 'content': MagicMock()} - mlflow.model_monitoring_repository.transform.return_value = transform_response - - data_mock.sort_values.return_value = data_mock - data_mock.drop_duplicates.return_value = data_mock - data_mock.pivot.return_value = data_mock - response_data = await mlflow.request_transform(input_data) - mlflow.model_monitoring_repository.transform.assert_called_once_with( - 'test_model', data_mock, {}, metadata['metadata'] + mlflow.mlflow_repository.get_cached_model.assert_called_once_with( + model_name='test_model', + alias='production', + retention_minutes=0, + metadata=metadata['metadata'], ) mock_from_dataframe.assert_called_once() assert response_data == mock_from_dataframe.return_value @@ -153,6 +166,8 @@ async def test_request_transform_success(mock_from_dataframe, mlflow): new_callable=AsyncMock, ) async 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) @@ -164,26 +179,22 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow): 'model_config': {}, } - transform_response = {'success': False, 'message': 'Transform failed'} - mlflow.model_monitoring_repository.transform.return_value = transform_response - data_mock.sort_values.return_value = data_mock data_mock.drop_duplicates.return_value = data_mock data_mock.pivot.return_value = data_mock - response_data = await mlflow.request_transform(input_data) + await mlflow.request_transform(input_data) mock_from_dataframe.assert_called_once_with( dataframe=None, minio_repo=mlflow.minio_repository, model_name='test_model', operation='transform', - status=transform_response, + status={'success': False, 'content': ANY}, workflow_metadata=metadata['metadata'], last_timestamp=payload.last_timestamp, logger=mlflow.logger, ) - assert response_data == mock_from_dataframe.return_value @mark.asyncio @@ -193,7 +204,13 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow): ) @patch('laborious.activities.mlflow.to_datetime') async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): + wrapper = MagicMock() + pred_df = MagicMock() + wrapper.predict.return_value = (pred_df, {}) + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + data_mock = MagicMock() + data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')]) payload = AsyncMock() payload.retrieve = AsyncMock(return_value=data_mock) @@ -204,20 +221,14 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): 'model_config': {}, } - predict_response = {'success': True, 'content': MagicMock()} - mlflow.model_monitoring_repository.predict.return_value = predict_response + pred_df.columns = MagicMock() + pred_df.__setitem__ = MagicMock() response_data = await mlflow.request_predict(input_data) data_mock.replace.assert_called_once_with(np.nan, None, inplace=True) - mock_to_datetime.assert_called_once_with( - data_mock.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ - ) - mock_to_datetime.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT) - - mlflow.model_monitoring_repository.predict.assert_called_once_with( - 'test_model', data_mock, {}, metadata['metadata'] - ) + mock_to_datetime.assert_called() + mlflow.mlflow_repository.get_cached_model.assert_called_once() mock_from_dataframe.assert_called_once() assert response_data == mock_from_dataframe.return_value @@ -229,6 +240,8 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow): ) @patch('laborious.activities.mlflow.to_datetime') async 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) @@ -240,34 +253,51 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml 'model_config': {}, } - predict_response = {'success': False, 'message': 'Predict failed'} - mlflow.model_monitoring_repository.predict.return_value = predict_response - - response_data = await mlflow.request_predict(input_data) + await mlflow.request_predict(input_data) mock_from_dataframe.assert_called_once_with( dataframe=None, minio_repo=mlflow.minio_repository, model_name='test_model', operation='predict', - status=predict_response, + status={'success': False, 'content': ANY}, workflow_metadata=metadata['metadata'], last_timestamp=payload.last_timestamp, logger=mlflow.logger, ) - assert response_data == mock_from_dataframe.return_value @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(mock_to_datetime, mlflow): - mlflow.model_monitoring_repository.retrain_model.return_value = { - 'success': True, - 'experiment': 'test_experiment', - 'message': 'Model retrained successfully.', - } +async def test_retrain_model_success_data_success_retrain( + mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow +): + mock_mkdtemp.return_value = '/tmp/x' + + mv_alias = MagicMock() + mv_alias.run_id = 'source-run' + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias + + wrapper = MagicMock() + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + + mock_cm = MagicMock() + mock_cm.__enter__.return_value = MagicMock(run_id='new-run', experiment_id='exp-1') + mock_cm.__exit__.return_value = False + mlflow.mlflow_repository.start_run.return_value = mock_cm + + ts = pd.Timestamp('2020-01-01', tz='UTC') + raw_data = pd.DataFrame( + { + 'variable': ['target', 'f1'], + 'timestamp': [ts, ts], + 'value': [1.0, 2.0], + } + ) - raw_data = MagicMock(columns=['variable', 'timestamp', 'value']) payload = AsyncMock() payload.retrieve = AsyncMock(return_value=raw_data) @@ -278,177 +308,93 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl 'model_name': 'test_model', 'model_config': { 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', }, } ) - timestamp = raw_data.__getitem__.return_value.max.return_value - - raw_data.sort_values.assert_not_called() - raw_data.drop_duplicates.assert_called_once_with(subset=['variable', 'timestamp'], keep='first') - raw_data = raw_data.drop_duplicates.return_value - - raw_data.drop.assert_has_calls( - [ - call(columns=['model_id'], inplace=True, errors='ignore'), - call(columns=['created_at'], inplace=True, errors='ignore'), - ] - ) - raw_data.pivot.assert_called_once_with(index='timestamp', columns='variable', values='value') - raw_data.pivot.return_value.fillna.assert_called_once_with(np.nan, inplace=True) - - raw_data = raw_data.pivot.return_value - - raw_data.__setitem__.assert_has_calls( - [ - call('timestamp', raw_data.index), - call('timestamp', mock_to_datetime.return_value.dt.strftime.return_value), - call('timestamp', mock_to_datetime.return_value), - ] - ) - mock_to_datetime.assert_has_calls( - [call(raw_data.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ)] - ) - mock_to_datetime.assert_has_calls( - [call(raw_data.__getitem__.return_value, format=DATETIME_FORMAT)] - ) - - mlflow.model_monitoring_repository.retrain_model.assert_called_once_with( - data=raw_data, - model_name='test_model', - model_config={ - 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - }, - metadata=metadata['metadata'], - ) - - assert response == { - 'success': True, - 'experiment': 'test_experiment', - 'message': 'Model retrained successfully.', - 'timestamp': timestamp, - } + wrapper.retrain.assert_called_once() + wrapper.store_model.assert_called_once_with(name='test_model') + assert response['success'] is True + 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(mock_to_datetime, mlflow): - mlflow.model_monitoring_repository.retrain_model.return_value = { - 'success': True, - 'experiment': 'test_experiment', - 'message': 'Model retrained successfully.', - } +async def test_retrain_model_success_with_payload_data( + mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow +): + mock_mkdtemp.return_value = '/tmp/x' + mv_alias = MagicMock(run_id='src') + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias + wrapper = MagicMock() + mlflow.mlflow_repository.get_cached_model.return_value = wrapper + mock_cm = MagicMock() + mock_cm.__enter__.return_value = MagicMock(run_id='r', experiment_id='e') + mock_cm.__exit__.return_value = False + mlflow.mlflow_repository.start_run.return_value = mock_cm 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) + pivoted = MagicMock() + raw_data.sort_values.return_value = raw_data + raw_data.drop_duplicates.return_value = raw_data + raw_data.pivot.return_value = pivoted + pivoted.fillna = MagicMock() + pivoted.columns.name = None + pivoted.index = MagicMock() + pivoted.__setitem__ = MagicMock() + response = await mlflow.retrain_model( { **metadata, 'data': payload, 'model_name': 'test_model', - 'model_config': { - 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - }, + 'model_config': {'target': 'target'}, } ) assert response['success'] is True - mlflow.minio_repository.download_file.assert_not_called() @mark.asyncio @patch('laborious.activities.mlflow.to_datetime') async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow): - mlflow.model_monitoring_repository.retrain_model.return_value = { - 'success': False, - 'traceback': 'test_traceback', - 'message': 'Model retrained failed.', - } + 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) + pivoted = MagicMock() + raw_data.sort_values.return_value = raw_data + raw_data.drop_duplicates.return_value = raw_data + raw_data.pivot.return_value = pivoted + pivoted.fillna = MagicMock() + pivoted.columns.name = None + pivoted.index = MagicMock() + pivoted.__setitem__ = MagicMock() + response = await mlflow.retrain_model( { **metadata, 'data': payload, 'model_name': 'test_model', - 'model_config': { - 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - }, + 'model_config': {'target': 'target'}, } ) - timestamp = raw_data.__getitem__.return_value.max.return_value - - raw_data.sort_values.assert_called_once_with('created_at', ascending=False) - raw_data.sort_values.return_value.drop_duplicates.assert_called_once_with( - subset=['variable', 'timestamp'], keep='first' - ) - raw_data = raw_data.sort_values.return_value.drop_duplicates.return_value - - raw_data.drop.assert_has_calls( - [ - call(columns=['model_id'], inplace=True, errors='ignore'), - call(columns=['created_at'], inplace=True, errors='ignore'), - ] - ) - raw_data.pivot.assert_called_once_with(index='timestamp', columns='variable', values='value') - raw_data.pivot.return_value.fillna.assert_called_once_with(np.nan, inplace=True) - - raw_data = raw_data.pivot.return_value - - raw_data.__setitem__.assert_has_calls( - [ - call('timestamp', raw_data.index), - call('timestamp', mock_to_datetime.return_value.dt.strftime.return_value), - call('timestamp', mock_to_datetime.return_value), - ] - ) - mock_to_datetime.assert_has_calls( - [call(raw_data.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ)] - ) - mock_to_datetime.assert_has_calls( - [call(raw_data.__getitem__.return_value, format=DATETIME_FORMAT)] - ) - - mlflow.model_monitoring_repository.retrain_model.assert_called_once_with( - data=raw_data, - model_name='test_model', - model_config={ - 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - }, - metadata=metadata['metadata'], - ) - - mlflow.send_notification_async.assert_called_once_with( - metadata=metadata['metadata'], - notification_id='RETRAIN_MODEL_ERROR', - message='Error retraining model test_model: Model retrained failed.', - block='retrain_model', - level=NotificationLevel.ERROR, - attachment_content=ANY, - ) - - assert response == { - 'success': False, - 'traceback': 'test_traceback', - 'message': 'Model retrained failed.', - 'timestamp': timestamp, - } + assert response['success'] is False + mlflow.send_notification_async.assert_called_once() + assert 'retrain failed' in response['message'] @mark.asyncio @@ -459,18 +405,40 @@ async def test_retrain_model_data_error(mlflow): 'model_name': 'test_model', 'model_config': { 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', }, } ) - assert response == { - 'success': False, - 'message': "Error loading retrain data: 'data'", - 'traceback': ANY, - 'timestamp': ANY, - } + assert response['success'] is False + assert 'data' in response['message'].lower() or 'loading' in response['message'].lower() + + +@mark.asyncio +async def test_retrain_model_missing_target(mlflow): + raw_data = MagicMock(columns=['variable', 'timestamp', 'value']) + raw_data.__getitem__.return_value.max.return_value = 'ts' + payload = AsyncMock() + payload.retrieve = AsyncMock(return_value=raw_data) + + pivoted = MagicMock() + raw_data.drop_duplicates.return_value = raw_data + raw_data.pivot.return_value = pivoted + pivoted.fillna = MagicMock() + pivoted.columns.name = None + pivoted.index = MagicMock() + pivoted.__setitem__ = MagicMock() + + response = await mlflow.retrain_model( + { + **metadata, + 'data': payload, + 'model_name': 'test_model', + 'model_config': {}, + } + ) + + assert response['success'] is False + assert 'target' in response['message'] @mark.asyncio @@ -485,96 +453,94 @@ async def test_retrain_model_data_error_no_minio_repository(mlflow): 'model_name': 'test_model', 'model_config': { 'target': 'target', - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', }, } ) - assert str(e.value) == 'Minio repository not initialized' + assert str(e.value) == 'Minio repository not initialized' @mark.asyncio async 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'), + ] + input_data = { **metadata, 'model_name': 'test_model', 'model_id': 1, - 'experiment': 'test', + 'experiment': {'run_id': 'run-x', 'experiment_id': 'e1'}, 'timestamp': 2, 'status': 'success', } response = await mlflow.update_production_model(input_data) - mlflow.model_monitoring_repository.update_production_model.assert_called_once_with( - experiment='test', model_name='test_model', metadata=metadata['metadata'] + mlflow.mlflow_repository.promote_to_alias.assert_called_once_with( + model_name='test_model', + version='3', + alias='production', + metadata=metadata['metadata'], ) - assert response == mlflow.model_monitoring_repository.update_production_model.return_value + assert response['model_name'] == 'test_model' + assert response['version'] == '3' @mark.asyncio async def test_update_production_model_error(mlflow): - mlflow.model_monitoring_repository.update_production_model.side_effect = Exception( - 'Error updating production model' - ) + mlflow.mlflow_repository._client.search_model_versions.return_value = [] input_data = { **metadata, 'model_name': 'test_model', 'model_id': 1, - 'experiment': 'test', + 'experiment': {'run_id': 'run-x', 'experiment_id': 'e1'}, 'timestamp': 2, 'status': 'success', } try: await mlflow.update_production_model(input_data) - except Exception as e: - assert str(e) == 'Error updating production model' + except Exception: mlflow.send_notification_async.assert_called_once_with( metadata=metadata['metadata'], notification_id='UPDATE_PRODUCTION_MODEL_ERROR', - message='Error updating production model test_model: Error updating production model', + message=ANY, block='update_production_model', level=NotificationLevel.ERROR, attachment_content=ANY, ) else: - raise AssertionError('No exception raised') + raise AssertionError('Expected exception') @mark.asyncio @patch('laborious.activities.mlflow.to_datetime') async def test_get_reference_data_success(mock_to_datetime, mlflow): - # Arrange input_data = { **metadata, 'model_name': 'test_model', } - # Mock reference data DataFrame + mv = MagicMock(run_id='run1') + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv + mock_reference_data = MagicMock() mock_reference_data.__getitem__.return_value = MagicMock() mock_to_datetime.return_value.dt.strftime.return_value = MagicMock() mock_reference_data.to_dict.return_value = [ {'timestamp': '2023-05-26 11:12:27', 'value': 1.0}, - {'timestamp': '2023-05-26 11:12:28', 'value': 2.0}, ] - mlflow.model_monitoring_repository.load_artifact_dataframe.return_value = mock_reference_data - - # Act - result = await mlflow.get_reference_data(input_data) - - # Assert - mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with( - model_name='test_model', - artifact_path='evaluation_data.csv', - metadata=metadata['metadata'], - ) - mock_to_datetime.assert_called_once_with(mock_reference_data.__getitem__.return_value) + with patch('laborious.activities.mlflow.pd.read_csv', return_value=mock_reference_data): + with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='/t'): + 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) mock_reference_data.to_dict.assert_called_once_with(orient='records') assert result == mock_reference_data.to_dict.return_value @@ -582,48 +548,30 @@ async def test_get_reference_data_success(mock_to_datetime, mlflow): @mark.asyncio async def test_get_reference_data_not_found(mlflow): - # Arrange input_data = { **metadata, 'model_name': 'test_model', } - mlflow.model_monitoring_repository.load_artifact_dataframe.return_value = None + mlflow.mlflow_repository._client.get_model_version_by_alias.side_effect = Exception('missing') - # Act result = await mlflow.get_reference_data(input_data) - # Assert - mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with( - model_name='test_model', - artifact_path='evaluation_data.csv', - metadata=metadata['metadata'], - ) - mlflow.warning.assert_called_once_with( - 'Reference data not found for model test_model', metadata['metadata'] - ) + mlflow.warning.assert_called() assert result is None @mark.asyncio async def test_get_reference_data_exception(mlflow): - # Arrange input_data = { **metadata, 'model_name': 'test_model', } - mlflow.model_monitoring_repository.load_artifact_dataframe.side_effect = Exception( - 'Error loading artifact' - ) + mv = MagicMock(run_id='run1') + mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv + mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail') - # Act & Assert - with raises(Exception) as e: - await mlflow.get_reference_data(input_data) + result = await mlflow.get_reference_data(input_data) - assert str(e.value) == 'Error loading artifact' - mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with( - model_name='test_model', - artifact_path='evaluation_data.csv', - metadata=metadata['metadata'], - ) + assert result is None diff --git a/tests/laborious/utils/repository/test_model_repository.py b/tests/laborious/utils/repository/test_model_repository.py deleted file mode 100644 index 15f8c6b..0000000 --- a/tests/laborious/utils/repository/test_model_repository.py +++ /dev/null @@ -1,1680 +0,0 @@ -from datetime import UTC, datetime -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch - -import mlflow as mlflow_lib -import numpy as np -import pytest -from pandas import DataFrame, Timestamp - -from laborious import metrics -from laborious.utils.repository.model_repository import MLFlowRepository, force_memory_release - - -@patch('laborious.utils.repository.model_repository.ctypes') -@patch('laborious.utils.repository.model_repository.gc') -def test_force_memory_release_success(gc, ctypes): - logger = MagicMock() - - force_memory_release(logger) - - gc.collect.assert_called_once() - ctypes.CDLL.return_value.malloc_trim.assert_called_once_with(0) - logger.info.assert_called_once_with('Memory released') - - -@patch('laborious.utils.repository.model_repository.ctypes') -@patch('laborious.utils.repository.model_repository.gc') -def test_force_memory_release_error(gc, ctypes): - logger = MagicMock() - ctypes.CDLL.return_value.malloc_trim.side_effect = Exception('error') - force_memory_release(logger) - - gc.collect.assert_called_once() - ctypes.CDLL.return_value.malloc_trim.assert_called_once_with(0) - - logger.info.assert_called_once_with('Memory release failed: error') - - -@pytest.fixture -def mlflow_repository(): - with patch('laborious.utils.repository.model_repository.mlflow'): - repo = MLFlowRepository( - host='http://localhost:5000', - username='admin', - password='admin', - logger=MagicMock(), - notification_handler=MagicMock(), - metrics_controller=AsyncMock(), - ) - repo.emit_metric = AsyncMock() - repo.observe_lag = AsyncMock() - repo.send_notification = MagicMock() - repo.send_notification_async = AsyncMock() - return repo - - -@pytest.fixture -def mlflow(): - with patch('laborious.utils.repository.model_repository.mlflow') as mlflow: - yield mlflow - - -metadata = { - 'metadata': { - 'model_id': 'test_model', - 'model_name': 'test_model', - 'workflow_name': 'test_workflow', - 'schema_name': 'test_schedule', - }, -} - - -class Any: - pass - - -def test_get_model_uri_prediction(mlflow, mlflow_repository): - mlflow.get_run.return_value = MagicMock(info=MagicMock(artifact_uri='test')) - output = mlflow_repository.get_model_uri('0', prediction=True) - assert output == 'test/prediction_model' - mlflow.get_run.assert_called_once_with('0') - - -def test_get_model_uri_transform(mlflow, mlflow_repository): - mlflow.get_run.return_value = MagicMock(info=MagicMock(artifact_uri='test')) - output = mlflow_repository.get_model_uri('0', prediction=False) - assert output == 'test/data_model' - mlflow.get_run.assert_called_once_with('0') - - -def test_get_model_run_id_not_registered_models(mlflow_repository): - mlflow_repository.client.search_registered_models.return_value = [] - with pytest.raises(mlflow_lib.exceptions.MlflowException) as e: - mlflow_repository.get_model_run_id('test') - - mlflow_repository.client.search_registered_models.assert_called_once_with( - filter_string="name='test'" - ) - - assert str(e.value) == "Model 'test' not found in the Model Registry." - - -def test_get_model_run_id_not_stage_versions(mlflow_repository): - mlflow_repository.client.search_registered_models.return_value = [MagicMock(name='test')] - - mlflow_repository.client.search_model_versions.return_value = [ - MagicMock(current_stage='Staging'), - MagicMock(current_stage='Staging'), - MagicMock(current_stage='Archived'), - ] - - with pytest.raises(mlflow_lib.exceptions.MlflowException) as e: - mlflow_repository.get_model_run_id('test') - - mlflow_repository.client.search_registered_models.assert_called_once_with( - filter_string="name='test'" - ) - mlflow_repository.client.search_model_versions.assert_called_once_with( - filter_string="name='test'" - ) - - assert str(e.value) == "Model 'test' in stage 'Production' not found in the Model Registry." - - -def test_get_model_run_id_success(mlflow_repository): - mlflow_repository.client.search_registered_models.return_value = [MagicMock(name='test')] - - mlflow_repository.client.search_model_versions.return_value = [ - MagicMock(current_stage='Production', version='1'), - MagicMock(current_stage='Production', version='2', source='runs/test/1'), - MagicMock(current_stage='Archived', version='3'), - ] - - output = mlflow_repository.get_model_run_id('test') - - mlflow_repository.client.search_registered_models.assert_called_once_with( - filter_string="name='test'" - ) - mlflow_repository.client.search_model_versions.assert_called_once_with( - filter_string="name='test'" - ) - - assert output == '1' - - -def test_get_next_run_name(mlflow, mlflow_repository): - mlflow.search_runs.return_value = [1, 2, 3] - output = mlflow_repository.get_next_run_name('run') - assert output == 'run-4' - mlflow.search_runs.assert_called_once_with( - experiment_names=['run'], - order_by=['start_time desc'], - ) - - -def test_get_experiment_experiment_exists(mlflow, mlflow_repository): - experiment = MagicMock(experiment_id='0') - - mlflow.get_experiment_by_name.return_value = experiment - - output = mlflow_repository.get_experiment('test') - - assert output == experiment - - -def test_get_experiment_none_create(mlflow, mlflow_repository): - experiment = MagicMock(experiment_id='0') - - mlflow.get_experiment_by_name.return_value = None - - mlflow.get_experiment.return_value = experiment - - output = mlflow_repository.get_experiment('test', create_if_not_exists=True) - - mlflow.create_experiment.assert_called_once_with('test') - mlflow.get_experiment.assert_called_once_with(mlflow.create_experiment.return_value) - - assert output == experiment - - -def test_get_experiment_none_not_create(mlflow, mlflow_repository): - mlflow.get_experiment_by_name.return_value = None - - with pytest.raises(ValueError) as e: - mlflow_repository.get_experiment('test', create_if_not_exists=False) - - assert str(e) == 'Experiment test not found' - - -def test_get_model_params(mlflow, mlflow_repository): - output = mlflow_repository.get_model_params('test') - assert output == mlflow.get_run.return_value.data.params - - -def test_check_artifact_exists_true(mlflow_repository): - artifact = MagicMock(path='test_artifact') - mlflow_repository.client.list_artifacts.return_value = [artifact] - - result = mlflow_repository.check_artifact_exists( - 'run_id', 'test_artifact', metadata['metadata'] - ) - - assert result is True - mlflow_repository.client.list_artifacts.assert_called_once_with('run_id') - - -def test_check_artifact_exists_false(mlflow_repository): - artifact = MagicMock(path='other_artifact') - mlflow_repository.client.list_artifacts.return_value = [artifact] - - result = mlflow_repository.check_artifact_exists( - 'run_id', 'test_artifact', metadata['metadata'] - ) - - assert result is False - mlflow_repository.client.list_artifacts.assert_called_once_with('run_id') - - -@pytest.mark.asyncio -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@patch('laborious.utils.repository.model_repository.makedirs') -async def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='test') - - path.exists.return_value = True - - output = await mlflow_repository.dowload_artifacts('test', {}, 'path') - - mlflow_repository.get_model_run_id.assert_called_once_with( - model_name='test', stage='Production' - ) - - path.join.assert_called_once_with('./tmp/artifacts/test', 'path') - - path.exists.assert_called_once_with(path.join.return_value) - - rmtree.assert_called_once_with(path.join.return_value) - - makedirs.assert_called_once_with('./tmp/artifacts/test', exist_ok=True) - - mlflow_repository.client.download_artifacts.assert_called_once_with( - mlflow_repository.get_model_run_id.return_value, 'path', './tmp/artifacts/test' - ) - - assert output == mlflow_repository.client.download_artifacts.return_value - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@patch('laborious.utils.repository.model_repository.makedirs') -async def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='test') - - path.exists.return_value = False - - output = await mlflow_repository.dowload_artifacts('test', {}, 'path') - - mlflow_repository.get_model_run_id.assert_called_once_with( - model_name='test', stage='Production' - ) - - path.join.assert_called_once_with('./tmp/artifacts/test', 'path') - - path.exists.assert_called_once_with(path.join.return_value) - - rmtree.assert_not_called() - - makedirs.assert_called_once_with('./tmp/artifacts/test', exist_ok=True) - - mlflow_repository.client.download_artifacts.assert_called_once_with( - mlflow_repository.get_model_run_id.return_value, 'path', './tmp/artifacts/test' - ) - - assert output == mlflow_repository.client.download_artifacts.return_value - - -@pytest.mark.asyncio -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@patch('laborious.utils.repository.model_repository.makedirs') -async def test_download_artifacts_error(makedirs, rmtree, path, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='test') - - path.exists.return_value = True - - mlflow_repository.client.download_artifacts.side_effect = ValueError('test') - - with pytest.raises(ValueError): - await mlflow_repository.dowload_artifacts('test', {}, 'path') - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -@patch('laborious.utils.repository.model_repository.mlflow') -@patch('laborious.utils.repository.model_repository.pd') -@patch('laborious.utils.repository.model_repository.StringIO') -async def test_load_artifact_dataframe_success(_stringio, pd, mlflow, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='run_id') - mlflow_repository.check_artifact_exists = MagicMock(return_value=True) - - mlflow.artifacts.load_text.return_value = 'col1,col2\n1,2\n3,4' - - result = await mlflow_repository.load_artifact_dataframe( - 'model_name', 'artifact_path', metadata['metadata'] - ) - - mlflow_repository.get_model_run_id.assert_called_once_with( - model_name='model_name', stage='Production' - ) - mlflow_repository.check_artifact_exists.assert_called_once_with( - 'run_id', 'artifact_path', metadata['metadata'] - ) - mlflow.artifacts.load_text.assert_called_once_with('runs:/run_id/artifact_path') - assert result == pd.read_csv.return_value - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_artifact_dataframe_not_exists(mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='run_id') - mlflow_repository.check_artifact_exists = MagicMock(return_value=False) - - result = await mlflow_repository.load_artifact_dataframe( - 'model_name', 'artifact_path', metadata['metadata'] - ) - - assert result is None - mlflow_repository.get_model_run_id.assert_called_once_with( - model_name='model_name', stage='Production' - ) - mlflow_repository.check_artifact_exists.assert_called_once_with( - 'run_id', 'artifact_path', metadata['metadata'] - ) - - -@pytest.mark.asyncio -@patch('laborious.utils.repository.model_repository.mlflow') -async def test_load_artifact_dataframe_error(mlflow, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock(return_value='run_id') - mlflow_repository.check_artifact_exists = MagicMock(return_value=True) - mlflow.artifacts.load_text.side_effect = ValueError('error') - - with pytest.raises(ValueError): - await mlflow_repository.load_artifact_dataframe( - 'model_name', 'artifact_path', metadata['metadata'] - ) - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -def test_get_experiment_error(mlflow, mlflow_repository): - mlflow.get_experiment_by_name.return_value = None - - try: - mlflow_repository.get_experiment('test') - except ValueError as e: - assert str(e) == 'Experiment test not found' - else: - raise AssertionError('Expected ValueError') - - -def test_get_experiment_create_error(mlflow, mlflow_repository): - mlflow.get_experiment_by_name.return_value = None - mlflow.create_experiment.return_value = None - mlflow.get_experiment.return_value = None - with pytest.raises(ValueError) as e: - mlflow_repository.get_experiment('test', create_if_not_exists=True) - - assert str(e) == 'Experiment test not found after creation, unknown reason' - - -@pytest.mark.asyncio -async def test_load_predict_model_sklearn(mlflow, mlflow_repository): - result = await mlflow_repository.load_predict_model('test_model', {}, 'sklearn') - - assert result == mlflow.sklearn.load_model.return_value - mlflow.sklearn.load_model.assert_called_once_with('models:/test_model/production') - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_predict_model_pyfunc(mlflow, mlflow_repository): - result = await mlflow_repository.load_predict_model('test_model', {}, 'pyfunc') - assert result == mlflow.pyfunc.load_model.return_value - mlflow.pyfunc.load_model.assert_called_once_with('models:/test_model/production') - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_predict_model_pytorch(mlflow, mlflow_repository): - result = await mlflow_repository.load_predict_model('test_model', {}, 'pytorch') - assert result == mlflow.pytorch.load_model.return_value - mlflow.pytorch.load_model.assert_called_once_with('models:/test_model/production') - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_predict_model_error(mlflow_repository): - with pytest.raises(ValueError) as e: - await mlflow_repository.load_predict_model('test_model', {}, 'invalid') - assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'." - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -def validate_common_load_transform_model_mocks(mlflow_repository, model_name): - mlflow_repository.get_model_run_id.assert_called_once_with( - model_name=model_name, stage='Production' - ) - mlflow_repository.get_model_uri.assert_called_once_with( - mlflow_repository.get_model_run_id.return_value, prediction=False - ) - - -@pytest.mark.asyncio -async def test_load_transform_model_sklearn(mlflow, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock() - mlflow_repository.get_model_uri = MagicMock() - - result = await mlflow_repository.load_transform_model('test_model', {}, 'sklearn') - - validate_common_load_transform_model_mocks(mlflow_repository, 'test_model') - - assert result == mlflow.sklearn.load_model.return_value - mlflow.sklearn.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value) - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_transform_model_pyfunc(mlflow, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock() - mlflow_repository.get_model_uri = MagicMock() - - result = await mlflow_repository.load_transform_model('test_model', {}, 'pyfunc') - - validate_common_load_transform_model_mocks(mlflow_repository, 'test_model') - - assert result == mlflow.pyfunc.load_model.return_value - mlflow.pyfunc.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value) - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_transform_model_pytorch(mlflow, mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock() - mlflow_repository.get_model_uri = MagicMock() - - result = await mlflow_repository.load_transform_model('test_model', {}, 'pytorch') - validate_common_load_transform_model_mocks(mlflow_repository, 'test_model') - - assert result == mlflow.pytorch.load_model.return_value - mlflow.pytorch.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value) - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_load_transform_model_error(mlflow_repository): - mlflow_repository.get_model_run_id = MagicMock() - mlflow_repository.get_model_uri = MagicMock() - - with pytest.raises(ValueError) as e: - await mlflow_repository.load_transform_model('test_model', {}, 'invalid') - assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'." - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -async def test_download_model_invalid_model_type(mlflow_repository): - with pytest.raises(ValueError) as e: - await mlflow_repository.download_model('test_model', {}, 'invalid', 'sklearn') - assert str(e) == "Invalid model_type. Use 'predict' or 'transform'." - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - 'model_type', [('predict', 'prediction_model'), ('transform', 'data_model')] -) -async def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type): - mlflow_repository.dowload_artifacts = AsyncMock() - - result = await mlflow_repository.download_model('test_model', {}, model_type[0], 'pyfunc', True) - - mlflow_repository.dowload_artifacts.assert_called_once_with('test_model', {}, model_type[1]) - - mlflow.pyfunc.load_model.assert_called_once_with( - mlflow_repository.dowload_artifacts.return_value - ) - - assert result == ( - mlflow.pyfunc.load_model.return_value._model_impl.python_model, - mlflow_repository.dowload_artifacts.return_value, - ) - - -@pytest.mark.asyncio -async def test_download_model_predict(mlflow_repository): - mlflow_repository.load_predict_model = AsyncMock() - mlflow_repository.load_transform_model = AsyncMock() - - result = await mlflow_repository.download_model('test_model', {}, 'predict', 'pyfunc', False) - - mlflow_repository.load_predict_model.assert_called_once_with('test_model', {}, 'pyfunc') - mlflow_repository.load_transform_model.assert_not_called() - - assert result == (mlflow_repository.load_predict_model.return_value, None) - - -@pytest.mark.asyncio -async def test_download_model_transform(mlflow_repository): - mlflow_repository.load_predict_model = AsyncMock() - mlflow_repository.load_transform_model = AsyncMock() - - result = await mlflow_repository.download_model('test_model', {}, 'transform', 'pyfunc', False) - - mlflow_repository.load_predict_model.assert_not_called() - mlflow_repository.load_transform_model.assert_called_once_with('test_model', {}, 'pyfunc') - - assert result == (mlflow_repository.load_transform_model.return_value, None) - - -def test_detect_and_parse_datetime_index_empty(mlflow_repository): - input_data = DataFrame() - response = mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata']) - assert response.equals(input_data) - - -invalid_cases = [ - ( - {'value': {'2024-01-01 12:00:00': 1, 2024: 2}}, - "Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S. Elements are , .", - ), - ( - {'value': {'2024-01-01': 1, '2024-01-02': 2}}, - 'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S. Unable to parse given date format', - ), - ( - {'value': {Any(): 1, Any(): 2}}, - 'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S. Got .', - ), -] - - -@pytest.mark.parametrize('data', invalid_cases) -def test_detect_and_parse_datetime_index_error_cases(mlflow_repository, data): - input_data = DataFrame(data[0]) - - message = data[1] - - with pytest.raises(ValueError) as e: - mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata']) - - assert str(e) == message - - -valid_cases = [ - ( - {'value': {'2024-01-01 12:00:00+0000': 1, '2024-01-02 12:00:00+0000': 2}}, - ['2024-01-01 12:00:00+0000', '2024-01-02 12:00:00+0000'], - ), - ( - { - 'value': { - datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC): 1, - datetime(2025, 1, 2, 12, 0, 0, tzinfo=UTC): 2, - } - }, - ['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000'], - ), - ( - { - 'value': { - datetime(2025, 1, 1, 12, 0, 0, tzinfo=None): 1, - datetime(2025, 1, 2, 12, 0, 0, tzinfo=None): 2, - } - }, - ['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000'], - ), - ( - { - 'value': { - Timestamp(2026, 1, 1, 12, 0, 0, tzinfo=UTC): 1, - Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=UTC): 2, - } - }, - ['2026-01-01 12:00:00+0000', '2026-01-02 12:00:00+0000'], - ), -] - - -@pytest.mark.parametrize('data,expected', valid_cases) -def test_detect_and_parse_datetime_index_valid_format(mlflow_repository, data, expected): - input_data = DataFrame(data) - - response = mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata']) - - assert response.index.tolist() == expected - - -@patch('laborious.utils.repository.model_repository.datetime') -def test_check_cache_retention_false(datetime_mock, mlflow_repository): - datetime_mock.now = MagicMock(return_value=datetime.strptime('2025-01-02', '%Y-%m-%d')) - cache = {'timestamp': datetime.strptime('2025-01-01', '%Y-%m-%d')} - - assert mlflow_repository.check_cache_retention(cache, 1) is False - - -@patch('laborious.utils.repository.model_repository.datetime') -def test_check_cache_retention_true(datetime_mock, mlflow_repository): - datetime_mock.now = MagicMock(return_value=datetime.strptime('2025-01-01', '%Y-%m-%d')) - cache = {'timestamp': datetime.strptime('2025-01-01', '%Y-%m-%d')} - assert mlflow_repository.check_cache_retention(cache, 1) is True - - -def test_handle_valid_model(mlflow_repository): - cache = {'target': {'model': 'model', 'artifact_path': 'test_artifact_path'}} - output = mlflow_repository.handle_valid_model('model_name', cache) - assert output == {'model': 'model', 'artifact_path': 'test_artifact_path'} - - -def test_handle_outdated_model(mlflow_repository): - mlflow_repository.model_cache = { - 'model_name_transform': { - 'target': {'model': 'model', 'artifact_path': 'test_artifact_path'} - } - } - mlflow_repository.handle_outdated_model('model_name', 'model_name_transform') - assert mlflow_repository.model_cache == {} - - -@pytest.mark.asyncio -async def test_get_model_retention_0(mlflow_repository): - model = MagicMock() - - mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path')) - output = await mlflow_repository.get_model('model_name', {}, 0, 'predict', 'pyfunc') - - assert output == model - mlflow_repository.download_model.assert_called_once_with( - model_name='model_name', - metadata={}, - model_type='predict', - flavor='pyfunc', - load_wrapper=False, - ) - - -@pytest.mark.asyncio -async def test_get_model_cached_valid(mlflow_repository): - mlflow_repository.check_cache_retention = MagicMock(return_value=True) - mlflow_repository.handle_valid_model = MagicMock() - mlflow_repository.handle_outdated_model = MagicMock() - mlflow_repository.model_cache = { - 'model_name_predict': { - 'target': 'cached_model', - } - } - - output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc') - - assert output == mlflow_repository.handle_valid_model.return_value - mlflow_repository.check_cache_retention.assert_called_once_with( - mlflow_repository.model_cache['model_name_predict'], 1 - ) - - mlflow_repository.handle_valid_model.assert_called_once_with( - model_name='model_name', cache=mlflow_repository.model_cache['model_name_predict'] - ) - - mlflow_repository.handle_outdated_model.assert_not_called() - - -@pytest.mark.asyncio -async def test_get_model_cached_outdated(mlflow_repository): - mlflow_repository.check_cache_retention = MagicMock(return_value=False) - mlflow_repository.handle_valid_model = MagicMock() - mlflow_repository.handle_outdated_model = MagicMock() - model = MagicMock() - mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path')) - cache = { - 'model_name_predict': { - 'target': 'cached_model', - } - } - mlflow_repository.model_cache = cache - - output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc') - assert output == model - mlflow_repository.check_cache_retention.assert_called_once_with( - { - 'target': 'cached_model', - }, - 1, - ) - mlflow_repository.handle_valid_model.assert_not_called() - mlflow_repository.handle_outdated_model.assert_called_once_with( - model_name='model_name', model_key='model_name_predict' - ) - - -@pytest.mark.asyncio -async def test_get_model_cached_not_found(mlflow_repository): - mlflow_repository.check_cache_retention = MagicMock(return_value=False) - mlflow_repository.handle_valid_model = MagicMock() - mlflow_repository.handle_outdated_model = MagicMock() - mlflow_repository.model_cache = {} - model = MagicMock() - mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path')) - output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc') - assert output == model - mlflow_repository.check_cache_retention.assert_not_called() - mlflow_repository.handle_valid_model.assert_not_called() - mlflow_repository.handle_outdated_model.assert_not_called() - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@pytest.mark.asyncio -async def test_get_cached_operation_retention_0(force_memory_release, mlflow_repository): - model = MagicMock() - data = MagicMock() - mlflow_repository.get_model = AsyncMock(return_value=model) - output = await mlflow_repository.get_cached_operation( - 'model_name', data, 'transform', 0, 'sklearn', {} - ) - assert output == model.predict.return_value - force_memory_release.assert_called_once_with(mlflow_repository.logger) - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@pytest.mark.asyncio -async def test_get_cached_predict_retention_not_0(force_memory_release, mlflow_repository): - model = MagicMock() - data = MagicMock() - mlflow_repository.get_model = AsyncMock(return_value=model) - output = await mlflow_repository.get_cached_operation( - 'model_name', data, 'predict', 1, 'sklearn', {} - ) - assert output == model.predict.return_value - force_memory_release.assert_not_called() - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@pytest.mark.asyncio -async def test_get_cached_operation_invalid_operation(force_memory_release, mlflow_repository): - data = MagicMock() - with pytest.raises(ValueError) as e: - await mlflow_repository.get_cached_operation( - 'model_name', data, 'invalid', 0, 'sklearn', {} - ) - assert str(e) == "Invalid operation. Use 'transform' or 'predict'." - - -@patch('laborious.utils.repository.model_repository.pd.merge') -@patch('laborious.utils.repository.model_repository.isinstance') -@pytest.mark.asyncio -async def test_fit_models_not_df_target_name_none_and_not_in_model( - isinstance_mock, pd_merge, mlflow_repository -): - isinstance_mock.return_value = False - - data_model = MagicMock() - prediction_model = MagicMock() - mlflow_repository.download_model = AsyncMock( - side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')], - ) - mlflow_repository.detect_and_parse_datetime_index = MagicMock( - return_value=MagicMock(drop_duplicates=MagicMock(return_value=MagicMock(columns=[]))) - ) - mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame()) - - data = MagicMock() - - output = await mlflow_repository.fit_models( - 'model_name', - data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - False, - 'pyfunc', - None, - ) - - mlflow_repository.download_model.assert_has_calls( - [ - call( - model_name='model_name', - metadata=metadata['metadata'], - model_type='transform', - flavor='sklearn', - load_wrapper=False, - ), - call( - model_name='model_name', - metadata=metadata['metadata'], - model_type='predict', - flavor='pyfunc', - load_wrapper=True, - ), - ] - ) - - data_model.fit.assert_called_once_with(data) - - data_model.fit.return_value.predict.assert_called_once_with(data) - - transformed_data = data_model.fit.return_value.predict.return_value - - transformed_data.__setitem__.assert_called_once_with('timestamp', transformed_data.index) - - mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with( - transformed_data, metadata['metadata'] - ) - - transformed_data = mlflow_repository.detect_and_parse_datetime_index.return_value - - transformed_data.drop_duplicates.assert_called_once_with(subset=['timestamp'], keep='first') - - transformed_data = transformed_data.drop_duplicates.return_value - - data.loc.__getitem__.assert_called_once_with(transformed_data.index) - - aligned_data = data.loc.__getitem__.return_value - - aligned_data.__getitem__.assert_called_once_with(data_model.fit.return_value.target_variable) - - pd_merge.assert_called_once_with( - transformed_data, aligned_data.__getitem__.return_value, left_index=True, right_index=True - ) - - prediction_model.fit.assert_called_once_with(pd_merge.return_value) - - mlflow_repository.get_prediction_data.assert_called_once_with( - prediction_model, - pd_merge.return_value, - data_model.fit.return_value.target_variable, - 'pyfunc', - ) - - assert output == { - 'prediction_model': {'model': prediction_model, 'artifact_path': 'artifact_path'}, - 'data_model': {'model': data_model.fit.return_value, 'artifact_path': 'artifact_path'}, - 'prediction_data': mlflow_repository.get_prediction_data.return_value, - } - - -@patch('laborious.utils.repository.model_repository.pd.merge') -@patch('laborious.utils.repository.model_repository.isinstance') -@pytest.mark.asyncio -async def test_fit_models_df_target_name_not_none_and_in_model( - isinstance_mock, pd_merge, mlflow_repository -): - isinstance_mock.return_value = True - - data_model = MagicMock( - target_variable='feat_2', - ) - prediction_model = MagicMock() - mlflow_repository.download_model = AsyncMock( - side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')], - ) - mlflow_repository.detect_and_parse_datetime_index = MagicMock( - return_value=MagicMock( - drop_duplicates=MagicMock(return_value=MagicMock(columns=['feat_1'])) - ) - ) - mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame()) - - data = MagicMock() - - output = await mlflow_repository.fit_models( - 'model_name', - data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - False, - 'pyfunc', - 'feat_1', - ) - - mlflow_repository.download_model.assert_has_calls( - [ - call( - model_name='model_name', - metadata=metadata['metadata'], - model_type='transform', - flavor='sklearn', - load_wrapper=False, - ), - call( - model_name='model_name', - metadata=metadata['metadata'], - model_type='predict', - flavor='pyfunc', - load_wrapper=True, - ), - ] - ) - - data_model.fit.assert_called_once_with(data) - - transformed_data = data_model.fit.return_value - - transformed_data.__setitem__.assert_called_once_with('timestamp', transformed_data.index) - - mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with( - transformed_data, metadata['metadata'] - ) - - transformed_data = mlflow_repository.detect_and_parse_datetime_index.return_value - - transformed_data.drop_duplicates.assert_called_once_with(subset=['timestamp'], keep='first') - - transformed_data = transformed_data.drop_duplicates.return_value - - data.loc.__getitem__.assert_not_called() - - pd_merge.assert_not_called() - - prediction_model.fit.assert_called_once_with(transformed_data) - - mlflow_repository.get_prediction_data.assert_called_once_with( - prediction_model, transformed_data, 'feat_1', 'pyfunc' - ) - - assert output == { - 'prediction_model': {'model': prediction_model, 'artifact_path': 'artifact_path'}, - 'data_model': {'model': data_model, 'artifact_path': 'artifact_path'}, - 'prediction_data': mlflow_repository.get_prediction_data.return_value, - } - - -@pytest.mark.asyncio -async def test_log_model_sklearn(mlflow, mlflow_repository): - model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'} - await mlflow_repository.log_model( - model_data, 'sklearn', 'prediction_model', metadata['metadata'] - ) - mlflow.sklearn.log_model.assert_called_once_with(model_data['model'], 'prediction_model') - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY - ) - - -@patch('laborious.utils.repository.model_repository.path') -@pytest.mark.asyncio -async def test_log_model_pyfunc(path, mlflow, mlflow_repository): - model_mock = MagicMock() - model_data = {'model': model_mock, 'artifact_path': 'artifact_path'} - await mlflow_repository.log_model( - model_data, 'pyfunc', 'prediction_model', metadata['metadata'] - ) - - mlflow.pyfunc.log_model.assert_not_called() - - path.join.assert_called_once_with('artifact_path', 'code', 'utils') - - model_mock.store_model.assert_called_once_with( - artifact_path='prediction_model', code_path=[path.join.return_value], to_disk=False - ) - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_log_model_pytorch(mlflow, mlflow_repository): - model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'} - await mlflow_repository.log_model( - model_data, 'pytorch', 'prediction_model', metadata['metadata'] - ) - mlflow.pytorch.log_model.assert_called_once_with(model_data['model'], 'prediction_model') - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_log_model_error(mlflow_repository): - model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'} - with pytest.raises(ValueError) as e: - await mlflow_repository.log_model( - model_data, 'invalid', 'prediction_model', metadata['metadata'] - ) - assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'." - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@pytest.mark.asyncio -async def test_create_new_experiment( - _rmtree, path, force_memory_release, mlflow, mlflow_repository -): - model_name = 'model_name' - data = MagicMock() - prediction_data = MagicMock(spec=DataFrame) - retrain_data = { - 'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'prediction_data': prediction_data, - } - - mlflow_repository.get_model_params = MagicMock( - return_value={ - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - 'target_name': 'target_name', - } - ) - - mlflow_repository.get_experiment = MagicMock() - mlflow_repository.get_next_run_name = MagicMock() - mlflow_repository.log_model = AsyncMock() - path.exists.return_value = True - path.join.return_value = './tmp/artifacts/model_name' - - report = await mlflow_repository.create_new_experiment( - model_name, - data, - retrain_data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - 'pyfunc', - ) - - path.join.assert_called_once_with('./tmp/artifacts', 'model_name') - - mlflow_repository.get_model_params.assert_called_once_with('latest_production_id') - mlflow_repository.get_experiment.assert_called_once_with(model_name, create_if_not_exists=True) - mlflow_repository.get_next_run_name.assert_called_once_with( - mlflow_repository.get_experiment.return_value.name - ) - - data.to_csv.assert_called_once_with('./tmp/artifacts/model_name/retrain_data.csv', index=False) - # Verify prediction_data.to_csv was called with correct arguments - prediction_data.to_csv.assert_called_once() - assert ( - prediction_data.to_csv.call_args[0][0] == './tmp/artifacts/model_name/evaluation_data.csv' - ) - assert prediction_data.to_csv.call_args[1]['index'] is False - - mlflow.start_run.assert_called_once_with( - experiment_id=mlflow_repository.get_experiment.return_value.experiment_id, - run_name=mlflow_repository.get_next_run_name.return_value, - description='Retrain model model_name with new data', - ) - - mlflow_repository.log_model.assert_has_calls( - [ - call(retrain_data['data_model'], 'sklearn', 'data_model', metadata['metadata']), - call( - retrain_data['prediction_model'], 'pyfunc', 'prediction_model', metadata['metadata'] - ), - ] - ) - - mlflow.log_params.assert_called_once_with( - { - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - 'target_name': 'target_name', - 'retrain': True, - 'retrain_date': ANY, - 'source_run_id': 'latest_production_id', - 'retrain_samples': data.shape.__str__.return_value, - } - ) - - mlflow.log_artifact.assert_has_calls( - [ - call('./tmp/artifacts/model_name/retrain_data.csv'), - call('./tmp/artifacts/model_name/evaluation_data.csv'), - ] - ) - - force_memory_release.assert_called_once_with(mlflow_repository.logger) - - assert report == { - 'run_id': mlflow.start_run.return_value.__enter__.return_value.info.run_id, - 'experiment_id': mlflow_repository.get_experiment.return_value.experiment_id, - 'experiment_name': mlflow_repository.get_experiment.return_value.name, - } - - mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY) - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY - ) - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@pytest.mark.asyncio -async def test_create_new_experiment_error( - _rmtree, path, force_memory_release, mlflow, mlflow_repository -): - mlflow.start_run.side_effect = ValueError('error') - model_name = 'model_name' - data = MagicMock() - prediction_data = MagicMock(spec=DataFrame) - retrain_data = { - 'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'prediction_data': prediction_data, - } - - mlflow_repository.get_model_params = MagicMock( - return_value={ - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - 'target_name': 'target_name', - } - ) - - mlflow_repository.get_experiment = MagicMock() - mlflow_repository.get_next_run_name = MagicMock() - mlflow_repository.log_model = AsyncMock() - path.exists.return_value = True - path.join.return_value = './tmp/artifacts/model_name' - - with pytest.raises(ValueError): - await mlflow_repository.create_new_experiment( - model_name, - data, - retrain_data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - 'pyfunc', - ) - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -async def test_update_production_model_by_run_id(mlflow, mlflow_repository): - mlflow_repository.client.get_registered_model.return_value = MagicMock( - latest_versions=[ - MagicMock(version='1'), - MagicMock(version='2'), - MagicMock(version='3'), - ] - ) - output = await mlflow_repository.update_production_model_by_run_id( - '0', 'test', metadata['metadata'] - ) - - mlflow.register_model.assert_called_once_with( - 'runs:/0/prediction_model', - 'test', - ) - - mlflow_repository.client.get_registered_model.assert_called_once_with('test') - mlflow_repository.client.transition_model_version_stage.assert_called_once_with( - name='test', - version='3', - stage='Production', - archive_existing_versions=True, - ) - - assert output == { - 'model_name': 'test', - 'version': '3', - 'mlflow_run_id': '0', - } - - mlflow_repository.observe_lag.assert_has_calls( - [ - call(ANY, metrics.MODEL_WRITE_LAG, ANY), - call(ANY, metrics.MODEL_WRITE_LAG, ANY), - ] - ) - mlflow_repository.emit_metric.assert_has_calls( - [ - call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY), - call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY), - ] - ) - - -@pytest.mark.asyncio -async def test_update_production_model_by_run_id_error_register_model(mlflow, mlflow_repository): - mlflow.register_model.side_effect = ValueError('error') - with pytest.raises(ValueError): - await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata']) - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -async def test_update_production_model_by_run_id_error_transition_model_version_stage( - mlflow, mlflow_repository -): - mlflow_repository.client.transition_model_version_stage.side_effect = ValueError('error') - with pytest.raises(ValueError): - await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata']) - - mlflow_repository.emit_metric.assert_called_once_with( - metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY - ) - mlflow_repository.observe_lag.assert_not_called() - - -@pytest.mark.asyncio -async def test_update_production_model_by_run_id_error(mlflow, mlflow_repository): - mlflow_repository.client.get_registered_model.return_value = MagicMock( - get_registered_model=MagicMock(return_value=MagicMock(latest_versions={})) - ) - - try: - await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata']) - except Exception as e: - assert str(e) == 'Model versions is not a list' - else: - raise AssertionError('Expected Exception') - - -@pytest.mark.asyncio -async def test_transform_success(mlflow_repository): - data = MagicMock() - model_name = 'model' - model_config = { - 'retention_minutes': 60, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - } - - mlflow_repository.get_cached_operation = AsyncMock(return_value=data) - - mlflow_repository.detect_and_parse_datetime_index = MagicMock() - - output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata']) - - mlflow_repository.get_cached_operation.assert_called_once_with( - model_name=model_name, - data=data, - operation='transform', - retention=60, - flavor='sklearn', - metadata=metadata['metadata'], - ) - - mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with( - mlflow_repository.get_cached_operation.return_value, metadata['metadata'] - ) - - assert output['success'] is True - assert output['content'] is mlflow_repository.detect_and_parse_datetime_index.return_value - - -@pytest.mark.asyncio -async def test_transform_error(mlflow_repository): - data = MagicMock() - model_name = 'model' - model_config = { - 'retention_minutes': 60, - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - } - - mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error')) - - output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata']) - - mlflow_repository.get_cached_operation.assert_called_once_with( - model_name=model_name, - data=data, - operation='transform', - retention=60, - flavor='sklearn', - metadata=metadata['metadata'], - ) - - assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}} - - -@pytest.mark.asyncio -async def test_predict_success_array(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} - model_name = 'model' - mlflow_repository.get_cached_operation = AsyncMock( - return_value=DataFrame({'prediction': {'index_1': 2, 'index_2': 3}}) - ) - - output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata']) - - mlflow_repository.get_cached_operation.assert_called_once_with( - model_name=model_name, - data=ANY, - operation='predict', - retention=60, - flavor='pyfunc', - metadata=metadata['metadata'], - ) - - assert output['success'] is True - content = output['content'] - assert isinstance(content, DataFrame) - assert 'prediction' in content.columns - assert 'response_time' in content.columns - assert list(content.columns) == ['prediction', 'response_time'] - assert content.index.tolist() == data.index.tolist() - - -@pytest.mark.asyncio -async def test_predict_success_df(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} - model_name = 'model' - - mlflow_repository.get_cached_operation = AsyncMock( - return_value=DataFrame({'feat_1': {'index_3': 2, 'index_4': 3}}) - ) - - output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata']) - - mlflow_repository.get_cached_operation.assert_called_once_with( - model_name=model_name, - data=ANY, - operation='predict', - retention=60, - flavor='pyfunc', - metadata=metadata['metadata'], - ) - - assert output['success'] is True - content = output['content'] - assert isinstance(content, DataFrame) - assert 'prediction' in content.columns - assert 'response_time' in content.columns - assert list(content.columns) == ['prediction', 'response_time'] - assert content.index.tolist() == data.index.tolist() - - -@pytest.mark.asyncio -async def test_predict_error(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_name = 'model' - model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} - - mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error')) - - output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata']) - - mlflow_repository.get_cached_operation.assert_called_once_with( - model_name=model_name, - data=ANY, - operation='predict', - retention=60, - flavor='pyfunc', - metadata=metadata['metadata'], - ) - - assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}} - - -@pytest.mark.asyncio -async def test_retrain_model(mlflow_repository): - data = MagicMock() - model_name = 'test' - model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'} - - mlflow_repository.get_model_run_id = MagicMock() - mlflow_repository.fit_models = AsyncMock() - mlflow_repository.create_new_experiment = AsyncMock() - - output = await mlflow_repository.retrain_model( - data, model_name, model_config, metadata['metadata'] - ) - - mlflow_repository.get_model_run_id.assert_called_once_with(model_name, stage='Production') - - mlflow_repository.fit_models.assert_called_once_with( - model_name=model_name, - data=data, - transform_flavor='sklearn', - skip_transform=False, - predict_flavor='pyfunc', - target_name='target', - metadata=metadata['metadata'], - latest_production_id=mlflow_repository.get_model_run_id.return_value, - ) - - mlflow_repository.create_new_experiment.assert_called_once_with( - model_name=model_name, - data=data, - retrain_data=mlflow_repository.fit_models.return_value, - transform_flavor='sklearn', - predict_flavor='pyfunc', - metadata=metadata['metadata'], - latest_production_id=mlflow_repository.get_model_run_id.return_value, - ) - - assert output == { - 'success': True, - 'experiment': mlflow_repository.create_new_experiment.return_value, - 'message': 'Model retrained successfully.', - } - - -@pytest.mark.asyncio -async def test_retrain_model_error(mlflow_repository): - data = MagicMock() - model_name = 'test' - model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'} - mlflow_repository.get_model_run_id = MagicMock(side_effect=Exception('error')) - output = await mlflow_repository.retrain_model( - data, model_name, model_config, metadata['metadata'] - ) - assert output == { - 'success': False, - 'experiment': None, - 'message': 'Error retraining model test: error', - 'traceback': ANY, - } - - -@pytest.mark.asyncio -async def test_update_production_model(mlflow_repository): - experiment = {'run_id': '0', 'experiment_id': '0'} - model_name = 'test' - mlflow_repository.update_production_model_by_run_id = AsyncMock() - mlflow_repository.update_production_model_by_run_id.return_value = { - 'model_name': 'test', - 'version': '3', - 'mlflow_run_id': '0', - } - - output = await mlflow_repository.update_production_model( - experiment, model_name, metadata['metadata'] - ) - - mlflow_repository.update_production_model_by_run_id.assert_called_once_with( - '0', 'test', metadata['metadata'] - ) - - assert output == { - 'model_name': 'test', - 'version': '3', - 'mlflow_run_id': '0', - 'mlflow_experiment_id': '0', - } - - -def test_get_prediction_data_dataframe(mlflow_repository): - prediction_model = MagicMock() - retrain_dataset = DataFrame({'feat_1': [1, 2], 'target': [3, 4]}, index=['idx1', 'idx2']) - prediction_model.predict.return_value = DataFrame({'pred': [5, 6]}, index=['idx1', 'idx2']) - target_name = 'target' - predict_flavor = 'sklearn' - - result = mlflow_repository.get_prediction_data( - prediction_model, retrain_dataset, target_name, predict_flavor - ) - - prediction_model.predict.assert_called_once_with(retrain_dataset) - assert 'prediction' in result.columns - assert 'target' in result.columns - assert 'timestamp' in result.columns - assert result.index.tolist() == [0, 1] - - -def test_get_prediction_data_array(mlflow_repository): - prediction_model = MagicMock() - retrain_dataset = DataFrame({'feat_1': [1, 2], 'target': [3, 4]}, index=['idx1', 'idx2']) - prediction_model.predict.return_value = [5, 6] - target_name = 'target' - predict_flavor = 'sklearn' - - result = mlflow_repository.get_prediction_data( - prediction_model, retrain_dataset, target_name, predict_flavor - ) - - prediction_model.predict.assert_called_once_with(retrain_dataset) - assert 'prediction' in result.columns - assert 'target' in result.columns - assert 'timestamp' in result.columns - assert result.index.tolist() == [0, 1] - - -def test_get_prediction_data_pyfunc(mlflow_repository): - prediction_model = MagicMock() - retrain_dataset = DataFrame({'feat_1': [1, 2], 'target': [3, 4]}, index=['idx1', 'idx2']) - prediction_model.predict.return_value = DataFrame({'pred': [5, 6]}, index=['idx1', 'idx2']) - target_name = 'target' - predict_flavor = 'pyfunc' - - result = mlflow_repository.get_prediction_data( - prediction_model, retrain_dataset, target_name, predict_flavor - ) - - prediction_model.predict.assert_called_once_with({}, retrain_dataset) - assert 'prediction' in result.columns - assert 'target' in result.columns - assert 'timestamp' in result.columns - assert result.index.tolist() == [0, 1] - - -@patch('laborious.utils.repository.model_repository.pd.merge') -@patch('laborious.utils.repository.model_repository.isinstance') -@pytest.mark.asyncio -async def test_fit_models_skip_transform(isinstance_mock, pd_merge, mlflow_repository): - isinstance_mock.return_value = True - - data_model = MagicMock(target_variable='feat_2') - prediction_model = MagicMock() - mlflow_repository.download_model = AsyncMock( - side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')], - ) - mlflow_repository.detect_and_parse_datetime_index = MagicMock( - return_value=MagicMock( - drop_duplicates=MagicMock(return_value=MagicMock(columns=['feat_1'])) - ) - ) - mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame()) - - data = MagicMock() - - output = await mlflow_repository.fit_models( - 'model_name', - data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - True, - 'pyfunc', - 'feat_1', - ) - - data_model.fit.assert_not_called() - - assert output['data_model'] == {'model': data_model, 'artifact_path': 'artifact_path'} - - -@patch('laborious.utils.repository.model_repository.force_memory_release') -@patch('laborious.utils.repository.model_repository.path') -@patch('laborious.utils.repository.model_repository.rmtree') -@pytest.mark.asyncio -async def test_create_new_experiment_path_not_exists( - _rmtree, path, force_memory_release, mlflow, mlflow_repository -): - model_name = 'model_name' - data = MagicMock() - prediction_data = MagicMock(spec=DataFrame) - retrain_data = { - 'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'}, - 'prediction_data': prediction_data, - } - - mlflow_repository.get_model_params = MagicMock( - return_value={ - 'transform_flavor': 'sklearn', - 'predict_flavor': 'pyfunc', - 'target_name': 'target_name', - } - ) - - mlflow_repository.get_experiment = MagicMock() - mlflow_repository.get_next_run_name = MagicMock() - mlflow_repository.log_model = AsyncMock() - path.exists.return_value = False - path.join.return_value = './tmp/artifacts/model_name' - - await mlflow_repository.create_new_experiment( - model_name, - data, - retrain_data, - 'latest_production_id', - metadata['metadata'], - 'sklearn', - 'pyfunc', - ) - - _rmtree.assert_not_called() - - -@pytest.mark.asyncio -async def test_update_production_model_by_run_id_transition_error(mlflow, mlflow_repository): - mlflow_repository.client.get_registered_model.return_value = MagicMock( - latest_versions=[ - MagicMock(version='1'), - MagicMock(version='2'), - ] - ) - mlflow_repository.client.transition_model_version_stage.side_effect = Exception( - 'transition error' - ) - - with pytest.raises(Exception, match='transition error'): - await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata']) - - mlflow_repository.emit_metric.assert_called_with( - metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY - ) - - -@pytest.mark.asyncio -async def test_predict_success_ndarray(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} - model_name = 'model' - mlflow_repository.get_cached_operation = AsyncMock(return_value=np.array([5.0, 6.0])) - - output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata']) - - assert output['success'] is True - content = output['content'] - assert isinstance(content, DataFrame) - assert 'prediction' in content.columns - assert 'response_time' in content.columns - assert content.index.tolist() == data.index.tolist() diff --git a/tests/laborious/utils/test_connectors_config.py b/tests/laborious/utils/test_connectors_config.py index b94bb8d..62a267d 100644 --- a/tests/laborious/utils/test_connectors_config.py +++ b/tests/laborious/utils/test_connectors_config.py @@ -4,6 +4,7 @@ from laborious.utils.connectors_config import ( build_minio_config, build_mlflow_config, build_opc_config, + build_plugin_store_config, ) @@ -18,12 +19,22 @@ def test_build_mlflow_config_with_env_vars(): config = build_mlflow_config() # Assert - assert config['host'] == 'http://test-host' - assert config['port'] == 8080 + assert config['url'] == 'http://test-host:8080' assert config['username'] == 'test-user' assert config['password'] == 'test-pass' +def test_build_mlflow_config_host_already_has_port(): + environ['MLFLOW_HOST'] = 'http://tracker.example.com:443' + environ['MLFLOW_PORT'] = '8080' + environ['MLFLOW_USERNAME'] = 'u' + environ['MLFLOW_PASSWORD'] = 'p' + + config = build_mlflow_config() + + assert config['url'] == 'http://tracker.example.com:443' + + def test_build_mlflow_config_with_defaults(): # Arrange # Clear any existing env vars @@ -36,12 +47,30 @@ def test_build_mlflow_config_with_defaults(): config = build_mlflow_config() # Assert - assert config['host'] == 'http://localhost' - assert config['port'] == 5080 + assert config['url'] == 'http://localhost:5080' assert config['username'] == 'aignosi' assert config['password'] == 'aignosi' +def test_build_plugin_store_config_defaults(): + environ.pop('STORE_BASE_URL', None) + environ.pop('STORE_OWNER', None) + environ.pop('STORE_REPO', None) + environ.pop('STORE_BRANCH', None) + environ.pop('STORE_USERNAME', None) + environ.pop('STORE_PASSWORD', None) + environ.pop('STORE_CACHE_TTL_SECONDS', None) + environ.pop('PYPI_SERVER', None) + environ.pop('PYPI_USERNAME', None) + environ.pop('PYPI_PASSWORD', None) + + cfg = build_plugin_store_config() + assert cfg['base_url'] == 'http://localhost:3000' + assert cfg['owner'] == 'sientia' + assert cfg['repo'] == 'model-library-store' + assert cfg['pypi_index_url'] == 'http://localhost:5000' + + def test_build_opc_config_with_env_vars(): # Arrange environ['OPC_CONFIG'] = '{"opc": {"name": "test-opc", "url": "opc.tcp://test:4840"}}' diff --git a/tests/laborious/workflows/test_minimal_retrain.py b/tests/laborious/workflows/test_minimal_retrain.py index 94d86da..e8da86e 100644 --- a/tests/laborious/workflows/test_minimal_retrain.py +++ b/tests/laborious/workflows/test_minimal_retrain.py @@ -34,8 +34,7 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain): 'table_name': 'test_table', 'model_config': { 'target': 'test_target', - 'transform_flavor': 'test_transform_flavor', - 'predict_flavor': 'test_predict_flavor', + 'retention_minutes': 0, }, } @@ -165,8 +164,7 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim 'table_name': 'test_table', 'model_config': { 'target': 'test_target', - 'transform_flavor': 'test_transform_flavor', - 'predict_flavor': 'test_predict_flavor', + 'retention_minutes': 0, }, } @@ -227,8 +225,7 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim 'table_name': 'test_table', 'model_config': { 'target': 'test_target', - 'transform_flavor': 'test_transform_flavor', - 'predict_flavor': 'test_predict_flavor', + 'retention_minutes': 0, }, } diff --git a/values.yaml b/values.yaml index 31dd4fa..2b1c49d 100644 --- a/values.yaml +++ b/values.yaml @@ -188,6 +188,32 @@ env: - 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