24 Commits

Author SHA1 Message Date
Eduardo Rios
695e4c07a6 SIENTIAPDE-2072: bump sientia_do pin to 1.12.2
Picks up the notification timestamp -> native datetime fix.
2026-08-17 16:42:31 -03:00
Bruno Domingues
6a8c41b328 fix(pytest): Workaround unraisableexception plugin crash
Disables the unraisableexception plugin in pytest due to a known bug in
pytest>=9.1 where it crashes with tracemalloc errors when multiple
unraisable exceptions occur close together.
2026-08-04 15:21:14 -03:00
Bruno Domingues
773c980fc3 fix(storage): Prevent AttributeError when closing uninitialized Postgres engine 2026-08-04 14:57:52 -03:00
Bruno Domingues
35efb67c89 fix(test_activities): Use aclose for async mock shutdown 2026-08-04 14:43:15 -03:00
Bruno Domingues
acd925cf2a refactor(opc): Rename async close method to aclose
Renamed the OPC.close asynchronous method to OPC.aclose to align with common Python conventions for asynchronous context managers and methods, improving clarity. All call sites and tests have been updated accordingly.
2026-08-04 12:01:06 -03:00
Bruno Domingues
310ceea0d8 ci(quality-gate): Configure push triggers, concurrency, and granular permissions 2026-08-04 11:41:23 -03:00
Bruno Domingues
2d73ec9ec2 Merge pull request #42 from Aignosi/feature/SIENTIAPDE-1945
SIENTIAPDE-1945: Remove values.yaml from sientia-module Helm chart
2026-07-10 11:49:36 -03:00
Bruno Domingues
2142143ab9 SIENTIAPDE-1945: Delete values.yaml configuration file for sientia-module Helm chart. 2026-07-08 22:12:50 -03:00
Bruno Domingues
c07457bfbe chore(sonar): update project key 2026-07-01 20:59:29 -03:00
vitor-aignosi
086b12492e Merge pull request #41 from Aignosi/feature/SIENTIAPDE-1646-legacy-laborious-worker
SIENTIAPDE-1646: Refactor Worker Task Queue Management and Update Dependencies
2026-05-21 15:30:09 -03:00
vitor-aignosi
4ea0754f0c SIENTIAPDE-1646
Enhance MLFlow run ID resolution with error handling for missing and invalid source URIs

- Added checks in `get_model_run_id` method to raise exceptions for models with missing or invalid source URIs.
- Introduced new test cases to validate error handling for these scenarios.
- Updated `requirements-light.txt` to include `mlflow` as a dependency.
2026-05-20 09:53:53 -03:00
vitor-aignosi
ddb1618209 SIENTIAPDE-1646
Update quality-gate workflow to use python-quality-gate template
2026-05-20 09:30:05 -03:00
vitor-aignosi
90f8bdda61 SIENTIAPDE-1646
Update requirements.txt to align with recent dependency changes and ensure compatibility across the project.
2026-05-20 09:26:35 -03:00
vitor-aignosi
2ccda3e440 SIENTIAPDE-1646
Refactor worker task queue management and update README

- Introduced runtime-scoped task queues for workflows, replacing legacy queue names.
- Updated worker implementation to utilize `sientia_do.temporal.worker.prepare_worker`.
- Added `RUNTIME` environment variable to configure task queue suffixes.
- Enhanced README documentation to reflect changes in task queue structure and worker setup.
2026-05-19 17:07:18 -03:00
vitor-aignosi
d856150e24 Update requirements.txt 2026-05-19 14:45:26 -03:00
vitor-aignosi
6569810756 Update requirements.txt 2026-05-19 14:41:58 -03:00
vitor-aignosi
7153f1da0d SIENTIAPDE-1646
Remove requirements-light.txt and update requirements.txt to specify versions for asyncua and new sientia dependencies.
2026-05-19 14:13:23 -03:00
vitor-aignosi
fcc8920a8b Merge pull request #40 from Aignosi/fix/SIENTIAPDE-1811-fix
SIENTIAPDE-1811: Update dependencies, gitignore, and strengthen OPC UA error handling
2026-05-19 14:01:37 -03:00
vitor-aignosi
cd1be2430a SIENTIAPDE-1811
Update .gitignore, requirements, and enhance OPC UA error handling

- Added new entries to .gitignore for openspec and cursor directories.
- Updated sientia-dataops-library dependency version in requirements-light.txt from 1.10.4 to 1.12.0.
- Enhanced OPC UA communication by refining reconnect logic and error handling in opc_repository.py, including the introduction of a reconnect flag and improved session management.
- Updated tests to cover new reconnect scenarios and ensure robust error handling for protocol states.
2026-05-19 11:45:08 -03:00
vitor-aignosi
7c1dae8ef6 Merge pull request #39 from Aignosi/fix/SIENTIAPDE-1811
Enhance OPC UA Communication and Metrics Tracking
2026-05-18 13:32:33 -03:00
vitor-aignosi
2a6def4056 SIENTIAPDE-1811
SIENTIAPDE-1811 Implement OPC write error handling and refactor tag writing logic

- Introduced a new function `_apply_opc_write_error` to manage session and reconnect flags based on OPC write error responses.
- Refactored the `_write_tags_from_config` method to streamline the writing of OPC tags for both prediction and confidence data.
- Enhanced unit tests to cover various scenarios for OPC write errors, including session bad and reconnect in progress cases.
- Updated existing tests to validate the new logic and ensure robust error handling.
2026-05-18 10:36:57 -03:00
vitor-aignosi
7d59fd7c8c SIENTIAPDE-1811
Update values.yaml to rename worker references and adjust GitHub branch for SIENTIAPDE-1811. Changed nameOverride, fullnameOverride, and service account name to "sientia-laborious-legacy-worker" and updated the GITHUB_BRANCH value to "fix/SIENTIAPDE-1811".
2026-05-18 10:17:52 -03:00
vitor-aignosi
e3636d4b88 SIENTIAPDE-1811
Enhance OPC UA testing framework and documentation

- Added a new marker in pyproject.toml for tests using the in-process OPC UA server.
- Updated opc-communication.md to clarify E2E test scenarios involving the real OPC server and mock server.
- Introduced an in-process asyncua OPC UA server fixture in conftest.py for E2E tests.
- Created a new fixture for activities using the real OpcRepository connected to the in-process server.
- Updated scenarios.md to include instructions for running OPC real-server tests.
2026-05-15 15:46:39 -03:00
vitor-aignosi
638d5b70b4 SIENTIAPDE-1811
Enhance OPC UA communication and metrics tracking

- Updated README.md to include new OPC UA Communication section and detailed metrics for session and write diagnostics.
- Added new metrics in laborious/metrics.py for tracking OPC UA session states and write attempts.
- Refactored OPC activity in laborious/activities/opc.py to handle session errors and improve error reporting.
- Updated e2e tests to cover new scenarios for OPC session/channel errors and reconnect handling.
- Modified .gitignore to include relatorio files and mlruns directory.
- Added ipykernel to requirements-dev.txt for Jupyter notebook support.
2026-05-15 15:28:28 -03:00
71 changed files with 6798 additions and 7213 deletions

View File

@@ -11,22 +11,6 @@ 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"

View File

@@ -1,15 +1,29 @@
name: Quality gate
on:
push:
branches:
- main
- 'release/**'
- 'feature/**'
pull_request:
branches:
- main
- 'release/**'
- 'feature/**'
types: [ opened, synchronize, reopened ]
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
quality-gate:
uses: Aignosi/github_workflow_templates/.github/workflows/dataops-module-quality-gate.yml@main
permissions: write-all
uses: Aignosi/github_workflow_templates/.github/workflows/python-quality-gate.yml@main
permissions:
contents: read
pull-requests: write
issues: write
with:
project_name: 'laborious'
repositories: 'sientia-dataops-library, sientia-mlops-library'

4
.gitignore vendored
View File

@@ -54,5 +54,5 @@ catboost_info/
mlruns/
relatorio*
openspec/*
.cursor/*
openspec/
.cursor/

135
README.md
View File

@@ -126,10 +126,15 @@ Laborious uses a Temporal-based architecture with strong separation of concerns
### Key Components
#### **Worker (`laborious/worker/worker.py`)**
- Temporal client setup, worker lifecycle, task queues
- Temporal client setup, four workers via `sientia_do.temporal.worker.prepare_worker`
- Runtime-scoped task queues: `{workflow}-{RUNTIME}-queue` for all workflows
- Metrics server initialization, notification handler setup
- Graceful shutdown and autoscaling-friendly behavior
**Breaking (schedulers):** drift and simple_metrics queues are no longer `drift-queue` /
`simple_metrics-queue`. Use `drift-{RUNTIME}-queue` and `simple_metrics-{RUNTIME}-queue`
matching the worker pod `RUNTIME` env (same as `predictions_batch` / `minimal_retrain`).
#### **Workflows (`laborious/workflows/`)**
- `predictions_batch.py`: Batch prediction entry point
- `sub_workflows/prediction_process.py`: Core prediction pipeline
@@ -164,7 +169,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
- ML models are loaded via `SientiaMLflowRepository` (wrapper-based, `@production` alias) constructed in `Activities` from `build_mlflow_config()`.
- `repository/model_repository.py`: MLFlow operations and retraining
- `repository/opc_repository.py`: OPC UA client, writes, session recovery (see [OPC UA Communication](#opc-ua-communication))
- `repository/minio_manager.py`: MinIO object storage operations
- `filters/conditional_filters.py` and `filters/mlflow_filters.py`
@@ -464,7 +469,7 @@ flowchart LR
"source_table_name": "laborious_data",
"target_table_name": "drift_metrics",
"interval": 60,
"model_config": { "target": "temperature", "alias": "production" },
"model_config": { "target": "temperature" },
"drift_metrics": ["kolmogorov_smirnov", "jensen_shannon", "wasserstein"],
"chunk_period": "min"
}
@@ -499,7 +504,7 @@ flowchart LR
"data_table_name": "laborious_data",
"target_table_name": "simple_metrics",
"interval_minutes": 60,
"model_config": { "target": "temperature", "alias": "production" },
"model_config": { "target": "temperature" },
"metrics": ["rmse", "mse", "mae", "r2"]
}
```
@@ -731,6 +736,7 @@ tests/
│ │ ├── test_conditional_filters.py
│ │ └── test_mlflow_filters.py
│ └── repository/
│ ├── test_model_repository.py
│ └── test_opc_repository.py
```
@@ -801,6 +807,7 @@ See [OPC UA Communication](#opc-ua-communication) for semantics, concurrency, an
|----------|-------------|---------|----------|
| `TEMPORAL_HOST` | Temporal server address | `localhost:7233` | Yes |
| `TEMPORAL_NAMESPACE` | Temporal namespace | `laborious` | No |
| `RUNTIME` | Task queue suffix for all workflows (`{workflow}-{RUNTIME}-queue`) | _(none)_ | Yes |
| `POSTGRES_HOST` | PostgreSQL hostname | `localhost` | Yes |
| `POSTGRES_PORT` | PostgreSQL port | `5432` | Yes |
| `POSTGRES_USER` | PostgreSQL username | `sientia` | Yes |
@@ -812,20 +819,7 @@ See [OPC UA Communication](#opc-ua-communication) for semantics, concurrency, an
| `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 |
@@ -969,7 +963,7 @@ For single OPC server, use individual environment variables:
### PI Web API Configuration
PI Web API configuration is built from environment variables using the `build_api_config` function from `sientia_do.utils.connectors_config`. The configuration includes:
PI Web API configuration is built from environment variables using the `build_api_config` function from `sientia_do.connectors_config`. The configuration includes:
- `PI_WEB_API_BASE_URL`: Base URL of the PI Web API server
- `PI_WEB_API_AUTH_TYPE`: Authentication type ('basic' or 'bearer')
@@ -1002,30 +996,9 @@ Where:
MongoDB pipeline configuration:
#### MongoDB input samples (updated)
#### Predictions Batch Workflow configuration sample
Updated examples are available in `input_sample.json` at the repository root.
The sample already reflects the runtime-aware and alias-based flow:
- `model_config` uses `target`, `retention_minutes`, and `alias`.
- `transform_flavor` / `predict_flavor` are not used anymore.
Example model document:
```json
{
"id": "4",
"name": "vcm-nox",
"active": false,
"model_config": {
"alias": "production",
"retention_minutes": 60,
"target": "CI-W3W01A3"
}
}
```
Example predictions_batch schedule document:
This is the configuration for the Predictions Batch Workflow, to be inserted into the MongoDB pipeline collection.
```json
{
@@ -1035,67 +1008,41 @@ Example predictions_batch schedule document:
"frequency": "30s",
"max_retry_policy": 1,
"query": "select * from sientia_data.laborious_data where model_id = 1 and \"timestamp\" > NOW() - INTERVAL '5 minutes' order by \"timestamp\" desc limit 30;",
"retention_time": 60,
"write_tags": [
{
"server_id": "1",
"server_id": "server1",
"type": "prediction",
"addr": "ns=2;i=5",
"data_type": "double"
},
{
"server_id": "1",
"server_id": "server1",
"type": "confidence",
"addr": "ns=2;i=5",
"addr": "ns=2;i=6",
"data_type": "double"
}
],
"input_filters": [
{
"filter_name": "EMPTY_DATA",
"policy": "STOP"
},
{
"filter_name": "SPECIFIC_VARIABLES_NULL_VALUES",
"policy": "CONTINUE",
"config": {
"variables": ["Counter"]
}
"input_filters": {
"EMPTY_DATA": {"POLICY": "STOP"},
"SPECIFIC_VARIABLES_NULL_VALUES": {
"POLICY": "CONTINUE",
"config": {"variables": ["Counter"]}
}
],
"mlflow_transform_filters": [
{
"filter_name": "API_ERROR",
"policy": "REPEAT"
},
{
"filter_name": "NAN_VALUES",
"policy": "STOP"
}
],
"mlflow_predict_filters": [
{
"filter_name": "API_ERROR",
"policy": "CONTINUE"
}
],
},
"mlflow_transform_filters": {
"API_ERROR": {"POLICY": "REPEAT"},
"NAN_VALUES": {"POLICY": "STOP"}
},
"mlflow_predict_filters": {
"API_ERROR": {"POLICY": "CONTINUE"}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"active": true,
"updated_at": {
"$date": "2026-01-27T17:35:01.600Z"
"$date": "2025-09-16T10:00:00.000Z"
},
"save_transform": false,
"pi_web_api_output_config": {
"endpoint": "/streamsets/value",
"prediction_tags": {},
"confidence_tags": {}
},
"model_config": {
"alias": "production",
"retention_minutes": 60,
"target": "CI-W3W01A3"
},
"datetime_columns": ["timestamp", "created_at"]
"datetime_columns": ["timestamp", "created_at"],
"predictions_storage_policy": "lts:1"
}
```
@@ -1113,8 +1060,11 @@ This is the configuration created by the Orchestrator in Temporal.
"EMPTY_DATA":{"config":{},"policy":"STOP"}
},
"model_config":{
"target":"sensor_or_label_column",
"retention_minutes":60
"is_compressed":true,
"predict_flavor":"pyfunc",
"retention_minutes":60,
"retention_target":"artifact",
"transform_function_keyword":"transform"
},
"model_id":"352",
"model_name":"courier",
@@ -1153,7 +1103,7 @@ laborious/
│ ├── prediction_process.py # Core prediction workflow
│ └── format_and_export_prediction.py # Export workflow
├── worker/ # Worker implementation
│ └── worker.py # Entrypoint; workers built via `sientia_do.temporal.worker.prepare_worker`
│ └── worker.py # Main worker orchestrator (uses sientia_do prepare_worker)
├── utils/ # Utility functions
│ ├── connectors_config.py # Environment-driven config builders
│ ├── models/ # Data models
@@ -1162,6 +1112,7 @@ 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
@@ -1188,7 +1139,7 @@ laborious/
2. **MLFlow Connection Issues**
- Verify MLFlow server is running and accessible
- Check authentication credentials and permissions
- Ensure model names exist and the expected alias (for example `production`) is registered
- Ensure model names and versions exist
3. **Database Connection Issues**
- Verify PostgreSQL service is running
@@ -1233,7 +1184,9 @@ export LOG_LEVEL=DEBUG
### Scaling Considerations
- **Horizontal Scaling**: Deploy multiple worker instances
- **Task Queue Distribution**: Use multiple task queues for different workflow types
- **Task Queue Distribution**: One worker pod per `RUNTIME`; queues are
`predictions_batch-{RUNTIME}-queue`, `minimal_retrain-{RUNTIME}-queue`,
`drift-{RUNTIME}-queue`, `simple_metrics-{RUNTIME}-queue`
- **Database Performance**: Optimize indexes and connection pooling
- **MLFlow Performance**: Configure appropriate model serving resources

View File

@@ -1,149 +0,0 @@
# E2E test run report
**Date:** 2026-05-08
**Command:** `source venv/bin/activate && rtk pytest e2e/ -v --tb=short`
**Environment:** Linux, Python 3.11.15, pytest 9.0.3
## Summary
| Metric | Count |
|--------|------:|
| Collected | 43 |
| **Passed** | **37** |
| **Failed** | **6** |
Full pytest output (compressed by `rtk`) was written to:
`~/.local/share/rtk/tee/1778268193_pytest.log`
---
## Failed tests (6)
1. `e2e/test_drift.py::test_drift_happy_path_persists_all_columns_with_reference_data`
2. `e2e/test_drift.py::test_drift_uses_30pct_fallback_when_reference_unavailable`
3. `e2e/test_drift.py::test_drift_chunk_period_seconds_preserves_seconds_in_chunk_start_date`
4. `e2e/test_predictions_batch_prediction_process.py::test_scenario_2_1_3_input_gate_triggers_repeat`
5. `e2e/test_predictions_batch_prediction_process.py::test_scenario_2_2_3_transform_gate_triggers_repeat`
6. `e2e/test_predictions_batch_prediction_process.py::test_scenario_2_3_3_predict_gate_triggers_repeat`
**Follow-up:** JensenShannon NULLs for the drift happy-path scenario are fully traced (histogram out-of-range + `density=True`, vs NannyMLs leftover bin) in [DRIFT_JS_NULL_INVESTIGATION.md](./DRIFT_JS_NULL_INVESTIGATION.md).
---
## Failure group A — Drift: `drift_metrics.value` NOT NULL (2 tests)
### Error
`psycopg2.errors.NotNullViolation`: null value in column `value` of relation `sientia_data.drift_metrics` violates not-null constraint.
Example failing row (from logs): `feature=sensor_2`, `method=jensen_shannon`, `value=null`, with `kolmogorov_smirnov` / `wasserstein` populated for the same chunk.
The bulk INSERT built by `Activities.export_data_to_postgres` includes parameters such as `'value__4': None` for `jensen_shannon` on a given chunk.
### Root cause
`calculate_drift` (real `DriftAnalysis` + `ModelMetrics.get_drift_metrics`) can emit **NaN / missing** values for some metric methods (here **JensenShannon**) on some features/chunks. Pandas/SQLAlchemy turns that into SQL `NULL`, while the E2E schema (mirroring production) defines:
```sql
value numeric NOT NULL
```
in `e2e/db_schema.sql` for `sientia_data.drift_metrics`.
### Recommended fixes (pick one consistent with product rules)
1. **Application layer (preferred if NULLs are never valid in production):** Before export, sanitize the drift dataframe — e.g. drop rows where `value` is null/NaN, or replace with a defined sentinel (only if product agrees), or skip emitting that method row when the statistic is undefined.
2. **Analytics layer:** Harden the JensenShannon (and similar) paths so they always return a finite float for the supported inputs, or explicitly map “undefined” to an agreed numeric convention.
3. **Schema (only if product allows missing metrics):** Align DDL with reality by making `value` nullable — **only** if production and downstream consumers already expect missing metrics; the E2E comment in `test_drift.py` suggests `feature`/`timestamp` nullable cases exist, but `value` is still listed in `NON_NULL_DRIFT_COLUMNS`.
### Tests affected
- `test_drift_happy_path_persists_all_columns_with_reference_data`
- `test_drift_uses_30pct_fallback_when_reference_unavailable`
---
## Failure group B — Drift: chunk period seconds (1 test)
### Error
`AssertionError: expected at least one drift row to be persisted` (`e2e/test_drift.py:493`).
The workflow run completed without failing the test via `WorkflowFailureError`, but **no rows** were found in `sientia_data.drift_metrics` for the model.
### Likely cause
In `Drift.run`, persistence runs only when `if drift_data:` is truthy (`laborious/workflows/drift.py`). An **empty** drift result skips `export_data_to_postgres`, so the table stays empty.
Probable reasons:
- With **`chunk_period='s'`** and only **three** target rows (30 s spacing), `calculate_drift` / `DriftAnalysis` may produce **no output rows** (insufficient data per chunk or internal filters).
- Less likely here: time-window mismatch — timestamps are built from `datetime.now(UTC)` with `interval: 60` minutes from `drift_base.json`, so data should still fall in the window.
### Recommended fixes
1. **Test data:** Increase the number of second-spaced points (and/or span multiple chunk boundaries) so the analyzer reliably emits at least one chunk row.
2. **Product code:** If sub-minute chunking is required to always produce metrics when any data exists, adjust `ModelMetrics` / `DriftAnalysis` integration for small-N second buckets.
3. **Diagnostics:** Run the same scenario with `pytest -s` and confirm logs for “empty `drift_data`” vs export errors.
Solução a ser aplicada:
Dividir em dois testes:
1. Teste com dados suficientes para produzir pelo menos uma linha de drift
2. Teste com dados insuficientes para produzir pelo menos uma linha de drift, mas ja esperando os erros e validando que nao foi persistido nada
### Test affected
- `test_drift_chunk_period_seconds_preserves_seconds_in_chunk_start_date`
---
## Failure group C — Prediction REPEAT path: unique constraint on predictions (3 tests)
### Error
`psycopg2.errors.UniqueViolation`: duplicate key value violates unique constraint **`unique_model_id_timestamp`** on `sientia_data.predictions` (`model_id`, `timestamp`).
### Root cause
REPEAT is handled by `Activities.repeat_last_prediction` (registered from `sientia_do.temporal.activities.postgres_sync.Postgres`, via `Storage` inheritance in `laborious/activities/storage.py`). The E2E tests seed a prior row with a **fixed** timestamp:
```python
# e2e/test_predictions_batch_prediction_process.py — insert_sample_prediction
VALUES ({model_id}, '2024-01-01 12:00:00+00:00', ...)
```
The helper `assert_repeat` expects **two** rows with the same `(model_id, prediction, prediction_confidence, prediction_status)` but compares **only** those columns — not `timestamp` (`e2e/helpers.py`). So the intended behavior is: **duplicate business payload**, not necessarily **duplicate primary unique key** `(model_id, timestamp)`.
If `repeat_last_prediction` **INSERT**s a copy using the **same** `timestamp` as the last prediction, Postgres correctly rejects the second insert.
### Recommended fixes
1. **Implement REPEAT as UPSERT:** Use `ON CONFLICT (model_id, timestamp) DO UPDATE` (or the projects existing `export_data_to_postgres` `on_conflict` pattern used in `FormatAndExportPrediction`) when writing the repeated prediction — if the product definition of REPEAT is “refresh same logical slot.”
2. **Insert with the current batch timestamp:** Copy numeric/status fields from the last row but set `timestamp` to the **new** batch instant (e.g. the workflows `last_timestamp` / slice timestamp). This matches `assert_repeat`, which does not assert on `timestamp`.
3. **Test-only change (weakest):** Relax assertions or change seed data — only if production behavior is “duplicate key is expected” (unlikely).
Solução a ser aplicada:
Alterar o teste para usar o timestamp da ultima execucao do batch, ao inves do timestamp do primeiro batch.
2. **Insert with the current batch timestamp:** Copy numeric/status fields from the last row but set `timestamp` to the **new** batch instant (e.g. the workflows `last_timestamp` / slice timestamp). This matches `assert_repeat`, which does not assert on `timestamp`.
### Tests affected
- `test_scenario_2_1_3_input_gate_triggers_repeat`
- `test_scenario_2_2_3_transform_gate_triggers_repeat`
- `test_scenario_2_3_3_predict_gate_triggers_repeat`
---
## Passing areas (sanity check)
All scenarios in `test_child_workflows_e2e.py`, `test_minimal_retrain.py`, `test_minio_offload.py`, `test_predictions_batch_format_export.py`, `test_predictions_batch_main_workflow.py` (except the three REPEAT cases above), and `test_simple_metrics.py` **passed** in this run.
---
## Suggested order of work
1. Fix **drift `value` NULL** — unblocks two high-value drift E2Es and may clarify the chunk-seconds scenario if exports start succeeding consistently.
2. Fix **`repeat_last_prediction` uniqueness** — unblocks three prediction-process E2Es; implementation likely lives in **`sientia_do`** Postgres activities, not in this repo.
3. Revisit **`test_drift_chunk_period_seconds`** data volume / expectations after drift export is stable.

View File

@@ -1,8 +1,6 @@
# OPC UA communication (Laborious)
Laborious exports predictions to OPC UA servers through `OpcRepository` ([`laborious/utils/repository/opc_repository.py`](../laborious/utils/repository/opc_repository.py)) and the synchronous Temporal activity layer in [`laborious/activities/opc.py`](../laborious/activities/opc.py). The repository uses `asyncua.sync.Client` (asyncio on a background thread) so activities remain blocking without `async def`.
OPC reconnect, write error classification (`opc_error_kind`), and activity confidence/comment behavior are converted from the **async** implementation on `main` at `fcc8920a8be4` (`asyncua.Client` + `asyncio` reconnect task → `threading` reconnect thread). Re-convert with `scripts/convert_opc_async_to_sync.py` when `main` OPC files change.
Laborious exports predictions to OPC UA servers through `OpcRepository` ([`laborious/utils/repository/opc_repository.py`](../laborious/utils/repository/opc_repository.py)) and the Temporal activity layer in [`laborious/activities/opc.py`](../laborious/activities/opc.py).
Implementation plan for session/channel recovery on Tier-1 `Bad*` errors: [`.cursor/plans/opc_bad_reconnect_ac4c6045.plan.md`](../.cursor/plans/opc_bad_reconnect_ac4c6045.plan.md).
@@ -14,7 +12,7 @@ Worker (long-lived)
├── connect / disconnect / validate_connection (read-only)
├── _connect_locked / _reconnect_locked (under _connection_lock)
├── write_data (single attempt per call)
└── background reconnect on Tier-1 Bad*, closed protocol, or stale session
└── background reconnect on Tier-1 Bad*, protocol closed, or stale session
Temporal activity write_opc_data
└── OPC.manage_output_tags → write_data per tag (sequential per activity)
@@ -28,9 +26,9 @@ One worker process holds one `OpcRepository` instance per configured server. Mul
|-------|----------|
| Startup | `init_opc()` creates repositories and calls `connect()``_connect_locked()` |
| Steady state | `validate_connection()` is read-only (`protocol.state` only); `_session_ready` is checked in `write_data` |
| Tier-1 Bad* / protocol closed / session not ready | `_start_reconnect(reason)``_run_reconnect` (thread)`_reconnect_locked()` (respects `reconnection_interval`) |
| Write | `write_data()` checks in-flight reconnect thread, `_session_ready`, validates, then one `get_node` + `write_value` |
| Shutdown | `disconnect()` sets `_allow_reconnect = False`, then tears down session |
| Tier-1 Bad* / protocol closed | `_start_reconnect``_run_reconnect``_reconnect_locked()` (respects `reconnection_interval`) |
| Write | `write_data()` checks reconnect task, `_session_ready`, validates protocol, then one `get_node` + `write_value` |
| Shutdown | `close()` disconnects all repositories |
### Session and channel timeouts
@@ -38,7 +36,7 @@ Requested session and secure-channel lifetime: **10 minutes** (`OPC_UA_SESSION_A
### Reconnection interval
`OPC_RECONNECTION_INTERVAL` is in **seconds** (default `120`). It gates **background** reconnect after Tier-1 `Bad*` (`last_reconnection_time` is updated only in `_reconnect_locked()`). It limits load on the OPC server when many workflows fail at once.
`OPC_RECONNECTION_INTERVAL` is in **seconds** (default `120`). It gates **background** reconnect after Tier-1 `Bad*`, closed protocol, or stale session (`last_reconnection_time` is updated only in `_reconnect_locked()`). It limits load on the OPC server when many workflows fail at once.
## Concurrency: connection lock and session readiness
@@ -46,9 +44,8 @@ To allow **multiple concurrent writes** when the session is healthy, but **block
| Primitive | Role |
|-----------|------|
| `_connection_lock` (`threading.Lock`) | Held for the entire `disconnect``connect` path. Only one connection-maintenance task at a time. |
| `_session_ready` (`threading.Event`) | Set when a session is ready for writes; cleared before reconnect starts and set again after a successful connect. |
| `_allow_reconnect` | Cleared in `disconnect()` so shutdown does not spawn reconnect threads |
| `_connection_lock` (`asyncio.Lock`) | Held for the entire `disconnect``connect` path. Only one connection-maintenance task at a time. |
| `_session_ready` (`asyncio.Event`) | Set when a session is ready for writes; cleared before reconnect starts and set again after a successful connect. |
**Connection methods (caller holds `_connection_lock` for `_*_locked` helpers):**
@@ -64,24 +61,33 @@ Public `connect()` / `disconnect()` acquire the lock and call `_connect_locked()
**Write path (`write_data`):**
1. If a reconnect **thread** is alive → `reconnect_in_progress`.
2. If `_session_ready` is cleared → schedule `SessionNotReady` reconnect; return `reconnect_in_progress` or `connection_lost`.
3. If `validate_connection()` fails (protocol closed) → schedule `ProtocolClosed` reconnect; return `connection_lost`.
4. Single `get_node` + `write_value` (no retry in the same call).
1. If a reconnect task is **in flight****fail immediately** (`opc_error_kind=reconnect_in_progress`).
2. If `_session_ready` is cleared and no task is running → schedule reconnect (`SessionNotReady`); fail with `connection_lost` or `reconnect_in_progress` if a task started.
3. `validate_connection()` checks `protocol.state` only (read-only). If closed → schedule reconnect (`ProtocolClosed`) and fail with `opc_error_kind=connection_lost`.
4. Single `get_node` + `write_value` (no retry). Tier-1 `Bad*` on write also schedules reconnect.
**Reconnect path (`_run_reconnect`):**
1. `_start_reconnect` clears `_session_ready` and starts a daemon thread when the interval allows and `_allow_reconnect` is true.
2. `with _connection_lock:``_reconnect_locked()`.
1. `_start_reconnect` clears `_session_ready` and schedules the task when the interval allows and `_allow_reconnect` is true.
2. `async with _connection_lock:``_reconnect_locked()`.
3. `_session_ready` is set on successful `_open_session()`.
4. `disconnect()` sets `_allow_reconnect=False` so shutdown does not respawn sessions.
A second `_connect_locked()` while a session is already open raises `OpcSessionAlreadyConnectedError` (disconnect first).
**asyncua note:** Concurrent `write_value` on the same session is only safe if the stack tolerates it. If production shows issues, serialize writes while keeping the connection lock semantics above.
**asyncua note:** Concurrent `write_value` on the same session is only safe if the stack tolerates it. If production shows issues, serialize writes with an optional `asyncio.Semaphore(1)` while keeping the connection lock semantics above.
## Tier-1 `Bad*` errors and reconnect
**Future threads:** replace `asyncio.Lock` / `Event` with `threading` primitives or route all OPC I/O through one dedicated loop.
When the server invalidates the session (e.g. `BadSessionIdInvalid`) but the client still sees transport as open, `write_data` fails once, records the OPC status in metrics, and **schedules** reconnect if:
## Reconnect triggers
Background reconnect is scheduled when:
- `validate_connection()` sees a closed or missing protocol (`ProtocolClosed`).
- `_session_ready` is clear after a failed reconnect (`SessionNotReady`).
- A write raises a Tier-1 `UaStatusCodeError` in `RECONNECTABLE_OPC_BAD_NAMES`.
For Tier-1 `Bad*` when the server invalidates the session (e.g. `BadSessionIdInvalid`) but the client still sees transport as open, `write_data` fails once, records the OPC status in metrics, and **schedules** reconnect if:
- The exception is a `UaStatusCodeError` whose name is in `RECONNECTABLE_OPC_BAD_NAMES` (see plan), and
- `reconnection_interval` has elapsed since `last_reconnection_time`, and

View File

@@ -1,55 +0,0 @@
# Specification: `sientia_model` — JensenShannon drift (`DriftAnalysis`)
This document describes what **`sientia_model.analytics.drift_analysis.DriftAnalysis`** should change so downstream consumers (e.g. Laborious `calculate_drift` → Postgres `drift_metrics.value NOT NULL`) no longer receive **NaN** for JensenShannon on valid finite data.
## Scope
- **File:** `sientia_model/analytics/drift_analysis.py`
- **Method:** `_jensen_shannon_distance(self, ref: np.ndarray, cur: np.ndarray, bins: int = 20) -> float`
- **Callers:** `detect_univariate_drift` uses this for the `jensen_shannon` method; results are written to `value` in the drift dataframe.
## Problem
The current implementation:
1. Builds bin edges from **`ref`** only: `np.histogram(ref, bins=bins, density=True)`.
2. Builds the chunk histogram with **`density=True`** on the same edges: `np.histogram(cur, bins=edges, density=True)`.
When **every** value in **`cur`** falls **outside** the closed support implied by those edges (typical case: production chunk drifted above the reference max or below the reference min), NumPy yields **all-zero counts** for `cur`. With **`density=True`**, normalization does **0/0**, producing **NaN** for the whole histogram, which propagates to **`float('nan')`** in `detect_univariate_drift`**SQL NULL** where `value` is `NOT NULL`.
This appears in Laborious when:
- `model_config.target` excludes the main target column from univariate features, so another feature (e.g. `sensor_2`) is compared chunk-by-chunk against the full reference series for that feature.
- Reference and current ranges do not overlap for some chunks (strong drift or different scaling).
Other methods (`kolmogorov_smirnov`, `wasserstein`) do not use this histogram+density path, so they can stay finite while **JensenShannon** alone becomes null.
## Required behavior
1. **Finite output** for finite `ref` and `cur` after removing non-finite values, whenever both sides have **at least one** usable sample.
2. **Explicit handling of out-of-range chunk mass:** probability mass from `cur` that does not fall into any bin defined from `ref` must still be represented (so the chunk distribution sums to 1), analogous to NannyMLs continuous JS approach (tail / “leftover” mass).
3. **Missing values:** drop `NaN` from `ref` and `cur` before computing. If either side is **empty** after that, return **`float('nan')`** (callers may filter or map; schema may still forbid null — product decision outside this spec).
## Recommended algorithm (replace current body)
1. `ref = np.asarray(ref, float); cur = np.asarray(cur, float)`.
2. `ref = ref[~np.isnan(ref)]; cur = cur[~np.isnan(cur)]`.
3. If `ref.size == 0` or `cur.size == 0`: return `float('nan')`.
4. `hist_ref, edges = np.histogram(ref, bins=bins)`**counts**, not `density=True`.
5. `p = hist_ref.astype(float) / ref.size` (reference bin probabilities).
6. `hist_cur, _ = np.histogram(cur, bins=edges)`; `q = hist_cur.astype(float) / cur.size`.
7. `leftover = 1.0 - float(np.sum(q))`. If `leftover > 1e-15` (tolerance for float noise), append **`leftover`** to `q` and **`0.0`** to `p` so both remain proper discrete distributions over the same extended support.
8. Apply small smoothing (existing module constant `EPSILON` is fine): add `EPSILON` to `p` and `q`, renormalize each to sum 1.
9. `m = 0.5 * (p + q)`; compute symmetric JS via KL terms as today, e.g. `inner = 0.5 * (sum(p*log(p/m)) + sum(q*log(q/m)))`.
10. Return `sqrt(max(inner, 0.0))` to guard against tiny negative `inner` from floating-point error.
## Non-goals / notes
- **Numerical parity** with the old `density=True` implementation is not required; parity with **NannyML** or **scipy** JS is desirable but optional. The priority is **finite, interpretable** drift when the chunk is outside the reference histogram range.
- **Multivariate** drift in the same file is unchanged by this spec.
- **Tests** in `sientia_model` should cover: (a) chunk entirely above reference max, (b) entirely below reference min, (c) overlapping range, (d) `ref` or `cur` all-NaN after cleaning.
## Reference (external)
- NannyML continuous JS uses count-based bin probabilities and a **leftover** mass bin; see `ContinuousJensenShannonDistance` in `nannyml/drift/univariate/methods.py` (`_calculate`, `leftover = 1 - np.sum(...)`).
- Historical NaN-in-reference issue: [NannyML#339](https://github.com/NannyML/nannyml/issues/339) / [#340](https://github.com/NannyML/nannyml/pull/340) (orthogonal to out-of-range mass, but relevant for input cleaning).

File diff suppressed because it is too large Load Diff

View File

@@ -1,109 +0,0 @@
-- =============================================================================
-- E2E test database schema for the ``sientia_data`` namespace.
--
-- Mirrors the production DDL one-to-one so any change in production can be
-- pasted directly into this file. The conftest fixture loads this SQL into the
-- testcontainers Postgres before each test run.
--
-- Notes on differences from production:
-- * Tables that are partitioned in production (e.g. ``simple_metrics``,
-- ``transformed_data``, ``drift_metrics``) are created as plain tables
-- here because the test suite does not exercise partition pruning.
-- * Indexes are intentionally omitted; tests rely on functional behavior,
-- not query plans.
-- =============================================================================
CREATE SCHEMA IF NOT EXISTS sientia_data;
-- -----------------------------------------------------------------------------
-- sientia_data.laborious_data
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.laborious_data (
model_id int4 NOT NULL,
variable text NOT NULL,
value numeric NULL,
"timestamp" timestamptz NOT NULL,
created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL,
CONSTRAINT unique_timestamp_variable
UNIQUE (model_id, "timestamp", variable)
);
-- -----------------------------------------------------------------------------
-- sientia_data.predictions
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.predictions (
model_id int4 NOT NULL,
prediction numeric NULL,
prediction_confidence numeric NOT NULL,
response_time numeric NOT NULL,
prediction_status text NOT NULL,
"timestamp" timestamptz NOT NULL,
created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL,
"comments" text NULL,
CONSTRAINT unique_model_id_timestamp
UNIQUE (model_id, "timestamp")
);
-- -----------------------------------------------------------------------------
-- sientia_data.transformed_data
-- Production: PARTITION BY RANGE (created_at). Tests use a plain table.
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.transformed_data (
id SERIAL NOT NULL,
model_id int4 NOT NULL,
variable text NOT NULL,
value numeric NULL,
"timestamp" timestamptz NOT NULL,
created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL,
PRIMARY KEY (id, created_at)
);
-- -----------------------------------------------------------------------------
-- sientia_data.drift_metrics
-- Production: PARTITION BY RANGE (created_at). Tests use a plain table.
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.drift_metrics (
id SERIAL NOT NULL,
model_id text NOT NULL,
feature text NULL,
method text NOT NULL,
value numeric NOT NULL,
alert bool NOT NULL,
chunk_index int4 NOT NULL,
chunk_start_date text NOT NULL,
chunk_end_date text NOT NULL,
accurate bool NOT NULL,
"timestamp" timestamptz NULL,
created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL,
PRIMARY KEY (id, created_at)
);
-- -----------------------------------------------------------------------------
-- sientia_data.simple_metrics
-- Production: PARTITION BY RANGE (created_at). Tests use a plain table.
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.simple_metrics (
id SERIAL NOT NULL,
model_id text NOT NULL,
metric text NOT NULL,
value numeric NOT NULL,
"timestamp" timestamptz NULL,
data_size int4 NOT NULL,
interval_minutes int4 NOT NULL,
created_at timestamptz DEFAULT CURRENT_TIMESTAMP NOT NULL,
PRIMARY KEY (id, created_at)
);
-- -----------------------------------------------------------------------------
-- sientia_data.log_retrain
-- No primary key in production; all columns nullable.
-- -----------------------------------------------------------------------------
CREATE TABLE IF NOT EXISTS sientia_data.log_retrain (
mlflow_experiment_id int8 NULL,
mlflow_run_id text NULL,
model_id text NULL,
model_name text NULL,
status text NULL,
"timestamp" timestamptz NULL,
"version" text NULL
);

View File

@@ -3,59 +3,13 @@ Shared helpers for E2E tests (Temporal workflows + PostgreSQL).
"""
import asyncio
import json
from datetime import datetime
from decimal import Decimal
from pathlib import Path
from typing import Any
from sqlalchemy import text
from sqlalchemy.engine import Engine
SCENARIO_INPUTS_DIR = Path(__file__).parent / 'scenario_inputs'
def _replace_template_values(payload: Any, model_id: int) -> Any:
"""
Replace string placeholders in scenario payloads with the concrete model id.
Args:
payload: JSON-like structure loaded from scenario input file.
model_id: Model id used to render template placeholders.
Return:
Any: Payload with ``{{MODEL_ID}}`` replaced where applicable.
"""
if isinstance(payload, dict):
return {key: _replace_template_values(value, model_id) for key, value in payload.items()}
if isinstance(payload, list):
return [_replace_template_values(item, model_id) for item in payload]
if isinstance(payload, str):
if payload == '{{MODEL_ID}}':
return model_id
return payload.replace('{{MODEL_ID}}', str(model_id))
return payload
def load_scenario_input(file_name: str, model_id: int | None = None) -> dict[str, Any]:
"""
Load a scenario input JSON from ``e2e/scenario_inputs``.
Args:
file_name: JSON file name inside ``e2e/scenario_inputs``.
model_id: Optional model id used to render ``{{MODEL_ID}}`` placeholders.
Return:
dict[str, Any]: Input payload ready to be passed to workflow/activity calls.
"""
file_path = SCENARIO_INPUTS_DIR / file_name
with file_path.open('r', encoding='utf-8') as f:
payload = json.load(f)
if model_id is not None:
return _replace_template_values(payload, model_id)
return payload
async def start_and_await_workflow(client, workflow_run, input_data: dict, workflow_id: str, timeout: float = 60.0):
"""
@@ -80,17 +34,7 @@ async def start_and_await_workflow(client, workflow_run, input_data: dict, workf
return await asyncio.wait_for(handle.result(), timeout=timeout)
DEFAULT_BATCH_TIMESTAMP = '2024-01-01 12:00:00+00:00'
DEFAULT_PREDICTION_HISTORY_TIMESTAMP = '2024-01-01 12:00:00+00:00'
def insert_sample_data(
postgres_engine: Engine,
model_id: int,
values: list[Any],
*,
data_timestamp: str = DEFAULT_BATCH_TIMESTAMP,
) -> None:
def insert_sample_data(postgres_engine: Engine, model_id: int, values: list[Any]) -> None:
"""
Replace laborious_data rows for a model_id with one row per value (sensor_1..n).
@@ -98,113 +42,29 @@ def insert_sample_data(
postgres_engine: SQLAlchemy engine.
model_id: Model id column value.
values: Per-sensor values; use string 'NULL' for SQL NULL.
data_timestamp: Timestamp and created_at for every inserted row; drives
``last_timestamp`` on the MinIO/query payload (max row time).
"""
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}'))
values_sql = []
for i, value in enumerate(values):
values_sql.append(f"""
({model_id}, 'sensor_{i + 1}', {value}, '{data_timestamp}', '{data_timestamp}')
({model_id}, 'sensor_{i + 1}', {value}, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00')
""")
insert_sql = f"""
INSERT INTO sientia_data.laborious_data (model_id, variable, value, timestamp, created_at)
INSERT INTO predictions_schema.laborious_data (model_id, variable, value, timestamp, created_at)
VALUES
{', '.join(values_sql)}
"""
conn.execute(text(insert_sql))
def insert_sample_prediction(
postgres_engine: Engine,
model_id: int,
*,
prediction_timestamp: str = DEFAULT_PREDICTION_HISTORY_TIMESTAMP,
) -> tuple[int, Decimal, Decimal, str]:
"""
Insert a single historical prediction row for REPEAT scenarios.
Args:
postgres_engine: SQLAlchemy engine.
model_id: Model id.
prediction_timestamp: Row ``timestamp`` (unique with model_id in tests).
Return:
tuple: (model_id, prediction, prediction_confidence, prediction_status) for assertions.
"""
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
insert_sql = f"""
INSERT INTO sientia_data.predictions (
model_id, timestamp, prediction, prediction_confidence, prediction_status, comments, response_time
)
VALUES (
{model_id}, '{prediction_timestamp}', 10, 0, 'Good', '', 0.1
)
"""
conn.execute(text(insert_sql))
return (model_id, Decimal(10), Decimal(0), 'Good')
def workflow_failure_message_chain(exc: BaseException) -> list[str]:
"""
Collect ``str()`` / ``message`` from an exception and its ``__cause__`` chain.
Args:
exc: Root exception (e.g. from ``pytest.raises``).
Return:
list[str]: Messages from root to innermost cause.
"""
messages: list[str] = []
current: BaseException | None = exc
while current is not None:
messages.append(getattr(current, 'message', None) or str(current) or repr(current))
current = current.__cause__
return messages
def assert_postgres_unique_violation_in_chain(exc: BaseException) -> None:
"""
Assert the exception chain mentions Postgres unique-constraint violation.
Args:
exc: Workflow or activity error from Temporal.
Raises:
AssertionError: If no link in the chain looks like UniqueViolation.
"""
chain = ' | '.join(workflow_failure_message_chain(exc))
assert 'UniqueViolation' in chain or 'unique_model_id_timestamp' in chain, (
f'Expected unique constraint violation in error chain, got: {chain}'
)
def assert_prediction_row_count(postgres_engine: Engine, model_id: int, expected: int) -> None:
"""
Assert how many prediction rows exist for a model_id.
Args:
postgres_engine: SQLAlchemy engine.
model_id: Model id filter.
expected: Expected row count.
"""
with postgres_engine.connect() as conn:
n = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = :m'),
{'m': model_id},
).scalar()
assert n == expected, f'Expected {expected} prediction rows, got {n}'
def assert_prediction(
postgres_engine: Engine,
model_id: int,
prediction: float = 0.5,
prediction_confidence: int | Decimal = 0,
prediction_status: str = 'Good',
comments: str = '',
comments: str | None = None,
comments_contains: str | None = None,
) -> None:
"""
@@ -216,14 +76,16 @@ def assert_prediction(
prediction: Expected prediction value.
prediction_confidence: Expected confidence (int or Decimal for numeric column).
prediction_status: Expected status string.
comments: Expected exact comments string (ignored when ``comments_contains`` is set).
comments_contains: When set, assert this substring appears in comments.
comments: Expected exact comments string (optional).
comments_contains: Substring expected in comments when queued (optional).
"""
import pytest
with postgres_engine.connect() as conn:
result_query = conn.execute(
text(
f'SELECT model_id, prediction, prediction_confidence, prediction_status, comments '
f'FROM sientia_data.predictions WHERE model_id = {model_id} '
f'FROM predictions_schema.predictions WHERE model_id = {model_id} '
f'ORDER BY created_at ASC'
)
)
@@ -238,13 +100,12 @@ def assert_prediction(
str(prediction_confidence)
), f'Expected prediction_confidence={prediction_confidence}, got {row[2]}'
assert row[3] == prediction_status, f"Expected prediction_status='{prediction_status}', got {row[3]}"
actual_comments = row[4] or ''
if comments is not None:
assert row[4] == comments, f"Expected comments='{comments}', got {row[4]}"
if comments_contains is not None:
assert comments_contains in actual_comments, (
f"Expected comments to contain '{comments_contains}', got '{actual_comments}'"
assert comments_contains in row[4], (
f"Expected comments to contain '{comments_contains}', got {row[4]}"
)
else:
assert actual_comments == comments, f"Expected comments='{comments}', got '{actual_comments}'"
def assert_continue(
@@ -258,7 +119,7 @@ def assert_continue(
result_query = conn.execute(
text(
f'SELECT model_id, prediction, prediction_confidence, prediction_status, comments '
f'FROM sientia_data.predictions WHERE model_id = {model_id}'
f'FROM predictions_schema.predictions WHERE model_id = {model_id}'
)
)
prediction_rows = result_query.fetchall()
@@ -278,7 +139,7 @@ def assert_stop(postgres_engine: Engine, model_id: int) -> None:
with postgres_engine.connect() as conn:
result_query = conn.execute(
text(f'SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = {model_id}')
text(f'SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = {model_id}')
)
count = result_query.scalar()
assert count == 0, f'Expected no predictions, but found {count} records'
@@ -301,7 +162,7 @@ def assert_repeat(postgres_engine: Engine, model_id: int, last_prediction: tuple
result_query = conn.execute(
text(
f'SELECT model_id, prediction, prediction_confidence, prediction_status '
f'FROM sientia_data.predictions WHERE model_id = {model_id} '
f'FROM predictions_schema.predictions WHERE model_id = {model_id} '
f'ORDER BY created_at ASC'
)
)
@@ -318,51 +179,3 @@ def assert_repeat(postgres_engine: Engine, model_id: int, last_prediction: tuple
def make_workflow_id(prefix: str) -> str:
"""Build a unique workflow id using a prefix and current timestamp."""
return f'{prefix}-{datetime.now().timestamp()}'
def insert_target_data_for_drift(
postgres_engine: Engine,
model_id: int,
timestamps: list[str],
variables_values: dict[str, list[float]],
) -> None:
"""
Insert one row per (timestamp, variable) pair into ``laborious_data``.
Used by drift scenarios that need wide-format input where the pivot keeps a
full row for every timestamp.
Args:
- postgres_engine: SQLAlchemy engine bound to the test container.
- model_id: Model id stamped on every row.
- timestamps: ISO-8601 strings used both as ``timestamp`` and ``created_at``.
- variables_values: Mapping of variable name to a list of values; each list
must be the same length as ``timestamps``.
"""
for var_name, values in variables_values.items():
if len(values) != len(timestamps):
raise ValueError(
f"Variable '{var_name}' has {len(values)} values but {len(timestamps)} timestamps"
)
rows_sql = []
for index, ts in enumerate(timestamps):
for var_name, values in variables_values.items():
rows_sql.append(
f"({model_id}, '{var_name}', {values[index]}, '{ts}', '{ts}')"
)
with postgres_engine.begin() as conn:
conn.execute(
text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}')
)
if rows_sql:
conn.execute(
text(
'INSERT INTO sientia_data.laborious_data '
'(model_id, variable, value, "timestamp", created_at) VALUES '
+ ', '.join(rows_sql)
)
)

View File

@@ -1,14 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"schema": "sientia_data",
"source_table_name": "laborious_data",
"target_table_name": "drift_metrics",
"interval": 60,
"drift_metrics": ["kolmogorov_smirnov", "jensen_shannon", "wasserstein"],
"chunk_period": "min",
"model_config": {
"target": "sensor_1"
}
}

View File

@@ -1,37 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"EMPTY_DATA": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": true,
"prediction_store_policy": "lts:1",
"model_config": {
"retention_minutes": 0,
"target": "sensor_1"
},
"datetime_columns": ["timestamp", "created_at"]
}

View File

@@ -1,45 +0,0 @@
{
"metadata": {
"metadata": {
"model_id": "{{MODEL_ID}}",
"model_name": "test_model",
"schedule_name": "test-schedule",
"workflow_name": "predictions_batch"
}
},
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"EMPTY_DATA": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": true,
"prediction_store_policy": "lts:1",
"model_config": {
"target": "sensor_1",
"retention_minutes": 0
},
"datetime_columns": ["timestamp", "created_at"]
}

View File

@@ -1,45 +0,0 @@
{
"metadata": {
"metadata": {
"model_id": "{{MODEL_ID}}",
"model_name": "test_model",
"schedule_name": "test-schedule",
"workflow_name": "predictions_batch"
}
},
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"EMPTY_DATA": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": true,
"prediction_store_policy": "lts:1",
"model_config": {
"target": "sensor_1",
"retention_minutes": 0
},
"datetime_columns": ["nonexistent_column"]
}

View File

@@ -1,16 +0,0 @@
{
"metadata": {
"metadata": {
"model_id": "{{MODEL_ID}}",
"model_name": "test_model",
"schedule_name": "test-schedule",
"workflow_name": "predictions_batch"
}
},
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data"
}

View File

@@ -1,44 +0,0 @@
{
"metadata": {
"metadata": {
"model_id": "{{MODEL_ID}}",
"model_name": "test_model",
"schedule_name": "test-schedule",
"workflow_name": "predictions_batch"
}
},
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT * FROM nonexistent_table WHERE invalid_syntax =",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"EMPTY_DATA": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": true,
"prediction_store_policy": "lts:1",
"model_config": {
"target": "sensor_1",
"retention_minutes": 0
}
}

View File

@@ -1,12 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "log_retrain",
"datetime_columns": ["timestamp", "created_at"],
"model_config": {
"target": "sensor_1"
}
}

View File

@@ -1,11 +0,0 @@
{
"metadata": {
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"workflow_name": "predictions_batch"
},
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"model_name": "test_model",
"datetime_columns": ["timestamp", "created_at"]
}

View File

@@ -1,37 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"EMPTY_DATA": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": false,
"prediction_store_policy": "lts:1",
"model_config": {
"retention_minutes": 0,
"target": "sensor_1"
},
"datetime_columns": ["timestamp", "created_at"]
}

View File

@@ -1,39 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
"schema": "sientia_data",
"table_name": "predictions",
"transform_table_name": "transformed_data",
"input_filters": {
"SPECIFIC_VARIABLES_NULL_VALUES": {
"POLICY": "CONTINUE",
"CONFIG": {
"variables": ["sensor_1"]
}
}
},
"mlflow_transform_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"mlflow_predict_filters": {
"API_ERROR": {
"POLICY": "STOP",
"CONFIG": {}
}
},
"path_priority": ["STOP", "CONTINUE", "REPEAT"],
"opc_output_config": {},
"pi_web_api_output_config": {},
"save_transform": true,
"prediction_store_policy": "lts:1",
"model_config": {
"retention_minutes": 0,
"target": "sensor_1"
},
"datetime_columns": ["timestamp", "created_at"]
}

View File

@@ -1,14 +0,0 @@
{
"schedule_name": "test-schedule",
"model_name": "test_model",
"model_id": "{{MODEL_ID}}",
"schema": "sientia_data",
"predictions_table_name": "predictions",
"data_table_name": "laborious_data",
"target_table_name": "simple_metrics",
"interval_minutes": 60,
"metrics": ["rmse", "mse", "mae", "r2"],
"model_config": {
"target": "sensor_target"
}
}

View File

@@ -1,455 +1,564 @@
# E2E Scenario Documentation - Predictions Batch
# Test Scenarios for Predictions Batch Workflow
This document describes the end-to-end scenarios for `predictions_batch` and its child workflows:
`prediction_process` and `format_and_export_prediction`.
This document describes all possible test scenarios for the `predictions_batch` workflow and its child workflows `prediction_process` and `format_and_export_prediction`.
It is a functional reference of scenario behavior, inputs, and expected outcomes.
## Running automated E2E tests (`e2e/`)
## Execution Context
- **Runtime**: Docker (or a Docker-compatible daemon) must be available so [testcontainers](https://testcontainers.com/) can start **PostgreSQL** and **MinIO** containers.
- **Dependencies**: install dev requirements (includes `testcontainers[postgres,minio]`).
- **Invocation**: run only integration-marked tests, for example: `pytest e2e/ -m integration`.
- **MinIO tests**: `e2e/test_minio_offload.py` exercises real S3 uploads; other E2E modules continue to mock MinIO on the worker used by most scenarios.
- **OPC tests (real server)**: `e2e/test_opc_real_server.py` uses an in-process **asyncua** server and real `OpcRepository` (`test_activities_real_opc`). Scenarios 3.1.2, 3.2.2, 3.2.4, and 3.2.5 are covered there. Other E2E modules keep the OPC mock.
- Run only OPC real-server tests: `pytest e2e/test_opc_real_server.py -m "integration and opc"`.
- Tests run under `e2e/` and are marked with `@pytest.mark.integration`.
- PostgreSQL and MinIO are provisioned with testcontainers.
- `test_minio_offload.py` uses real MinIO I/O; other scenario suites may use stubs/mocks for optional outputs.
- Real OPC UA scenarios use `@pytest.mark.opc` and an in-process asyncua server (`e2e/test_opc_real_server.py`).
## Workflow Overview
### Local validation
Use the existing project virtualenv and the shared `validate` script for unit/quality gates; run E2E separately (Docker required).
```bash
source ./venv/bin/activate
# Auto-fix + static checks (no pytest)
validate --fix --project-name=laborious
# Full unit + quality gate
validate --project-name=laborious
# E2E (integration)
pytest e2e/ --override-ini testpaths=e2e -m integration
# E2E (real OPC server only)
pytest e2e/test_opc_real_server.py --override-ini testpaths=e2e -m opc
```
The `predictions_batch` workflow:
1. Loads data using a custom SQL query
2. Prepares prediction configuration
3. Delegates to `prediction_process` child workflow which:
- Retrieves last timestamp for incremental processing
- Applies input data quality gates
- Executes MLFlow transform operation
- Validates transform response
- Executes MLFlow predict operation
- Validates predict response
- Delegates to `format_and_export_prediction` child workflow
4. The `format_and_export_prediction` workflow:
- Formats prediction data (normal or default)
- Exports to PI Web API (optional)
- Exports to OPC server (optional)
- Exports to PostgreSQL
- Writes metrics
---
## 1. Main Workflow Scenarios
Source: `e2e/test_predictions_batch_main_workflow.py`
## 1. Predictions Batch - Main Workflow Scenarios
### 1.1.1 Happy Path - Complete Success
**Summary**: Full workflow succeeds with valid query and default gate behavior.
### 1.1 Success Scenarios
**Description**:
- Query returns rows for a model.
- `prediction_process` runs transform and predict paths.
- Final prediction and transformed data are persisted.
#### Scenario 1.1.1: Happy Path - Complete Success
**Description**: Workflow completes successfully with valid SQL query and all activities succeed
**Expected Outcome**:
- Exactly one prediction row is created.
- Transform rows are created.
- Confidence/status/comments are success values.
**Input**:
- Valid `schedule_name`, `model_name`, `model_id`
- Valid `query` returning non-empty DataFrame
- Valid `schema`, `table_name`, `transform_table_name`
- Optional `datetime_columns` for timestamp parsing
- Optional `input_filters`, `mlflow_transform_filters`, `mlflow_predict_filters`
- Optional `path_priority`, `opc_output_config`, `pi_web_api_output_config`
### 1.2.1 SQL Query Execution Error
**Summary**: Invalid SQL leads to no persisted prediction.
**Expected Behavior**:
- `load_custom_query` returns DataFrame with data
- Workflow prepares prediction input with all configurations
- `prediction_process` child workflow executes successfully
- All gates pass with no issues
- Transform and predict operations succeed
- Data exported to PostgreSQL
- Metrics written
**Description**:
- Input query is invalid.
- Load step fails and workflow follows error/short-circuit path.
**Expected Outcome**:
- No prediction rows for the model.
- Workflow does not require retry-loop assumptions in assertions.
### 1.2.2 Missing Required Parameters
**Summary**: Missing required fields prevent workflow completion path.
**Description**:
- Required input key (e.g. `query`) is omitted.
- Workflow fails to produce actionable input for child flow.
**Expected Outcome**:
- No prediction rows are persisted.
- Workflow handle may require explicit terminate in E2E harness.
### 1.2.3 Invalid Datetime Column Specification (de-prioritized)
**Summary**: Legacy invalid datetime-column case is retained only as low-priority legacy coverage.
**Description**:
- `datetime_columns` references non-existing columns.
- Behavior may vary by query shape and parser fallback.
**Expected Outcome**:
- No predictions persisted in the covered legacy assertion path.
- Scenario is not considered primary behavior coverage.
**Assertions**:
- SQL query executed once
- `prediction_process` workflow called with correct parameters
- Data exists in PostgreSQL (predictions table)
- Metrics recorded
- No errors raised
---
## 2. Prediction Process Scenarios
Source: `e2e/test_predictions_batch_prediction_process.py`
### 1.2 Error Scenarios
### 2.1 Input Gate Path Decisions
#### Scenario 1.2.1: SQL Query Execution Error
**Description**: SQL query fails due to syntax error or connection issue
#### 2.1.1 CONTINUE
**Summary**: Input filter flags quality issue but allows continuation via default path.
**Input**:
- Invalid SQL query (syntax error)
- Or database connection unavailable
**Description**:
- Input gate returns `CONTINUE`.
- MLFlow transform/predict are skipped.
- Export path persists default-style prediction with warning context.
**Expected Behavior**:
- `load_custom_query` raises exception (caught by Temporal retry policy)
- Notification sent with SQL error details
- After retries, activity may return empty data or workflow may fail
- If empty data returned, workflow completes with early exit via input gate
#### 2.1.2 STOP
**Summary**: Input filter blocks processing.
**Description**:
- Input gate returns `STOP`.
- Workflow exits without export.
#### 2.1.3 REPEAT with history
**Summary**: Prior prediction is reused.
**Description**:
- Input gate returns `REPEAT`.
- `repeat_last_prediction` path is executed using existing historical row.
#### 2.1.4 REPEAT without history
**Summary**: Repeat requested but no previous prediction exists.
**Description**:
- Input gate returns `REPEAT`.
- No prior row is available to duplicate.
**Expected Outcome**:
- No new prediction rows are created for the model.
### 2.2 Transform Gate Decisions
#### 2.2.1 CONTINUE on transform response error
**Summary**: Transform response is degraded, but workflow continues.
#### 2.2.2 STOP on transform response error
**Summary**: Transform response error blocks downstream processing.
#### 2.2.3 REPEAT on transform response error
**Summary**: Transform response error triggers repeat-last-prediction path.
#### 2.2.4 STOP on transform content NaN
**Summary**: Content gate (`NAN_VALUES`) blocks on all-NaN transform payload.
### 2.3 Predict Gate Decisions
#### 2.3.1 CONTINUE on predict response error
**Summary**: Predict response degraded; workflow exports with degraded metadata.
#### 2.3.2 STOP on predict response error
**Summary**: Predict response error blocks export.
#### 2.3.3 REPEAT on predict response error
**Summary**: Predict response error routes to repeat-last-prediction.
### 2.4.1 Priority Conflict Resolution
**Summary**: Deterministic selection when multiple filters produce different flags.
**Description**:
- Multiple filters may produce `STOP`, `CONTINUE`, and/or `REPEAT`.
- `path_priority` defines precedence.
**Expected Outcome**:
- Highest-priority flag is applied consistently.
- Executed branch matches configured priority ordering.
**Assertions**:
- Error notification sent
- Workflow completes (either fails or exits early)
- No data in predictions table
---
## 3. Format and Export Scenarios
Source: `e2e/test_predictions_batch_format_export.py`
#### Scenario 1.2.2: Missing Required Parameters
**Description**: Essential parameters missing from input
### 3.1 Output Combination Scenarios
**Input**:
- Missing `query` or `model_id` or `schema` or `table_name`
#### 3.1.1 Default prediction export
**Summary**: Non-`None` path flag uses `format_default_prediction`.
**Expected Behavior**:
- Workflow or activity raises KeyError or validation error
- Workflow fails immediately
**Description**:
- Default prediction is generated.
- Transform export is skipped.
- Optional outputs (PI/OPC) still execute when configured.
#### 3.1.2 OPC only
**Summary**: Postgres + OPC writes, PI Web API disabled.
#### 3.1.3 PI Web API only
**Summary**: Postgres + PI writes, OPC disabled.
#### 3.1.4 Postgres only
**Summary**: Both optional outputs disabled; only Postgres persistence and metrics.
#### 3.1.5 No transformed data export
**Summary**: Prediction is persisted; transformed table is not written.
### 3.2 Degraded-but-successful Completion
#### 3.2.1 PI Web API write error
**Summary**: PI write failure does not fail workflow.
**Expected Outcome**:
- Workflow completes.
- Prediction persisted with degraded confidence/comments (PI error semantics).
#### 3.2.2 OPC write error
**Summary**: OPC write failure does not fail workflow.
**Expected Outcome**:
- Workflow completes.
- Prediction persisted with OPC degraded confidence/comments.
#### 3.2.3 PI Web API partial write error
**Summary**: Partial PI acknowledgement is treated as degraded success.
**Expected Outcome**:
- Workflow completes.
- Prediction persisted with PI error confidence and descriptive comment.
#### 3.2.4 OPC session / channel error (confidence 14)
**Summary**: Tier-1 `BadSessionIdInvalid` (or equivalent session error) degrades the prediction without failing the workflow.
**Sources**:
- Mock: `e2e/test_predictions_batch_format_export.py::test_scenario_3_2_4_opc_session_bad_mock`
- Real server: `e2e/test_opc_real_server.py::test_scenario_3_2_4_opc_session_bad_real_server` (`@pytest.mark.opc`)
**Expected Outcome**:
- Workflow completes.
- `prediction_confidence` is 14.
- Comments contain `OPC UA session/channel error: BadSessionIdInvalid`.
#### 3.2.5 OPC write blocked during reconnect (confidence 14)
**Summary**: While reconnect holds the repository connection lock, writes fail fast with `reconnect_in_progress`.
**Sources**:
- Mock: `e2e/test_predictions_batch_format_export.py::test_scenario_3_2_5_opc_reconnect_in_progress_mock`
- Real server: `e2e/test_opc_real_server.py::test_scenario_3_2_5_opc_write_blocked_during_reconnect_real_server` (`@pytest.mark.opc`)
**Expected Outcome**:
- Workflow completes.
- `prediction_confidence` is 14.
- Comments contain `OPC UA reconnect in progress`.
### 3.3.1 Combined Optional Outputs (PI + OPC)
**Summary**: Both external output channels are enabled together.
**Description**:
- PI Web API and OPC configs are both present.
- Output mutation order matters for final persisted payload.
**Expected Outcome**:
- PI write executes before OPC write in workflow sequence.
- Final Postgres payload reflects any confidence/comment updates.
- OPC metrics are emitted when tag writes return response times.
**Assertions**:
- Workflow fails with parameter error
- Error notification sent
- No child workflow called
---
## 4. MinIO Offload Scenarios
Source: `e2e/test_minio_offload.py`
#### Scenario 1.2.3: Invalid Datetime Column Specification
**Description**: Datetime column specified doesn't exist in query results
### 4.1.1 Forced offload to MinIO
**Summary**: Very low threshold forces parquet upload.
**Input**:
- `datetime_columns: ['nonexistent_column']`
- Query results don't have this column
**Description**:
- Payload is offloaded (`object_key` present, inline data absent/empty).
- Object is present in MinIO under `prediction_datasets/...`.
- Retrieval reconstructs the dataframe.
**Expected Behavior**:
- `load_custom_query` may raise KeyError or warning
- Depending on implementation, workflow may fail or continue
- Error notification sent
### 4.1.2 Full workflow with offloaded load payload
**Summary**: Offload path works during full `predictions_batch` execution.
**Expected Outcome**:
- Workflow completes.
- Prediction row is persisted.
### 4.2.1 Inline payload below threshold
**Summary**: Data remains inline when threshold is not exceeded.
**Expected Outcome**:
- Payload stores inline `data`.
- `object_key` is `None`.
- Downstream persistence behavior matches offload scenario semantics.
**Assertions**:
- Error raised or warning logged
- Workflow behavior depends on error handling policy
---
## 5. Drift Workflow Scenarios
Source: `e2e/test_drift.py`
## 2. Prediction Process - Child Workflow Scenarios
The drift suite drives the **real** `sientia_model.analytics.drift_analysis.DriftAnalysis`
analyzer (no stubs / mocks). Each scenario exercises the full pipeline:
### 2.1 Input gate Early Exit Scenarios
```
laborious_data (Postgres) -> load_custom_query
-> calculate_drift (DriftAnalysis univariate + multivariate)
-> export_data_to_postgres (sientia_data.drift_metrics)
```
#### Scenario 2.1.1: Input Gate Triggers CONTINUE
**Description**: Input gate determines data should use previous prediction
The `mlflow_repository_stub` provides the reference-data CSV via
`download_artifacts`, and tests assert postgres rows in
`sientia_data.drift_metrics` against this canonical schema:
**Input**:
- Data that should continue with input data as prediction
- `input_filters` configured with `POLICY: 'CONTINUE'`
- `path_priority` includes CONTINUE
`id, model_id, feature, method, value, alert, chunk_index, chunk_start_date, chunk_end_date, accurate, timestamp, created_at`.
**Expected Behavior**:
- `input_gate` returns `path_flag='CONTINUE'`
- `path_flag_handler` calls export workflow with input data directly
- MLFlow transform and predict skipped
- Data exported as-is
Tests assert behavioral / structural properties (column presence, NOT NULL
constraints, business-key invariants like uniform `timestamp` and stamped
`model_id`) rather than exact numeric drift scores, since those depend on
the real analyzer implementation and the synthetic data fed in.
**Assertions**:
- `input_gate` called
- MLFlow operations NOT called
- Export workflow called with original data
- Workflow completes
### 5.1 Happy paths
#### D.1.1 Full pipeline persists all columns with reference data
**Summary**: 10 minutes of target data are inserted; a 10-row reference CSV
is configured via the MLflow stub. The `DriftAnalysis` runs end-to-end.
#### Scenario 2.1.2: Input Gate Triggers STOP
**Description**: Input data quality gate fails with STOP policy
**Expected Outcome**:
- One row per `(chunk_index, feature, method)` plus a `multivariate` block
per chunk is persisted.
- Every column in the DDL is populated; `feature` is the only nullable column
per the new schema.
- `accurate=True` for every row (reference path).
- All three default univariate methods reach the analyzer.
- `model_id` is stamped as `text` and uniform across rows.
- `timestamp` equals `max(target_data.timestamp)` and is uniform across rows.
- `chunk_start_date` / `chunk_end_date` are persisted as ISO text and ordered.
- `p_value` is dropped before persistence.
**Input**:
- Data with EMPTY_DATA or other critical issues
- `input_filters` configured with `POLICY: 'STOP'`
#### D.1.2 30% fallback when reference data is unavailable
**Summary**: MLflow alias resolution is forced to fail so
`get_reference_data` returns `None`; `calculate_drift` falls back to the
first 30% of target rows as reference.
**Expected Behavior**:
- `input_gate` returns `path_flag='STOP'`
- `path_flag_handler` detects STOP
- Workflow returns early without calling MLFlow
- No prediction exported
**Expected Outcome**:
- All persisted rows carry `accurate=False`.
- A `MODEL_METRICS_REFERENCE_DATA_WARNING` notification is emitted to MongoDB.
**Assertions**:
- `input_gate` called
- `path_flag_handler` returns True (early exit)
- MLFlow transform NOT called
- Export workflow NOT called
- Workflow completes without error
### 5.2 Failure paths
#### D.3.1 Empty target data short-circuits the workflow
**Summary**: `load_custom_query` returns no rows.
#### Scenario 2.1.3: Input Gate Triggers REPEAT
**Description**: Input gate determines data should repeat last prediction
**Expected Outcome**:
- The workflow returns early and writes nothing to `sientia_data.drift_metrics`.
**Input**:
- Data with quality issues that require using previous prediction
- `input_filters` configured with `POLICY: 'REPEAT'`
- `path_priority` includes REPEAT
### 5.3 Configuration paths
**Expected Behavior**:
- `input_gate` returns `path_flag='REPEAT'`
- `path_flag_handler` calls `repeat_last_prediction` activity
- MLFlow transform and predict skipped
- Last prediction repeated and exported
#### D.4.2 Invalid `chunk_period` raises ValueError
**Summary**: Anything other than `min` / `s` is rejected by `calculate_drift`.
**Expected Outcome**:
- The workflow surfaces the `ValueError` ("Invalid chunk period: ...").
- No rows are persisted.
#### D.4.3 `chunk_period='s'` preserves seconds in `chunk_start_date`
**Summary**: Target data spans two minutes with samples at second-30
boundaries; the activity is configured with `chunk_period='s'`.
**Expected Outcome**:
- At least one persisted `chunk_start_date` carries `seconds=30`, proving
that the analyzer chunked at sub-minute granularity and the ISO-text
serialization preserved the boundary.
**Assertions**:
- `input_gate` called
- MLFlow operations NOT called
- `repeat_last_prediction` activity called
- Workflow completes
---
## 6. Simple Metrics Workflow Scenarios
Source: `e2e/test_simple_metrics.py`
### 2.2 Transform gate Early Exit Scenarios
Validates `sientia_data.simple_metrics` columns:
`id, model_id, metric, value, timestamp, data_size, interval_minutes, created_at`.
Note: ``timestamp`` is now nullable per the new DDL and ``model_id`` is ``text``.
#### Scenario 2.2.1: Transform Gate Triggers CONTINUE
**Description**: Transform response gate determines data should continue despite issues
### 6.1 Happy paths
**Input**:
- Valid input data
- Transform response has quality issues but policy is CONTINUE
- `mlflow_transform_filters` configured with `POLICY: 'CONTINUE'`
- `path_priority` includes CONTINUE
#### S.1.1 rmse/mse/mae/r2 happy path
**Summary**: Prediction/target pairs are inserted; the activity computes all
four metrics with closed-form expected values.
**Expected Behavior**:
- `request_transform` succeeds
- `mlflow_response_gate` for transform returns `path_flag='CONTINUE'`
- `path_flag_handler` calls export workflow with transform data
- MLFlow predict skipped
- Transform data exported as-is
**Expected Outcome**:
- One row per metric is persisted; all columns populated.
- `data_size` matches the joined row count and `interval_minutes=60`.
#### S.1.2 Subset metrics
**Summary**: Requesting `metrics=['rmse']` writes only the rmse row.
### 6.2 Edge cases
#### S.2.1 Zero-variance target returns r2=0
**Summary**: When all targets are equal, `ss_tot=0`; the activity must guard
against division by zero and return `r2=0`.
### 6.3 Failure paths
#### S.3.1 No overlapping data short-circuits persistence
**Summary**: With no `laborious_data` rows for the configured target variable
the workflow exits before `calculate_simple_metrics` and writes nothing.
**Assertions**:
- Transform completed
- `mlflow_response_gate` called for transform
- MLFlow predict NOT called
- Export workflow called with transform data
- Workflow completes
---
## 7. Minimal Retrain Workflow Scenarios
Source: `e2e/test_minimal_retrain.py`
#### Scenario 2.2.2: Transform Gate Triggers STOP
**Description**: Transform response validation fails with STOP policy
The MLflow registry is fully mocked (no real artifacts in test container).
Validates `sientia_data.log_retrain` columns:
`mlflow_experiment_id, mlflow_run_id, model_id, model_name, status, timestamp, version`.
Note: the new DDL drops the legacy ``id`` and ``created_at`` columns,
``mlflow_experiment_id`` is now ``int8`` and ``model_id`` is ``text``.
**Input**:
- Valid input data
- Transform response has critical errors
- `mlflow_transform_filters` configured with `POLICY: 'STOP'`
### 7.1 Happy path
**Expected Behavior**:
- `request_transform` succeeds but response invalid
- `mlflow_response_gate` for transform returns `path_flag='STOP'`
- Workflow exits without calling predict or export
#### MR.1.1 Successful retrain + promotion
**Summary**: Training data loads via MinIO offload, `wrapper.retrain` succeeds,
the new version is promoted to the `production` alias.
**Expected Outcome**:
- Report row has success status, `version='7'`, `mlflow_run_id='retrain-run-id'`,
`mlflow_experiment_id=4242` (`int8`).
- `mlflow.log_artifact` is called with the input CSV.
- `promote_to_alias` is called once with the resolved version and alias.
### 7.2 Failure paths
#### MR.2.1 Wrapper retrain raises
**Summary**: `wrapper.retrain` raises `RuntimeError`. The activity returns
`success=False`, `update_production_model` is NOT invoked.
**Expected Outcome**:
- Report row carries the error message and `version`/`mlflow_*` columns are NULL.
#### MR.2.2 Missing `model_config.target`
**Summary**: Empty model config short-circuits before any MLflow call.
**Expected Outcome**:
- Report row carries the explicit guard message.
- `get_cached_model` is never invoked.
#### MR.3.1 No training data
**Summary**: The training query returns no rows; the workflow does not
persist any report row. The current code raises plain `ValueError` from the
workflow function, which Temporal treats as a workflow-task failure (see
`CODE_ISSUES.md` issue MR-1).
**Assertions**:
- Transform completed but validation failed
- `mlflow_response_gate` called for transform
- MLFlow predict NOT called
- Export workflow NOT called
- Workflow completes without error
---
## Input Contract Reference
#### Scenario 2.2.3: Transform Gate Triggers REPEAT
**Description**: Transform response gate determines data should repeat last prediction
Common scenario input fields:
- `schedule_name`
- `model_name`
- `model_id`
- `query`
- `schema`
- `table_name`
- `transform_table_name`
- `input_filters`
- `mlflow_transform_filters`
- `mlflow_predict_filters`
- `path_priority` (default order: `STOP`, `CONTINUE`, `REPEAT`)
- `save_transform`
- `prediction_store_policy`
- `model_config.target`
- `datetime_columns` (when query returns temporal fields)
**Input**:
- Valid input data
- Transform response has quality issues that require using previous prediction
- `mlflow_transform_filters` configured with `POLICY: 'REPEAT'`
- `path_priority` includes REPEAT
Optional outputs:
- `opc_output_config`
- `pi_web_api_output_config`
**Expected Behavior**:
- `request_transform` succeeds but response has issues
- `mlflow_response_gate` for transform returns `path_flag='REPEAT'`
- `path_flag_handler` calls `repeat_last_prediction` activity
- MLFlow predict skipped
- Last prediction repeated and exported
**Assertions**:
- Transform completed but validation triggered REPEAT
- `mlflow_response_gate` called for transform
- MLFlow predict NOT called
- `repeat_last_prediction` activity called
- Workflow completes
---
### 2.3 Predict gate Early Exit Scenarios
#### Scenario 2.3.1: Predict Gate Triggers CONTINUE
**Description**: Predict response gate determines data should continue despite issues
**Input**:
- Valid input and transform data
- Predict response has quality issues but policy is CONTINUE
- `mlflow_predict_filters` configured with `POLICY: 'CONTINUE'`
- `path_priority` includes CONTINUE
**Expected Behavior**:
- `request_predict` succeeds
- `mlflow_response_gate` for predict returns `path_flag='CONTINUE'`
- `path_flag_handler` calls export workflow with predict data
- Prediction exported despite quality issues
**Assertions**:
- Transform and predict completed
- `mlflow_response_gate` called for predict
- Export workflow called with predict data
- Workflow completes
---
#### Scenario 2.3.2: Predict Gate Triggers STOP
**Description**: Prediction validation fails with STOP policy
**Input**:
- Valid input and transform
- Predict response has critical errors
- `mlflow_predict_filters` configured with `POLICY: 'STOP'`
**Expected Behavior**:
- `request_predict` succeeds but response invalid
- `mlflow_response_gate` for predict returns `path_flag='STOP'`
- Workflow exits without export
**Assertions**:
- Transform completed
- Predict completed but validation failed
- Export workflow NOT called
- Workflow completes without error
---
#### Scenario 2.3.3: Predict Gate Triggers REPEAT
**Description**: Predict response gate determines data should repeat last prediction
**Input**:
- Valid input and transform data
- Predict response has quality issues that require using previous prediction
- `mlflow_predict_filters` configured with `POLICY: 'REPEAT'`
- `path_priority` includes REPEAT
**Expected Behavior**:
- `request_predict` succeeds but response has issues
- `mlflow_response_gate` for predict returns `path_flag='REPEAT'`
- `path_flag_handler` calls `repeat_last_prediction` activity
- Last prediction repeated and exported
**Assertions**:
- Transform and predict completed but validation triggered REPEAT
- `mlflow_response_gate` called for predict
- `repeat_last_prediction` activity called
- Export workflow NOT called with current prediction
- Workflow completes
---
## 3. Format and Export Prediction - Child Workflow Scenarios
### 3.1 Success Scenarios
#### Scenario 3.1.1: Default Prediction Export
**Description**: Error prediction path creates default prediction
**Input**:
- `path_flag: 'ERROR'` or other non-None value (not STOP/CONTINUE/REPEAT)
- `comment` provided with error details
**Expected Behavior**:
- `format_default_prediction` called instead of `format_prediction`
- Default prediction created with error metadata
- Exported to PostgreSQL only
- Transformed data NOT processed
- Metrics written
**Assertions**:
- `format_default_prediction` called
- `format_prediction` NOT called
- `format_transformed_data` NOT called
- One PostgreSQL export only
- Default values in prediction data
- Comment included
---
#### Scenario 3.1.2: Export with OPC only
**Description**: Export to PostgreSQL and OPC server only (no PI Web API)
**Input**:
- `path_flag: None`
- `opc_output_config` configured with valid OPC settings
- `pi_web_api_output_config: None` or `{}`
**Expected Behavior**:
- Normal formatting
- PostgreSQL export executed
- OPC export executed
- PI Web API activity skipped
- Metrics written with OPC metrics
**Assertions**:
- PI Web API activity NOT called
- OPC activity called
- PostgreSQL export called
- Metrics written with `opc_metrics` populated
---
#### Scenario 3.1.3: Export with PI Web API only
**Description**: Export to PostgreSQL and PI Web API only (no OPC)
**Input**:
- `path_flag: None`
- `pi_web_api_output_config` configured with valid PI Web API settings
- `opc_output_config: None` or `{}`
**Expected Behavior**:
- Normal formatting
- PostgreSQL export executed
- PI Web API export executed
- OPC activity skipped
- Metrics written without OPC metrics
**Assertions**:
- OPC activity NOT called
- PI Web API activity called
- PostgreSQL export called
- Metrics written with empty `opc_metrics`
---
#### Scenario 3.1.4: Export Without Optional Outputs
**Description**: Export only to PostgreSQL (no OPC or PI Web API)
**Input**:
- `path_flag: None`
- `opc_output_config: None` or `{}`
- `pi_web_api_output_config: None` or `{}`
**Expected Behavior**:
- Normal formatting
- Only PostgreSQL export executed
- OPC and PI Web API activities skipped
- Metrics written without OPC metrics
**Assertions**:
- PI Web API activity NOT called
- OPC activity NOT called
- PostgreSQL export called
- Metrics written with empty `opc_metrics`
---
#### Scenario 3.1.5: Export Without Transformed Data
**Description**: Only prediction exported, no transform table
**Input**:
- `path_flag: None`
- `transformed_data: None` or `save_transform: False`
- `opc_output_config: None` or `{}`
- `pi_web_api_output_config: None` or `{}`
**Expected Behavior**:
- Only prediction formatted and exported
- Transform export skipped
- Single PostgreSQL write
**Assertions**:
- `format_transformed_data` NOT called
- One PostgreSQL export
- Transform table remains empty
---
### 3.2 Error Scenarios
These paths do **not** rely on Temporal activity retries for export failures: the write activities run once, errors are handled inside the activity, and the **workflow completes successfully** with degraded metadata on the persisted prediction (`prediction_confidence` and `comments`).
#### Scenario 3.2.1: PI Web API Write Error
**Description**: PI Web API export fails
**Input**:
- Valid prediction
- PI Web API service unavailable or invalid config
**Expected Behavior**:
- `write_pi_web_api_data` surfaces the failure (exception handled in the activity layer)
- Notification may be sent
- Workflow **completes** (does not fail)
- Prediction row is still written to PostgreSQL with error confidence **13** and a comment describing the PI error
- Subsequent steps (e.g. OPC, Postgres) still run per workflow order with the updated prediction payload
**Assertions**:
- PI Web API error notification sent (when applicable)
- Workflow completes
- PostgreSQL contains the prediction with `prediction_confidence` 13 and expected `comments`
---
#### Scenario 3.2.2: OPC Write Error
**Description**: OPC server write fails
**Input**:
- Valid prediction
- OPC server unavailable or invalid configuration
**Expected Behavior**:
- `write_opc_data` reports failure without aborting the workflow
- Notification may be sent
- Workflow **completes** (does not fail)
- Prediction row is written to PostgreSQL with OPC error confidence **12** and a comment indicating OPC write issues
**Assertions**:
- OPC error notification sent (when applicable)
- Workflow completes
- PostgreSQL contains the prediction with `prediction_confidence` 12 and expected `comments`
---
#### Scenario 3.2.4: OPC Session / Channel Bad* (Tier-1)
**Description**: OPC write fails with a Tier-1 session or channel status (e.g. `BadSessionIdInvalid`) while transport may still appear open on the client
**Input**:
- Valid prediction and OPC output config
- Mock or server returning Tier-1 `UaStatusCodeError` on write (no write retry in the same activity)
**Expected Behavior**:
- `write_opc_data` fails forward for affected tags; background reconnect may be scheduled if `OPC_RECONNECTION_INTERVAL` allows
- Workflow **completes**
- PostgreSQL row uses **`prediction_confidence` 14** and comment prefix `OPC UA session/channel error:` (including OPC status name)
- `opc_write_attempts_total` records `result=BadSessionIdInvalid` (or matching status); no second write attempt in the same activity
**Assertions**:
- Workflow completes
- `prediction_confidence = 14`
- `comments` matches `OPC UA session/channel error:%`
- Generic OPC error confidence **12** is not used for this case
**Reference**: [docs/opc-communication.md](../docs/opc-communication.md), plan `.cursor/plans/opc_bad_reconnect_ac4c6045.plan.md`
---
#### Scenario 3.2.5: OPC Write Blocked During Reconnect
**Description**: A write is attempted while the repository is reconnecting (session not ready)
**Input**:
- Valid prediction
- Simulated slow reconnect (e.g. delayed `connect`) or concurrent writes where the first triggers reconnect
**Expected Behavior**:
- Second write (or parallel write) is rejected **immediately** when reconnect is in progress or `_session_ready` is cleared — **without** calling `write_value`
- No wait/sleep on the write path; no duplicate `connect` from parallel writers (connection lock)
- `prediction_confidence = 14`, `comments = OPC UA reconnect in progress` (distinguish from Tier-1 `Bad*` via comment prefix in SQL)
**Assertions**:
- At most one reconnect sequence (`disconnect` + `connect`) for the overlapping window
- No write retry after failure
- Tests in `test_opc_repository` (unit) and optional e2e in `test_predictions_batch_format_export.py`
---
#### Scenario 3.2.3: PI Web API Partial Write Error
**Description**: Two prediction tags attempt to be written to PI Web API, but only one succeeds
**Input**:
- Valid prediction
- Two prediction tags configured
- PI Web API returns partial success (one tag succeeds, one fails)
**Expected Behavior**:
- `write_pi_web_api_data` processes response
- `process_pi_web_api_response` detects partial failure
- Error confidence set (13)
- Notification sent for failed tag
- Workflow completes with error confidence (single activity attempt; no retry loop)
**Assertions**:
- One tag written successfully
- One tag failed
- Error confidence set in prediction
- Error notification sent
- Workflow completes
---

View File

@@ -26,7 +26,7 @@ async def test_format_and_export_prediction_default_path_e2e(
client = temporal_test_env.client
model_id = 401
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}'))
metadata = {
'metadata': {
@@ -44,7 +44,7 @@ async def test_format_and_export_prediction_default_path_e2e(
'timestamp': '2024-01-01 12:00:00+00:00',
'model_id': model_id,
'model_name': 'test_model',
'schema': 'sientia_data',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'comment': 'e2e child workflow default path',
@@ -64,7 +64,7 @@ async def test_format_and_export_prediction_default_path_e2e(
row = conn.execute(
text(
f'SELECT prediction, prediction_confidence, prediction_status, comments '
f'FROM sientia_data.predictions WHERE model_id = {model_id}'
f'FROM predictions_schema.predictions WHERE model_id = {model_id}'
)
).fetchone()
assert row is not None

View File

@@ -1,600 +0,0 @@
"""
End-to-end tests for the Drift workflow.
The drift suite drives the **real** ``sientia_model.analytics.drift_analysis.DriftAnalysis``
analyzer (no mocking). Each scenario exercises the full pipeline:
laborious_data (Postgres)
-> load_custom_query
-> calculate_drift (DriftAnalysis univariate + multivariate)
-> export_data_to_postgres (sientia_data.drift_metrics)
Coverage focus:
- Happy path persists every column required by ``sientia_data.drift_metrics``
with a valid reference dataset downloaded from MLflow.
- 30% fallback path activates when the MLflow reference is unavailable and
emits the ``MODEL_METRICS_REFERENCE_DATA_WARNING`` notification.
- Empty target data short-circuits the workflow without persisting anything.
- Invalid ``chunk_period`` is rejected by ``calculate_drift``.
- ``chunk_period='s'`` preserves second-level precision in
``chunk_start_date``.
Tests assert behavioral / structural properties (column presence, NOT NULL
constraints, business-key invariants) rather than exact numeric values, since
those depend on the real analyzer implementation and synthetic data.
"""
from datetime import datetime, timedelta, timezone
from pathlib import Path
from unittest.mock import MagicMock, patch
import pandas as pd
import pytest
from sqlalchemy import text
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from e2e.helpers import (
insert_target_data_for_drift,
load_scenario_input,
make_workflow_id,
start_and_await_workflow,
)
from laborious.activities.activities import Activities
from laborious.activities.model_metrics import ModelMetrics
from laborious.workflows.drift import Drift
from sientia_model.analytics.drift_analysis import DriftAnalysis
# Drift columns persisted on every row in ``sientia_data.drift_metrics`` —
# mirrors the production DDL.
EXPECTED_DRIFT_COLUMNS = [
'id',
'model_id',
'feature',
'method',
'value',
'alert',
'chunk_index',
'chunk_start_date',
'chunk_end_date',
'accurate',
'timestamp',
'created_at',
]
# Columns the DDL marks as NOT NULL. ``feature`` and ``timestamp`` are
# nullable in the production schema (multivariate rows do not bind to a
# single feature; ``timestamp`` is allowed to be empty when upstream data has
# no usable instant).
NON_NULL_DRIFT_COLUMNS = {
'id',
'model_id',
'method',
'value',
'alert',
'chunk_index',
'chunk_start_date',
'chunk_end_date',
'accurate',
'created_at',
}
DEFAULT_DRIFT_METHODS = ['kolmogorov_smirnov', 'jensen_shannon', 'wasserstein']
def _chunk_dataframe_skip_empty_groups(
self: DriftAnalysis,
df: pd.DataFrame,
timestamp_col: str,
chunk_period: str,
) -> list[tuple[int, pd.DataFrame]]:
"""
Same as ``DriftAnalysis._chunk_dataframe`` but omit empty time buckets.
``pd.Grouper(freq='s')`` yields every second between min and max timestamp;
empty buckets still appear in the groupby iterator and produce invalid
drift rows (e.g. NaT timestamps) that ``calculate_drift`` later filters out
entirely. Production fix belongs in ``sientia_model``; this shim keeps the
e2e honest about second-level chunk boundaries with sparse samples.
"""
grouped = df.groupby(pd.Grouper(key=timestamp_col, freq=chunk_period), dropna=True)
chunks: list[tuple[int, pd.DataFrame]] = []
idx = 0
for _, chunk in grouped:
if chunk.empty:
continue
chunks.append((idx, chunk.copy()))
idx += 1
return chunks
def _drift_input(model_id: int, **overrides) -> dict:
"""Load the base drift scenario JSON and apply ad-hoc overrides."""
input_data = load_scenario_input('drift_base.json', model_id=model_id)
input_data.update(overrides)
return input_data
def _recent_minute_timestamps(count: int, offset_minutes: int = 6) -> list[str]:
"""
Build ``count`` consecutive UTC minute timestamps placed in the recent past.
The Drift workflow filters target rows with ``timestamp > NOW() - INTERVAL``,
so timestamps must be recent for tests to retrieve any data. Snapping to
minute precision keeps the helper deterministic regardless of clock skew.
Args:
- count (int): How many consecutive minute timestamps to generate.
- offset_minutes (int): Minutes ago for the EARLIEST generated timestamp.
Return:
list[str]: ISO strings with ``+0000`` offset, one per minute.
"""
base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(
minutes=offset_minutes
)
return [
(base + timedelta(minutes=i)).strftime('%Y-%m-%d %H:%M:%S%z') for i in range(count)
]
def _configure_reference_csv(mlflow_repository_stub, reference_rows: pd.DataFrame) -> None:
"""
Wire ``mlflow_repository_stub`` so ``get_reference_data`` returns
``reference_rows`` by writing them to ``dst_path/retrain_input.csv``.
Args:
- mlflow_repository_stub: External MLflow repository fixture.
- reference_rows (pd.DataFrame): Rows to expose as the production reference.
"""
def _download(run_id: str, artifact_path: str, dst_path: str, metadata=None):
target = Path(dst_path) / artifact_path
target.parent.mkdir(parents=True, exist_ok=True)
reference_rows.to_csv(target, index=False)
mlflow_repository_stub._client.get_model_version_by_alias.return_value = MagicMock(
run_id='fake-reference-run'
)
file_info = MagicMock()
file_info.path = 'retrain_input.csv'
mlflow_repository_stub._client.list_artifacts.return_value = [file_info]
mlflow_repository_stub.download_artifacts.side_effect = _download
def _force_reference_unavailable(mlflow_repository_stub) -> None:
"""Make ``get_reference_data`` return ``None`` by failing alias resolution."""
mlflow_repository_stub._client.get_model_version_by_alias.side_effect = Exception(
'no production alias registered'
)
def _select_drift_rows(postgres_engine, model_id: int) -> list[dict]:
"""Read every persisted drift row for ``model_id`` ordered by chunk/feature/method."""
with postgres_engine.connect() as conn:
rows = (
conn.execute(
text(
'SELECT * FROM sientia_data.drift_metrics '
'WHERE model_id = :m '
'ORDER BY chunk_index, feature, method'
),
{'m': str(model_id)},
)
.mappings()
.all()
)
return [dict(row) for row in rows]
def _assert_required_columns_populated(rows: list[dict]) -> None:
"""Validate column presence and NOT NULL constraints on every row."""
assert rows, 'expected at least one drift row to be persisted'
seen_columns = set(rows[0].keys())
for column in EXPECTED_DRIFT_COLUMNS:
assert column in seen_columns, f'Missing drift column in postgres: {column}'
for row in rows:
for column in NON_NULL_DRIFT_COLUMNS:
assert row[column] is not None, f"Column '{column}' is NULL in {row}"
assert 'p_value' not in row, 'p_value must not be persisted to drift_metrics'
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_happy_path_persists_all_columns_with_reference_data(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario D.1.1: Happy path with reference data downloaded from MLflow.
Drives the full pipeline against the real ``DriftAnalysis``. Asserts:
- One row is persisted per ``(chunk_index, feature, method)`` combination
plus the multivariate row block, with every column required by
``sientia_data.drift_metrics`` populated.
- The three default univariate methods are forwarded to the analyzer.
- ``model_id`` and ``timestamp`` are stamped by the activity (not by the
analyzer); ``timestamp`` equals ``max(target_data.timestamp)`` and is
identical on every persisted row.
- ``chunk_start_date`` / ``chunk_end_date`` are persisted as ISO text so
the analyzer's nanosecond-precision boundaries survive the ``text``
column type.
- ``accurate=True`` because the reference dataset was available.
"""
client = temporal_test_env.client
model_id = 411
target_timestamps = _recent_minute_timestamps(count=10)
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': [10.0 + i * 0.1 for i in range(10)],
'sensor_2': [20.0 + i * 0.5 for i in range(10)],
},
)
reference_df = pd.DataFrame(
{
'timestamp': [
f'2023-12-31 11:{minute:02d}:00+00:00' for minute in range(10)
],
'sensor_1': [9.0 + i * 0.05 for i in range(10)],
'sensor_2': [18.0 + i * 0.25 for i in range(10)],
}
)
_configure_reference_csv(mlflow_repository_stub, reference_df)
input_data = _drift_input(model_id)
await start_and_await_workflow(
client, Drift.run, input_data, make_workflow_id('test-drift-happy-path')
)
rows = _select_drift_rows(postgres_engine, model_id)
exported_csv_path = '/tmp/test_drift_happy_path_exported.csv'
pd.DataFrame(rows).to_csv(exported_csv_path, index=False)
print(
f'\n[test_drift_happy_path] Exported drift dataframe '
f'({len(rows)} rows) -> {exported_csv_path}'
)
_assert_required_columns_populated(rows)
# The activity drops the target column from the feature list, so only
# ``sensor_2`` participates in univariate analysis (``sensor_1`` is the
# configured target). Multivariate produces one row per chunk regardless.
univariate_rows = [r for r in rows if r['feature'] != 'multivariate']
multivariate_rows = [r for r in rows if r['feature'] == 'multivariate']
assert univariate_rows, 'expected univariate drift rows for non-target features'
assert multivariate_rows, 'expected one multivariate drift row per chunk'
# All three default methods must reach the analyzer.
assert {r['method'] for r in univariate_rows} == set(DEFAULT_DRIFT_METHODS)
assert all(r['method'] == 'multivariate' for r in multivariate_rows)
assert {r['feature'] for r in univariate_rows} == {'sensor_2'}
# ``timestamp`` is stamped uniformly with ``max(target_data.timestamp)``.
expected_timestamp = pd.to_datetime(max(target_timestamps), utc=True)
persisted_timestamps = {pd.to_datetime(r['timestamp'], utc=True) for r in rows}
assert len(persisted_timestamps) == 1, (
'timestamp must be uniform across all drift rows '
f'(got {len(persisted_timestamps)} distinct values)'
)
assert pd.Timestamp(persisted_timestamps.pop()) == expected_timestamp, (
'timestamp must equal max(target_data.timestamp)'
)
# ``model_id`` is stamped by ``calculate_drift`` (not produced by the analyzer).
assert all(r['model_id'] == str(model_id) for r in rows), (
'model_id must be stamped on every drift row'
)
# Reference path → accurate=True.
assert all(r['accurate'] is True for r in rows), (
'reference path should mark all rows as accurate'
)
# ISO text serialization preserves ordering between start/end of each chunk.
for row in rows:
assert 'T' in row['chunk_start_date'], (
f"chunk_start_date should be ISO text, got {row['chunk_start_date']!r}"
)
assert 'T' in row['chunk_end_date'], (
f"chunk_end_date should be ISO text, got {row['chunk_end_date']!r}"
)
assert row['chunk_start_date'] <= row['chunk_end_date'], (
f'chunk_start_date must precede chunk_end_date '
f"(start={row['chunk_start_date']}, end={row['chunk_end_date']})"
)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_uses_30pct_fallback_when_reference_unavailable(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
notification_inserts,
):
"""
Scenario D.1.2: ``get_reference_data`` returns ``None`` (production alias
missing), so ``calculate_drift`` falls back to using the first 30% of
target rows as reference. Persisted rows must report ``accurate=False``
and a ``MODEL_METRICS_REFERENCE_DATA_WARNING`` notification must be
emitted to mongo.
"""
client = temporal_test_env.client
model_id = 412
target_timestamps = _recent_minute_timestamps(count=10)
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': [10.0 + i * 0.1 for i in range(10)],
'sensor_2': [20.0 + i * 0.5 for i in range(10)],
},
)
_force_reference_unavailable(mlflow_repository_stub)
input_data = _drift_input(model_id)
await start_and_await_workflow(
client, Drift.run, input_data, make_workflow_id('test-drift-fallback')
)
rows = _select_drift_rows(postgres_engine, model_id)
_assert_required_columns_populated(rows)
assert all(r['accurate'] is False for r in rows), (
'fallback path must mark all rows as inaccurate'
)
fallback_warnings = [
call
for call in notification_inserts.call_args_list
if call.args
and isinstance(call.args[0], dict)
and call.args[0].get('notification_id') == 'MODEL_METRICS_REFERENCE_DATA_WARNING'
]
assert len(fallback_warnings) >= 1, 'expected reference fallback warning notification'
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_empty_target_data_short_circuits_workflow(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario D.3.1: When ``load_custom_query`` returns no rows the workflow
must return early without invoking the analyzer or writing any drift rows.
"""
client = temporal_test_env.client
model_id = 431
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
_force_reference_unavailable(mlflow_repository_stub)
input_data = _drift_input(model_id)
await start_and_await_workflow(
client, Drift.run, input_data, make_workflow_id('test-drift-empty-target')
)
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.drift_metrics WHERE model_id = :m'),
{'m': str(model_id)},
).scalar()
assert count == 0
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_invalid_chunk_period_raises_value_error(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario D.4.2: ``calculate_drift`` validates ``chunk_period`` and rejects
anything other than ``min`` / ``s``. The workflow must surface the
``ValueError`` and persist nothing.
"""
client = temporal_test_env.client
model_id = 442
target_timestamps = _recent_minute_timestamps(count=5)
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': [1.0, 2.0, 3.0, 4.0, 5.0],
'sensor_2': [10.0, 20.0, 30.0, 40.0, 50.0],
},
)
_force_reference_unavailable(mlflow_repository_stub)
input_data = _drift_input(model_id, chunk_period='hour')
with pytest.raises(Exception) as excinfo:
await start_and_await_workflow(
client,
Drift.run,
input_data,
make_workflow_id('test-drift-bad-chunk-period'),
)
# Temporal wraps the activity ValueError in WorkflowFailureError; the
# message may live on ``.message`` or ``str(exc)`` depending on the SDK
# error class, so walk the cause chain looking for the guard text.
cause_descriptions = []
current: BaseException | None = excinfo.value
while current is not None:
cause_descriptions.append(
getattr(current, 'message', None) or str(current) or repr(current)
)
current = current.__cause__
assert any('Invalid chunk period' in msg for msg in cause_descriptions), (
f'Expected ValueError about chunk period in chain, got: {cause_descriptions}'
)
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.drift_metrics WHERE model_id = :m'),
{'m': str(model_id)},
).scalar()
assert count == 0
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_empty_merge_skips_export_without_insufficient_notification(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario D.4.3a: When the analyzer returns an empty merged frame (no metric rows),
``calculate_drift`` yields ``[]``; the workflow skips export. Real insufficient-data
cases are signaled by ``DriftInsufficientDataError`` inside ``sientia_model``, not by
empty output alone.
"""
client = temporal_test_env.client
model_id = 444
target_timestamps = _recent_minute_timestamps(count=5)
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': [1.0 + i * 0.1 for i in range(5)],
'sensor_2': [10.0 + i * 0.5 for i in range(5)],
},
)
_force_reference_unavailable(mlflow_repository_stub)
input_data = _drift_input(model_id, chunk_period='min')
empty_merge = pd.DataFrame(
columns=[
'timestamp',
'feature',
'method',
'value',
'alert',
'chunk_index',
'chunk_start_date',
'chunk_end_date',
'threshold',
'drift_type',
]
)
with patch.object(ModelMetrics, 'get_drift_metrics', return_value=empty_merge):
await start_and_await_workflow(
client,
Drift.run,
input_data,
make_workflow_id('test-drift-empty-merge'),
)
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.drift_metrics WHERE model_id = :m'),
{'m': str(model_id)},
).scalar()
assert count == 0
@pytest.mark.asyncio
@pytest.mark.integration
async def test_drift_chunk_period_seconds_sufficient_data_preserves_seconds_in_chunk_start_date(
temporal_test_env: WorkflowEnvironment,
temporal_worker_drift: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario D.4.3b: With enough sub-minute samples and ``chunk_period='s'``, drift rows
persist and ``chunk_start_date`` keeps second-level precision (incl. second=30).
``DriftAnalysis._chunk_dataframe`` is patched to skip empty ``pd.Grouper(freq='s')``
buckets so sparse seconds between samples do not flood the pipeline with NaT rows;
the durable fix belongs in ``sientia_model``.
"""
client = temporal_test_env.client
model_id = 443
base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(minutes=6)
target_timestamps = []
sensor_1_vals = []
sensor_2_vals = []
for minute_offset in range(6):
t0 = base + timedelta(minutes=minute_offset)
t1 = t0 + timedelta(seconds=30)
target_timestamps.append(t0.strftime('%Y-%m-%d %H:%M:%S%z'))
target_timestamps.append(t1.strftime('%Y-%m-%d %H:%M:%S%z'))
v0 = 1.0 + minute_offset * 0.1
v1 = v0 + 0.05
sensor_1_vals.extend([v0, v1])
sensor_2_vals.extend([10.0 + v0, 10.0 + v1])
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': sensor_1_vals,
'sensor_2': sensor_2_vals,
},
)
_force_reference_unavailable(mlflow_repository_stub)
input_data = _drift_input(model_id, chunk_period='s')
with patch.object(DriftAnalysis, '_chunk_dataframe', _chunk_dataframe_skip_empty_groups):
await start_and_await_workflow(
client,
Drift.run,
input_data,
make_workflow_id('test-drift-chunk-seconds-sufficient'),
)
with postgres_engine.connect() as conn:
rows = (
conn.execute(
text(
'SELECT chunk_start_date FROM sientia_data.drift_metrics '
'WHERE model_id = :m ORDER BY chunk_index'
),
{'m': str(model_id)},
)
.mappings()
.all()
)
assert rows, 'expected at least one drift row to be persisted'
seconds_present = {pd.Timestamp(r['chunk_start_date']).second for r in rows}
assert 30 in seconds_present, (
f'expected at least one chunk_start_date with seconds=30, got {seconds_present}'
)

View File

@@ -1,430 +0,0 @@
"""
End-to-end tests for the MinimalRetrain workflow.
The MLflow registry is fully stubbed because no real artifacts exist in a
test container; we only validate that the workflow:
- Loads training data via ``load_query_with_minio_offload``.
- Calls ``retrain_model`` with a payload pointing at MinIO.
- Calls ``update_production_model`` only when retrain succeeds.
- Persists ``sientia_data.log_retrain`` rows with all required columns;
success rows carry the new ``version`` / ``mlflow_run_id`` /
``mlflow_experiment_id`` while failure rows leave them ``NULL``.
The production DDL drops the legacy ``id`` / ``created_at`` columns and
moves ``mlflow_experiment_id`` to ``int8`` and ``model_id`` to ``text``.
The stubs used here therefore emit ``experiment_id`` as an integer to fit
the new column type.
"""
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, patch
import pytest
import pandas as pd
from sqlalchemy import text
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from e2e.helpers import (
insert_target_data_for_drift,
load_scenario_input,
make_workflow_id,
start_and_await_workflow,
)
from laborious.activities.activities import Activities
from laborious.workflows.minimal_retrain import MinimalRetrain
from sientia_model.wrappers.sientia_model import SientiaModel
# Columns defined by the production DDL for ``sientia_data.log_retrain``.
# The legacy ``retrain_reports`` table had ``id`` and ``created_at``; the new
# DDL drops both. ``mlflow_experiment_id`` is ``int8`` and ``model_id`` is
# ``text``.
EXPECTED_RETRAIN_REPORT_COLUMNS = [
'mlflow_experiment_id',
'mlflow_run_id',
'model_id',
'model_name',
'status',
'timestamp',
'version',
]
# Matches ``retrain_model`` return ``message`` when ``success`` is True (also written to ``log_retrain.status``).
RETRAIN_ACTIVITY_SUCCESS_MESSAGE = 'Model retrained successfully.'
class _FakeSientiaModelForMinimalRetrain(SientiaModel):
"""
Fake SientiaModel that uses the real SientiaModel lifecycle to surface
index-alignment issues during ``retrain()``.
It intentionally performs strict alignment inside ``_retrain_model``:
``y.loc[x.index]``.
"""
def __init__(self, *, target: str = 'sensor_1'):
super().__init__(
model_type='FakeMinimalRetrain',
model_version='0.0.0',
model=object(),
transformer=object(),
)
self.target = target
self.model_is_fitted = True
self.force_retrain_error = False
def store_model( # type: ignore[override]
self,
name: str,
signature=None,
pip_requirements=None,
code_path=None,
) -> None:
# No-op: E2E tests validate workflow persistence, not real MLflow artifacts.
return None
def _predict(self, data: pd.DataFrame):
pred = pd.DataFrame({'prediction': [0.5] * len(data)}, index=data.index)
return pred, {}
def _transform(self, data: pd.DataFrame):
out = data.drop(columns=[self.target], errors='ignore').copy()
out.index = data.index
return out, {}
def _train_transformer(self, train_data: pd.DataFrame, val_data: pd.DataFrame) -> None:
return None
def _train_model(
self,
x: pd.DataFrame,
y: pd.DataFrame,
x_val: pd.DataFrame | None = None,
y_val: pd.DataFrame | None = None,
) -> None:
return None
def _retrain_transformer(self, data: pd.DataFrame) -> None:
return None
def _retrain_model(self, x: pd.DataFrame, y: pd.DataFrame | None) -> None:
if self.force_retrain_error:
raise RuntimeError('training did not converge')
if y is None:
return
# Strict alignment on purpose to reproduce the production failure mode.
_ = y.loc[x.index]
@pytest.fixture
def mlflow_repository_stub():
"""
Override the shared E2E fixture: return a real fake ``SientiaModel`` wrapper
instead of a MagicMock wrapper.
"""
repo = MagicMock()
repo._client = MagicMock()
wrapper = _FakeSientiaModelForMinimalRetrain(target='sensor_1')
repo.get_cached_model = MagicMock(return_value=wrapper)
return repo
def _retrain_input(model_id: int, **overrides) -> dict:
"""Load and override the minimal-retrain base scenario."""
payload = load_scenario_input('minimal_retrain_base.json', model_id=model_id)
payload.update(overrides)
return payload
def _seed_retrain_training_rows(postgres_engine, model_id: int) -> None:
"""
Insert training rows in long format that pivot cleanly into
``index=timestamp`` / ``columns={sensor_1, sensor_2}`` for ``retrain_model``.
"""
target_timestamps = [
(datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(minutes=10 - i))
.strftime('%Y-%m-%d %H:%M:%S%z')
for i in range(5)
]
insert_target_data_for_drift(
postgres_engine,
model_id=model_id,
timestamps=target_timestamps,
variables_values={
'sensor_1': [10.0, 11.0, 12.0, 13.0, 14.0],
'sensor_2': [20.0, 21.0, 22.0, 23.0, 24.0],
},
)
def _configure_retrain_happy_path(mlflow_repository_stub) -> None:
"""
Wire ``mlflow_repository_stub`` so retrain + update_production succeed.
Mocks (in order of consumption):
- ``_client.get_model_version_by_alias``: returns ``mv`` with a stable
``run_id`` (used as ``source_run_id``).
- ``start_run``: returns a context manager yielding a ``run_info`` with
run/experiment ids.
- ``log_params``: inert.
- ``_client.search_model_versions``: returns one registry entry whose
``version`` is promoted by ``update_production_model``.
- ``promote_to_alias``: inert success.
"""
mv_src = MagicMock()
mv_src.run_id = 'source-run-id'
new_version = MagicMock()
new_version.version = '7'
new_version.run_id = 'retrain-run-id'
mlflow_repository_stub._client.get_model_version_by_alias.return_value = mv_src
@contextmanager
def fake_start_run(**kwargs):
run_info = MagicMock()
run_info.run_id = 'retrain-run-id'
# ``mlflow_experiment_id`` is ``int8`` in the new DDL, so we feed an
# integer-compatible id from the stubbed run info.
run_info.experiment_id = 4242
yield run_info
mlflow_repository_stub.start_run.side_effect = fake_start_run
mlflow_repository_stub.log_params = MagicMock(return_value=None)
mlflow_repository_stub._client.search_model_versions.return_value = [new_version]
mlflow_repository_stub.promote_to_alias = MagicMock(return_value=None)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_minimal_retrain_happy_path_writes_success_report(
temporal_test_env: WorkflowEnvironment,
temporal_worker_minimal_retrain: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario MR.1.1: Retrain succeeds. ``sientia_data.log_retrain`` must
contain a success row with version/mlflow_run_id/mlflow_experiment_id
populated and the registry must have been told to promote the new version
to the configured alias.
"""
client = temporal_test_env.client
model_id = 711
_seed_retrain_training_rows(postgres_engine, model_id)
_configure_retrain_happy_path(mlflow_repository_stub)
input_data = _retrain_input(model_id)
with patch('laborious.activities.mlflow.mlflow.log_artifact') as log_artifact_mock:
await start_and_await_workflow(
client,
MinimalRetrain.run,
input_data,
make_workflow_id('test-retrain-happy'),
)
with postgres_engine.connect() as conn:
rows = (
conn.execute(
text(
'SELECT * FROM sientia_data.log_retrain '
'WHERE model_id = :m'
),
{'m': str(model_id)},
)
.mappings()
.all()
)
assert len(rows) == 1
row = rows[0]
for column in EXPECTED_RETRAIN_REPORT_COLUMNS:
assert column in row, f'Missing log_retrain column: {column}'
assert row['status'] == RETRAIN_ACTIVITY_SUCCESS_MESSAGE, (
"Expected retrain_model to return success (experiment_response['success'] is True). "
'Persisted log_retrain.status is the activity message; when success is False the run '
'never reaches mlflow.log_artifact — diagnose the retrain failure from status below, '
'not from a skipped artifact upload. '
f"Got status={row['status']!r}, version={row.get('version')!r}, "
f"mlflow_run_id={row.get('mlflow_run_id')!r}."
)
assert log_artifact_mock.called, (
'After a successful retrain, retrain_model must call mlflow.log_artifact for the '
'input CSV inside start_run.'
)
# ``model_id`` is now ``text``; compare against the stringified id.
assert row['model_id'] == str(model_id)
assert row['model_name'] == 'test_model'
assert row['version'] == '7'
assert row['mlflow_run_id'] == 'retrain-run-id'
# ``mlflow_experiment_id`` is now ``int8``; assert the integer value
# provided by the stubbed run info.
assert row['mlflow_experiment_id'] == 4242
assert row['timestamp'] is not None
mlflow_repository_stub.promote_to_alias.assert_called_once()
promote_kwargs = mlflow_repository_stub.promote_to_alias.call_args.kwargs
assert promote_kwargs['model_name'] == 'test_model'
assert promote_kwargs['version'] == '7'
assert promote_kwargs['alias'] == 'production'
@pytest.mark.asyncio
@pytest.mark.integration
async def test_minimal_retrain_failure_writes_report_without_version_columns(
temporal_test_env: WorkflowEnvironment,
temporal_worker_minimal_retrain: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario MR.2.1: ``wrapper.retrain`` raises. The activity must catch the
error, return ``success=False`` so ``update_production_model`` is skipped,
and ``format_retrain_report`` must produce a row with the error message
and NULL version columns.
"""
client = temporal_test_env.client
model_id = 721
_seed_retrain_training_rows(postgres_engine, model_id)
_configure_retrain_happy_path(mlflow_repository_stub)
mlflow_repository_stub.get_cached_model.return_value.force_retrain_error = True
input_data = _retrain_input(model_id)
with patch('laborious.activities.mlflow.mlflow.log_artifact'):
await start_and_await_workflow(
client,
MinimalRetrain.run,
input_data,
make_workflow_id('test-retrain-failure'),
)
with postgres_engine.connect() as conn:
rows = (
conn.execute(
text(
'SELECT * FROM sientia_data.log_retrain '
'WHERE model_id = :m'
),
{'m': str(model_id)},
)
.mappings()
.all()
)
assert len(rows) == 1
row = rows[0]
assert row['model_id'] == str(model_id)
assert row['model_name'] == 'test_model'
assert 'training did not converge' in row['status']
assert row['version'] is None
assert row['mlflow_run_id'] is None
assert row['mlflow_experiment_id'] is None
mlflow_repository_stub.promote_to_alias.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.integration
async def test_minimal_retrain_missing_target_writes_failure_report(
temporal_test_env: WorkflowEnvironment,
temporal_worker_minimal_retrain: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario MR.2.2: ``model_config`` does not declare ``target``. The retrain
activity must short-circuit before any MLflow call and the report row must
carry the explicit guard message.
"""
client = temporal_test_env.client
model_id = 722
_seed_retrain_training_rows(postgres_engine, model_id)
_configure_retrain_happy_path(mlflow_repository_stub)
input_data = _retrain_input(model_id, model_config={})
with patch('laborious.activities.mlflow.mlflow.log_artifact'):
await start_and_await_workflow(
client,
MinimalRetrain.run,
input_data,
make_workflow_id('test-retrain-missing-target'),
)
with postgres_engine.connect() as conn:
row = (
conn.execute(
text(
'SELECT * FROM sientia_data.log_retrain '
'WHERE model_id = :m'
),
{'m': str(model_id)},
)
.mappings()
.first()
)
assert row is not None
assert 'target' in row['status'].lower(), (
f"expected target-missing message, got status={row['status']!r}"
)
assert row['version'] is None
mlflow_repository_stub.get_cached_model.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.integration
async def test_minimal_retrain_no_training_data_does_not_persist_report(
temporal_test_env: WorkflowEnvironment,
temporal_worker_minimal_retrain: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
Scenario MR.3.1: When the training query returns no rows the workflow must
not persist any report row. The workflow currently raises plain
``ValueError`` which Temporal treats as a workflow-task failure (causing
indefinite retries until the test environment times out), so the assertion
here is constrained to the persistence side-effect. See ``e2e/CODE_ISSUES.md``
issue MR-1 for the recommended ``ApplicationError`` fix.
"""
client = temporal_test_env.client
model_id = 731
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
_configure_retrain_happy_path(mlflow_repository_stub)
input_data = _retrain_input(model_id)
with pytest.raises(Exception), patch('laborious.activities.mlflow.mlflow.log_artifact'):
await start_and_await_workflow(
client,
MinimalRetrain.run,
input_data,
make_workflow_id('test-retrain-no-data'),
)
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.log_retrain WHERE model_id = :m'),
{'m': str(model_id)},
).scalar()
assert count == 0

View File

@@ -9,12 +9,7 @@ from sqlalchemy import text
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from e2e.helpers import (
insert_sample_data,
load_scenario_input,
make_workflow_id,
start_and_await_workflow,
)
from e2e.helpers import insert_sample_data, make_workflow_id, start_and_await_workflow
from laborious.activities.activities import Activities
from laborious.utils.models import minio_dataframe_payload as mdp
from laborious.workflows.predictions_batch import PredictionsBatch
@@ -34,17 +29,33 @@ async def test_load_query_with_minio_offload_writes_object_to_bucket(
"""
model_id = 501
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}'))
insert_sample_data(postgres_engine, model_id, [1.0, 2.0])
scenario_input = load_scenario_input('minio_offload_load_query.json', model_id=model_id)
metadata = {'metadata': scenario_input['metadata']}
metadata = {
'metadata': {
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': model_id,
'workflow_name': 'predictions_batch',
}
}
with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1):
payload = test_activities_real_minio.load_query_with_minio_offload(scenario_input)
payload = await test_activities_real_minio.load_query_with_minio_offload(
{
**metadata,
'query': (
'SELECT timestamp, variable, value, created_at '
f'FROM predictions_schema.laborious_data WHERE model_id = {model_id}'
),
'model_name': 'test_model',
'datetime_columns': ['timestamp', 'created_at'],
}
)
assert payload.object_key, 'offloaded payload must reference a MinIO object'
assert payload.data is None or payload.data == {}, 'large payloads should not inline tabular dict'
df = payload.retrieve(test_activities_real_minio.minio_repository, metadata['metadata'])
df = await payload.retrieve(test_activities_real_minio.minio_repository, metadata['metadata'])
assert len(df) >= 1
client = minio_container.get_client()
@@ -66,12 +77,37 @@ async def test_predictions_batch_with_minio_offload_path(
"""
model_id = 502
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM sientia_data.transformed_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.transformed_data WHERE model_id = {model_id}'))
insert_sample_data(postgres_engine, model_id, [10.0, 20.0, 30.0])
input_data = load_scenario_input('minio_offload_workflow.json', model_id=model_id)
input_data = {
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': model_id,
'query': (
'SELECT timestamp, variable, value, created_at '
f'FROM predictions_schema.laborious_data WHERE model_id = {model_id}'
),
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}},
'mlflow_transform_filters': {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}},
'mlflow_predict_filters': {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
'datetime_columns': ['timestamp', 'created_at'],
}
with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1):
await start_and_await_workflow(
@@ -83,26 +119,6 @@ async def test_predictions_batch_with_minio_offload_path(
with postgres_engine.connect() as conn:
count = conn.execute(
text(f'SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = {model_id}')
text(f'SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = {model_id}')
).scalar()
assert count == 1
@pytest.mark.asyncio
@pytest.mark.integration
async def test_load_query_with_inline_payload_when_below_threshold(
postgres_engine,
test_activities_real_minio: Activities,
):
"""Scenario 4.2.1: payload stays inline when threshold is high enough."""
model_id = 503
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
insert_sample_data(postgres_engine, model_id, [1.0, 2.0])
scenario_input = load_scenario_input('minio_offload_load_query.json', model_id=model_id)
with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 10**9):
payload = test_activities_real_minio.load_query_with_minio_offload(scenario_input)
assert payload.object_key is None
assert payload.data is not None

View File

@@ -6,8 +6,6 @@ Mock-based OPC tests remain in test_predictions_batch_format_export.py.
"""
import asyncio
import threading
import time
import pytest
from temporalio.testing import WorkflowEnvironment
@@ -22,7 +20,7 @@ from laborious.utils.repository.opc_repository import OpcRepository
from laborious.workflows.predictions_batch import PredictionsBatch
def _slow_reconnect_under_lock(repo: OpcRepository, hold_seconds: float = 0.75) -> None:
async def _slow_reconnect_under_lock(repo: OpcRepository, hold_seconds: float = 0.75) -> None:
"""
Hold the connection lock briefly so concurrent writes see reconnect_in_progress.
@@ -30,9 +28,9 @@ def _slow_reconnect_under_lock(repo: OpcRepository, hold_seconds: float = 0.75)
repo (OpcRepository): Connected repository.
hold_seconds (float): Time to keep the lock before reconnecting.
"""
with repo._connection_lock:
time.sleep(hold_seconds)
repo._reconnect_locked()
async with repo._connection_lock:
await asyncio.sleep(hold_seconds)
await repo._reconnect_locked()
@pytest.mark.asyncio
@@ -172,12 +170,7 @@ async def test_scenario_3_2_5_opc_write_blocked_during_reconnect_real_server(
repo = test_activities_real_opc.opc_repository['1']
repo._session_ready.clear()
reconnect_thread = threading.Thread(
target=_slow_reconnect_under_lock,
args=(repo,),
daemon=True,
)
reconnect_thread.start()
reconnect_task = asyncio.create_task(_slow_reconnect_under_lock(repo))
input_data = get_base_input_data(model_id)
input_data['opc_output_config'] = build_opc_output_config(opc_e2e_server.node_ids)
@@ -191,7 +184,7 @@ async def test_scenario_3_2_5_opc_write_blocked_during_reconnect_real_server(
make_workflow_id('test-opc-real-reconnect-block'),
)
finally:
reconnect_thread.join(timeout=5.0)
await reconnect_task
assert_prediction(
postgres_engine,

View File

@@ -4,7 +4,7 @@ End-to-end tests for PredictionsBatch workflow - Format and Export scenarios.
from decimal import Decimal
from typing import Any, cast
from unittest.mock import call
from unittest.mock import ANY, AsyncMock, call
import pytest
from sientia_do.notifications.models import NotificationLevel
@@ -12,18 +12,48 @@ from sqlalchemy import text
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from e2e.helpers import (
assert_prediction,
insert_sample_data,
load_scenario_input,
make_workflow_id,
start_and_await_workflow,
)
from e2e.helpers import assert_prediction, insert_sample_data, make_workflow_id, start_and_await_workflow
from laborious.activities.activities import Activities
from laborious.workflows.predictions_batch import PredictionsBatch
base_input_data = {
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 301,
'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 301',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {
'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_transform_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_predict_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
'datetime_columns': ['timestamp', 'created_at'],
}
base_query = "SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {model_id}"
def get_base_input_data(model_id):
return load_scenario_input('format_export_base.json', model_id=model_id)
return {
**base_input_data,
'model_id': model_id,
'query': base_query.format(model_id=model_id),
}
@pytest.mark.asyncio
@@ -49,8 +79,8 @@ async def test_scenario_3_1_1_default_prediction_export(
model_id = 311
with postgres_engine.begin() as conn:
conn.execute(text(f"DELETE FROM sientia_data.predictions WHERE model_id = {model_id}"))
conn.execute(text(f"DELETE FROM sientia_data.transformed_data WHERE model_id = {model_id}"))
conn.execute(text(f"DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}"))
conn.execute(text(f"DELETE FROM predictions_schema.transformed_data WHERE model_id = {model_id}"))
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
input_data = get_base_input_data(model_id)
@@ -146,7 +176,7 @@ async def test_scenario_3_1_1_default_prediction_export(
with postgres_engine.connect() as conn:
tf_count = conn.execute(
text(f"SELECT COUNT(*) FROM sientia_data.transformed_data WHERE model_id = {model_id}")
text(f"SELECT COUNT(*) FROM predictions_schema.transformed_data WHERE model_id = {model_id}")
).scalar()
assert tf_count == 0, 'transform export must be skipped when path_flag is set'
@@ -413,7 +443,7 @@ async def test_scenario_3_1_5_export_without_transformed_data(
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
with postgres_engine.begin() as conn:
conn.execute(text(f"DELETE FROM sientia_data.transformed_data WHERE model_id = {model_id}"))
conn.execute(text(f"DELETE FROM predictions_schema.transformed_data WHERE model_id = {model_id}"))
input_data = get_base_input_data(model_id)
input_data['save_transform'] = False # Don't save transformed data
@@ -503,7 +533,7 @@ async def test_scenario_3_1_5_export_without_transformed_data(
with postgres_engine.connect() as conn:
result_query = conn.execute(
text(f"SELECT COUNT(*) FROM sientia_data.transformed_data WHERE model_id = {model_id}")
text(f"SELECT COUNT(*) FROM predictions_schema.transformed_data WHERE model_id = {model_id}")
)
count = result_query.scalar()
assert count == 0, f"Expected transform table to be empty, but found {count} records"
@@ -634,16 +664,19 @@ async def test_scenario_3_2_2_opc_write_error(
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_3_2_4_opc_session_bad_mock(
async def test_scenario_3_2_4_opc_session_bad_error(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario 3.2.4 (mock): Tier-1 session error maps to confidence 14.
Scenario 3.2.4: OPC session/channel Tier-1 Bad* (e.g. BadSessionIdInvalid).
PostgreSQL stores prediction_confidence 14 and a stable session error comment.
"""
client = temporal_test_env.client
model_id = 324
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
@@ -652,87 +685,42 @@ async def test_scenario_3_2_4_opc_session_bad_mock(
opc_write_data.return_value = (
False,
{
'opc_error_kind': 'session_bad',
'opc_status': 'BadSessionIdInvalid',
'notification_id': 'OPC_WRITE_DATA_ERROR_1',
'message': 'OPC session invalid',
'message': 'BadSessionIdInvalid',
'block': 'opc_repository',
'level': NotificationLevel.ERROR,
'attachment_content': 'BadSessionIdInvalid',
'opc_error_kind': 'session_bad',
'opc_status': 'BadSessionIdInvalid',
},
)
input_data = get_base_input_data(model_id)
input_data['opc_output_config'] = {
'1': {
'prediction_tags': {'addr_1': {'data_type': 'float'}},
'confidence_tags': {'addr_2': {'data_type': 'float'}},
'prediction_tags': {
'addr_1': {
'data_type': 'float',
}
},
'confidence_tags': {
'addr_2': {
'data_type': 'float',
}
},
}
}
input_data['pi_web_api_output_config'] = None
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-opc-session-bad-mock')
client, PredictionsBatch.run, input_data, make_workflow_id('test-opc-session-bad')
)
assert_prediction(
postgres_engine,
model_id,
prediction_confidence=14,
comments_contains='OPC UA session/channel error: BadSessionIdInvalid',
)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_3_2_5_opc_reconnect_in_progress_mock(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario 3.2.5 (mock): reconnect_in_progress maps to confidence 14.
"""
client = temporal_test_env.client
model_id = 325
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
opc_write_data.return_value = (
False,
{
'opc_error_kind': 'reconnect_in_progress',
'notification_id': 'OPC_WRITE_DATA_ERROR_1',
'message': 'OPC reconnect in progress',
'block': 'opc_repository',
'level': NotificationLevel.ERROR,
'attachment_content': 'reconnect',
},
)
input_data = get_base_input_data(model_id)
input_data['opc_output_config'] = {
'1': {
'prediction_tags': {'addr_1': {'data_type': 'float'}},
'confidence_tags': {'addr_2': {'data_type': 'float'}},
}
}
input_data['pi_web_api_output_config'] = None
await start_and_await_workflow(
client,
PredictionsBatch.run,
input_data,
make_workflow_id('test-opc-reconnect-mock'),
)
assert_prediction(
postgres_engine,
model_id,
prediction_confidence=14,
comments_contains='OPC UA reconnect in progress',
comments='OPC UA session/channel error: BadSessionIdInvalid',
)
@@ -756,8 +744,8 @@ async def test_scenario_3_2_3_pi_web_api_partial_write_error(
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
test_activities.pi_web_api_client.set_side_effect(
[
test_activities.pi_web_api_client.write_value = AsyncMock(
side_effect=[
# Prediction batch: two web_ids requested, only one acknowledged.
[{'WebId': 'web_id_1', 'Errors': []}],
# Confidence write succeeds.
@@ -795,42 +783,4 @@ async def test_scenario_3_2_3_pi_web_api_partial_write_error(
prediction_confidence=13,
comments="The number of written tags does not match the number of tag names: Expected ['tag_1', 'tag_3'] tags, but ['tag_1'] tags were written.",
)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_3_3_1_combined_pi_and_opc_outputs(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario 3.3.1: PI and OPC enabled together.
"""
client = temporal_test_env.client
model_id = 333
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
input_data = get_base_input_data(model_id)
input_data['pi_web_api_output_config'] = {
'endpoint': 'test_endpoint',
'prediction_tags': {'tag_1': 'web_id_1'},
'confidence_tags': {'tag_2': 'web_id_2'},
}
input_data['opc_output_config'] = {
'1': {
'prediction_tags': {'addr_1': {'data_type': 'float'}},
'confidence_tags': {'addr_2': {'data_type': 'float'}},
}
}
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-pi-opc-combined')
)
assert test_activities.pi_web_api_client.write_value.call_count == 2
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
assert opc_write_data.call_count == 2
assert_prediction(postgres_engine, model_id)

View File

@@ -10,7 +10,7 @@ from temporalio.worker import Worker
import pytest
from e2e.helpers import load_scenario_input, make_workflow_id, start_and_await_workflow
from e2e.helpers import make_workflow_id, start_and_await_workflow
from laborious.activities.activities import Activities
from laborious.workflows.predictions_batch import PredictionsBatch
@@ -27,9 +27,9 @@ async def test_scenario_1_1_1_happy_path_complete_success(
client = temporal_test_env.client
with postgres_engine.begin() as conn:
conn.execute(text('DELETE FROM sientia_data.laborious_data WHERE model_id = 123'))
conn.execute(text('DELETE FROM predictions_schema.laborious_data WHERE model_id = 123'))
insert_sql = """
INSERT INTO sientia_data.laborious_data (model_id, variable, value, timestamp, created_at)
INSERT INTO predictions_schema.laborious_data (model_id, variable, value, timestamp, created_at)
VALUES
(123, 'sensor_1', 23.5, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00'),
(123, 'sensor_2', 78.2, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00'),
@@ -37,7 +37,43 @@ async def test_scenario_1_1_1_happy_path_complete_success(
"""
conn.execute(text(insert_sql))
input_data = load_scenario_input('main_happy_path.json', model_id=123)
input_data = {
'metadata': {
'metadata': {
'model_id': 123,
'model_name': 'test_model',
'schedule_name': 'test-schedule',
'workflow_name': 'predictions_batch',
}
},
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 123,
'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 123',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {
'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_transform_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_predict_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
'datetime_columns': ['timestamp', 'created_at'],
}
await start_and_await_workflow(
client,
@@ -46,7 +82,7 @@ async def test_scenario_1_1_1_happy_path_complete_success(
make_workflow_id('test-predictions-batch'),
)
schema_name = 'sientia_data'
schema_name = 'predictions_schema'
with postgres_engine.connect() as conn:
result_query = conn.execute(
text(
@@ -90,7 +126,42 @@ async def test_scenario_1_2_1_sql_query_execution_error(
"""Invalid SQL: workflow may complete with early exit; no prediction rows."""
client = temporal_test_env.client
input_data = load_scenario_input('main_sql_error.json', model_id=128)
input_data = {
'metadata': {
'metadata': {
'model_id': 128,
'model_name': 'test_model',
'schedule_name': 'test-schedule',
'workflow_name': 'predictions_batch',
}
},
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 128,
'query': 'SELECT * FROM nonexistent_table WHERE invalid_syntax =',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {
'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_transform_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_predict_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
}
await start_and_await_workflow(
client,
@@ -101,7 +172,7 @@ async def test_scenario_1_2_1_sql_query_execution_error(
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 128')
text('SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = 128')
).scalar()
assert count == 0
@@ -117,7 +188,22 @@ async def test_scenario_1_2_2_missing_required_parameters(
"""Missing query: workflow does not produce predictions and is terminated explicitly."""
client = temporal_test_env.client
input_data = load_scenario_input('main_missing_required.json', model_id=129)
input_data = {
'metadata': {
'metadata': {
'model_id': 129,
'model_name': 'test_model',
'schedule_name': 'test-schedule',
'workflow_name': 'predictions_batch',
}
},
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 129,
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
}
handle = await client.start_workflow(
PredictionsBatch.run,
@@ -131,7 +217,7 @@ async def test_scenario_1_2_2_missing_required_parameters(
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 129')
text('SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = 129')
).scalar()
assert count == 0
@@ -150,17 +236,53 @@ async def test_scenario_1_2_3_invalid_datetime_column_specification(
client = temporal_test_env.client
with postgres_engine.begin() as conn:
conn.execute(text('DELETE FROM sientia_data.laborious_data WHERE model_id = 130'))
conn.execute(text('DELETE FROM predictions_schema.laborious_data WHERE model_id = 130'))
conn.execute(
text(
"""
INSERT INTO sientia_data.laborious_data (model_id, variable, value, timestamp, created_at)
INSERT INTO predictions_schema.laborious_data (model_id, variable, value, timestamp, created_at)
VALUES (130, 'sensor_1', 23.5, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00')
"""
)
)
input_data = load_scenario_input('main_invalid_datetime.json', model_id=130)
input_data = {
'metadata': {
'metadata': {
'model_id': 130,
'model_name': 'test_model',
'schedule_name': 'test-schedule',
'workflow_name': 'predictions_batch',
}
},
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 130,
'query': 'SELECT timestamp, variable, value FROM predictions_schema.laborious_data WHERE model_id = 130',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {
'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_transform_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_predict_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
'datetime_columns': ['nonexistent_column'],
}
handle = await client.start_workflow(
PredictionsBatch.run,
@@ -174,7 +296,7 @@ async def test_scenario_1_2_3_invalid_datetime_column_specification(
with postgres_engine.connect() as conn:
count = conn.execute(
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 130')
text('SELECT COUNT(*) FROM predictions_schema.predictions WHERE model_id = 130')
).scalar()
assert count == 0

View File

@@ -14,54 +14,90 @@ from temporalio.worker import Worker
from e2e.helpers import (
assert_continue,
assert_postgres_unique_violation_in_chain,
assert_prediction,
assert_prediction_row_count,
assert_repeat,
assert_stop,
insert_sample_data,
insert_sample_prediction,
load_scenario_input,
make_workflow_id,
start_and_await_workflow,
)
from laborious.activities.activities import Activities
from laborious.utils.models import minio_dataframe_payload as minio_payload_module
from laborious.workflows.predictions_batch import PredictionsBatch
DISTINCT_BATCH_TIMESTAMP = '2024-01-01 13:00:00+00:00'
HISTORY_TIMESTAMP = '2024-01-01 12:00:00+00:00'
base_input_data = {
'schedule_name': 'test-schedule',
'model_name': 'test_model',
'model_id': 201,
'query': 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = 201',
'schema': 'predictions_schema',
'table_name': 'predictions',
'transform_table_name': 'transformed_data',
'input_filters': {
'SPECIFIC_VARIABLES_NULL_VALUES': {
'POLICY': 'CONTINUE',
'CONFIG': {'variables': ['sensor_1']},
},
},
'mlflow_transform_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'mlflow_predict_filters': {
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
},
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
'opc_output_config': {},
'pi_web_api_output_config': {},
'save_transform': True,
'prediction_store_policy': 'lts:1',
'model_config': {
'retention_minutes': 0,
'transform_flavor': 'sklearn',
'predict_flavor': 'sklearn',
},
'datetime_columns': ['timestamp', 'created_at'],
}
base_query = 'SELECT timestamp, variable, value, created_at FROM predictions_schema.laborious_data WHERE model_id = {model_id}'
def get_base_input_data(model_id):
return load_scenario_input('prediction_process_base.json', model_id=model_id)
return {
**base_input_data,
'model_id': model_id,
'query': base_query.format(model_id=model_id),
}
def insert_sample_prediction(postgres_engine, model_id):
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}'))
insert_sql = f"""
INSERT INTO predictions_schema.predictions (model_id, timestamp, prediction, prediction_confidence, prediction_status, comments, response_time)
VALUES
({model_id}, '2024-01-01 12:00:00+00:00', 10, 0, 'Good', '', 0.1)
"""
conn.execute(text(insert_sql))
return (model_id, Decimal(10), Decimal(0), 'Good')
@pytest.fixture
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
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
@pytest.fixture
def bad_predict_model(mlflow_repository_stub):
wrapper = mlflow_repository_stub.stub_wrapper
def bad_predict_model(patch_mlflow, mock_mlflow_models):
model = MagicMock(predict=MagicMock(side_effect=Exception('Bad predict 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, {}
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
wrapper.transform.side_effect = _good_transform
wrapper.predict = MagicMock(side_effect=Exception('Bad predict model'))
return wrapper
patch_mlflow.sklearn = MagicMock()
patch_mlflow.sklearn.load_model = MagicMock(side_effect=mock_sklearn_load_model)
return model
@pytest.mark.asyncio
@@ -71,7 +107,7 @@ async def test_scenario_2_1_1_input_gate_triggers_continue(
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
mock_mlflow_models,
):
"""Input gate CONTINUE: export default prediction; MLflow transform/predict not used."""
client = temporal_test_env.client
@@ -82,8 +118,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)
mlflow_repository_stub.stub_wrapper.transform.assert_not_called()
mlflow_repository_stub.stub_wrapper.predict.assert_not_called()
mock_mlflow_models['transform_model'].predict.assert_not_called()
mock_mlflow_models['predict_model'].predict.assert_not_called()
@pytest.mark.asyncio
@@ -93,7 +129,7 @@ async def test_scenario_2_1_2_input_gate_triggers_stop(
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
mock_mlflow_models,
):
"""Input gate STOP: no export, no MLflow."""
client = temporal_test_env.client
@@ -105,60 +141,30 @@ 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)
mlflow_repository_stub.stub_wrapper.transform.assert_not_called()
mock_mlflow_models['transform_model'].predict.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_1_3_input_gate_repeat_batch_timestamp_equals_history_fails(
async def test_scenario_2_1_3_input_gate_triggers_repeat(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
mock_mlflow_models,
):
"""
REPEAT uses ``last_timestamp`` from the batch payload as the new row's ``timestamp``.
When it equals the only historical prediction row, Postgres rejects the duplicate key.
"""
"""Input gate REPEAT with existing history."""
client = temporal_test_env.client
model_id = 213
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2], data_timestamp=HISTORY_TIMESTAMP)
insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
input_data = get_base_input_data(model_id)
input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'REPEAT'
with pytest.raises(Exception) as excinfo:
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-input-repeat-collision')
)
assert_postgres_unique_violation_in_chain(excinfo.value)
assert_prediction_row_count(postgres_engine, model_id, 1)
mlflow_repository_stub.stub_wrapper.transform.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_1_3_input_gate_repeat_distinct_batch_timestamp_inserts_second_row(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""REPEAT succeeds when batch ``last_timestamp`` differs from the historical prediction row."""
client = temporal_test_env.client
model_id = 2131
insert_sample_data(
postgres_engine, model_id, ['NULL', 78.2], data_timestamp=DISTINCT_BATCH_TIMESTAMP
)
data = insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
data = insert_sample_prediction(postgres_engine, model_id)
input_data = get_base_input_data(model_id)
input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'REPEAT'
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-input-repeat-ok')
client, PredictionsBatch.run, input_data, make_workflow_id('test-input-repeat')
)
assert_repeat(postgres_engine, model_id, data)
mlflow_repository_stub.stub_wrapper.transform.assert_not_called()
mock_mlflow_models['transform_model'].predict.assert_not_called()
@pytest.mark.asyncio
@@ -174,7 +180,7 @@ async def test_scenario_2_1_4_input_gate_repeat_without_prior_prediction(
model_id = 214
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.predictions WHERE model_id = {model_id}'))
input_data = get_base_input_data(model_id)
input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'REPEAT'
await start_and_await_workflow(
@@ -216,7 +222,7 @@ async def test_scenario_2_2_2_transform_gate_triggers_stop(
test_activities: Activities,
postgres_engine,
bad_data_model,
mlflow_repository_stub,
mock_mlflow_models,
):
client = temporal_test_env.client
model_id = 222
@@ -227,12 +233,12 @@ 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)
mlflow_repository_stub.stub_wrapper.predict.assert_not_called()
mock_mlflow_models['predict_model'].predict.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_2_3_transform_gate_repeat_batch_timestamp_equals_history_fails(
async def test_scenario_2_2_3_transform_gate_triggers_repeat(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
@@ -241,37 +247,12 @@ async def test_scenario_2_2_3_transform_gate_repeat_batch_timestamp_equals_histo
):
client = temporal_test_env.client
model_id = 223
insert_sample_data(postgres_engine, model_id, [60.0, 78.2], data_timestamp=HISTORY_TIMESTAMP)
insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
input_data = get_base_input_data(model_id)
input_data['mlflow_transform_filters']['API_ERROR']['POLICY'] = 'REPEAT'
with pytest.raises(Exception) as excinfo:
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-repeat-collision')
)
assert_postgres_unique_violation_in_chain(excinfo.value)
assert_prediction_row_count(postgres_engine, model_id, 1)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_2_3_transform_gate_repeat_distinct_batch_timestamp_inserts_second_row(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
bad_data_model,
):
client = temporal_test_env.client
model_id = 2231
insert_sample_data(
postgres_engine, model_id, [60.0, 78.2], data_timestamp=DISTINCT_BATCH_TIMESTAMP
)
data = insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
insert_sample_data(postgres_engine, model_id, [60.0, 78.2])
data = insert_sample_prediction(postgres_engine, model_id)
input_data = get_base_input_data(model_id)
input_data['mlflow_transform_filters']['API_ERROR']['POLICY'] = 'REPEAT'
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-repeat-ok')
client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-repeat')
)
assert_repeat(postgres_engine, model_id, data)
@@ -283,20 +264,19 @@ async def test_scenario_2_2_4_transform_content_gate_nan_values_stop(
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
mock_mlflow_models,
):
"""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):
result = pd.DataFrame(
{'feature_1': [np.nan] * len(data), 'feature_2': [np.nan] * len(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.index = data.index
return result, {}
return result
mlflow_repository_stub.stub_wrapper.transform = MagicMock(side_effect=all_nan_transform)
mock_mlflow_models['transform_model'].predict = 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)
@@ -308,7 +288,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)
mlflow_repository_stub.stub_wrapper.predict.assert_not_called()
mock_mlflow_models['predict_model'].predict.assert_not_called()
@pytest.mark.asyncio
@@ -358,7 +338,7 @@ async def test_scenario_2_3_2_predict_gate_triggers_stop(
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_3_3_predict_gate_repeat_batch_timestamp_equals_history_fails(
async def test_scenario_2_3_3_predict_gate_triggers_repeat(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
@@ -367,39 +347,13 @@ async def test_scenario_2_3_3_predict_gate_repeat_batch_timestamp_equals_history
):
client = temporal_test_env.client
model_id = 233
insert_sample_data(postgres_engine, model_id, [23.5, 78.2], data_timestamp=HISTORY_TIMESTAMP)
insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
input_data = get_base_input_data(model_id)
input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'REPEAT'
input_data['path_priority'] = ['REPEAT', 'STOP', 'CONTINUE']
with pytest.raises(Exception) as excinfo:
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-predict-repeat-collision')
)
assert_postgres_unique_violation_in_chain(excinfo.value)
assert_prediction_row_count(postgres_engine, model_id, 1)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_3_3_predict_gate_repeat_distinct_batch_timestamp_inserts_second_row(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
bad_predict_model,
):
client = temporal_test_env.client
model_id = 2331
insert_sample_data(
postgres_engine, model_id, [23.5, 78.2], data_timestamp=DISTINCT_BATCH_TIMESTAMP
)
data = insert_sample_prediction(postgres_engine, model_id, prediction_timestamp=HISTORY_TIMESTAMP)
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
data = insert_sample_prediction(postgres_engine, model_id)
input_data = get_base_input_data(model_id)
input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'REPEAT'
input_data['path_priority'] = ['REPEAT', 'STOP', 'CONTINUE']
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-predict-repeat-ok')
client, PredictionsBatch.run, input_data, make_workflow_id('test-predict-repeat')
)
assert_repeat(postgres_engine, model_id, data)
@@ -416,73 +370,10 @@ async def test_scenario_2_4_1_input_empty_data_stop(
client = temporal_test_env.client
model_id = 241
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM predictions_schema.laborious_data WHERE model_id = {model_id}'))
input_data = get_base_input_data(model_id)
input_data['input_filters'] = {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}}
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-empty-data-stop')
)
assert_stop(postgres_engine, model_id)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_scenario_2_4_1_priority_conflict_resolution(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
bad_data_model,
):
"""Conflicting filter outputs must honor configured path_priority order."""
client = temporal_test_env.client
model_id = 242
insert_sample_data(postgres_engine, model_id, [60.0, 78.2])
input_data = get_base_input_data(model_id)
input_data['mlflow_transform_filters'] = {'API_ERROR': {'POLICY': 'CONTINUE', 'CONFIG': {}}}
input_data['mlflow_predict_filters'] = {'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}}}
input_data['path_priority'] = ['STOP', 'CONTINUE', 'REPEAT']
await start_and_await_workflow(
client, PredictionsBatch.run, input_data, make_workflow_id('test-priority-conflict')
)
assert_continue(
postgres_engine=postgres_engine,
model_id=model_id,
prediction_confidence=Decimal(10),
comments='Unknown MLFlow API error',
)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_e2e_request_predict_inline_minio_payload_with_datetimeindex(
temporal_test_env: WorkflowEnvironment,
temporal_worker: Worker,
test_activities: Activities,
postgres_engine,
mlflow_repository_stub,
):
"""
High offload threshold forces inline tabular dicts; ``DatetimeIndex`` must serialize as JSON
(string index keys via ``MinioDataFramePayload.from_dataframe``) so ``request_predict`` completes.
"""
client = temporal_test_env.client
model_id = 252
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM sientia_data.transformed_data WHERE model_id = {model_id}'))
insert_sample_data(postgres_engine, model_id, [60.0, 78.2])
input_data = get_base_input_data(model_id)
with patch.object(minio_payload_module, 'OFFLOAD_THRESHOLD_BYTES', 10**9):
await start_and_await_workflow(
client,
PredictionsBatch.run,
input_data,
make_workflow_id('test-predict-inline-json-datetimeindex'),
)
assert_prediction(postgres_engine, model_id, prediction=0.5, prediction_confidence=0)
mlflow_repository_stub.stub_wrapper.predict.assert_called()

View File

@@ -1,333 +0,0 @@
"""
End-to-end tests for the SimpleMetrics workflow.
Coverage focus:
- Happy path computes rmse/mse/mae/r2 from predictions joined against ``laborious_data``
and persists rows to ``sientia_data.simple_metrics`` with all required columns.
- Subset metric selection (only rmse) writes exactly the requested rows.
- Zero-variance target produces ``r2=0`` per division-by-zero guard.
- Empty join (no overlapping data) short-circuits without persisting anything.
"""
from datetime import datetime, timedelta, timezone
from decimal import Decimal
import math
import pytest
from sqlalchemy import text
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from e2e.helpers import load_scenario_input, make_workflow_id, start_and_await_workflow
from laborious.activities.activities import Activities
from laborious.workflows.simple_metrics import SimpleMetrics
# Columns defined by the production DDL for ``sientia_data.simple_metrics``.
EXPECTED_SIMPLE_METRICS_COLUMNS = [
'id',
'model_id',
'metric',
'value',
'timestamp',
'data_size',
'interval_minutes',
'created_at',
]
# ``timestamp`` is now nullable per the new DDL (production code may write it
# null when the upstream data has no usable instant); skip the non-null check
# for it while still validating presence.
NULLABLE_SIMPLE_METRICS_COLUMNS = {'timestamp'}
def _simple_metrics_input(model_id: int, **overrides) -> dict:
"""Load and override the simple-metrics base scenario."""
payload = load_scenario_input('simple_metrics_base.json', model_id=model_id)
payload.update(overrides)
return payload
def _seed_predictions_and_targets(
postgres_engine,
model_id: int,
pairs: list[tuple[float, float]],
target_name: str = 'sensor_target',
offset_minutes: int = 6,
) -> list[str]:
"""
Insert matching prediction/target rows used by the SimpleMetrics SQL JOIN.
For each ``(prediction, target)`` pair we write a row in ``predictions`` and
a matching row in ``laborious_data`` with ``variable=target_name`` so the
inner join in the workflow query yields one row per pair.
Args:
- postgres_engine: SQLAlchemy engine bound to the test container.
- model_id: Model id stamped on every row.
- pairs: ``(prediction, target)`` pairs, one per minute.
- target_name: Variable name in ``laborious_data`` representing the target.
- offset_minutes: Earliest row sits this many minutes ago so timestamps fall
inside the workflow's recent-data window.
Return:
List of timestamp strings written for the inserted rows.
"""
base = datetime.now(timezone.utc).replace(second=0, microsecond=0) - timedelta(
minutes=offset_minutes
)
timestamps = [
(base + timedelta(minutes=i)).strftime('%Y-%m-%d %H:%M:%S%z')
for i in range(len(pairs))
]
prediction_rows = []
target_rows = []
for index, (prediction, target_value) in enumerate(pairs):
ts = timestamps[index]
prediction_rows.append(
f"({model_id}, {prediction}, 0, 0, 'Good', '{ts}', '{ts}')"
)
target_rows.append(
f"({model_id}, '{target_name}', {target_value}, '{ts}', '{ts}')"
)
with postgres_engine.begin() as conn:
conn.execute(text(f'DELETE FROM sientia_data.predictions WHERE model_id = {model_id}'))
conn.execute(text(f'DELETE FROM sientia_data.laborious_data WHERE model_id = {model_id}'))
# The SimpleMetrics SQL JOIN only filters ``predictions.model_id``; it does
# NOT filter ``laborious_data.model_id`` (see ``e2e/CODE_ISSUES.md`` issue
# SM-1). Without this cross-model cleanup, a previous test's target rows
# under the same variable name would join into this test's predictions
# whenever timestamps happened to overlap.
conn.execute(
text(
"DELETE FROM sientia_data.laborious_data "
"WHERE variable IN (:sensor_default, :target_name) "
"AND timestamp >= NOW() - INTERVAL '120 minutes'"
),
{'sensor_default': 'sensor_target', 'target_name': target_name},
)
if prediction_rows:
conn.execute(
text(
'INSERT INTO sientia_data.predictions '
'(model_id, prediction, prediction_confidence, response_time, '
'prediction_status, "timestamp", created_at) VALUES '
+ ', '.join(prediction_rows)
)
)
conn.execute(
text(
'INSERT INTO sientia_data.laborious_data '
'(model_id, variable, value, "timestamp", created_at) VALUES '
+ ', '.join(target_rows)
)
)
return timestamps
@pytest.mark.asyncio
@pytest.mark.integration
async def test_simple_metrics_happy_path_persists_all_metrics_and_columns(
temporal_test_env: WorkflowEnvironment,
temporal_worker_simple_metrics: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario S.1.1: rmse/mse/mae/r2 are calculated from a deterministic
prediction/target pair set and written one row per metric. Every column
expected by ``sientia_data.simple_metrics`` must be populated (except the
nullable ``timestamp`` column) and the numerical values must match
closed-form expectations.
"""
client = temporal_test_env.client
model_id = 511
pairs = [
(1.0, 2.0),
(2.0, 4.0),
(3.0, 5.0),
(4.0, 9.0),
(5.0, 12.0),
]
diffs = [target - prediction for prediction, target in pairs]
n = len(diffs)
expected_rmse = math.sqrt(sum(d * d for d in diffs) / n)
expected_mse = sum(d * d for d in diffs) / n
expected_mae = sum(abs(d) for d in diffs) / n
target_mean = sum(t for _, t in pairs) / n
ss_res = sum((target - prediction) ** 2 for prediction, target in pairs)
ss_tot = sum((t - target_mean) ** 2 for _, t in pairs)
expected_r2 = 1.0 - (ss_res / ss_tot)
_seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs)
input_data = _simple_metrics_input(model_id)
await start_and_await_workflow(
client,
SimpleMetrics.run,
input_data,
make_workflow_id('test-simple-metrics-happy'),
)
with postgres_engine.connect() as conn:
rows = (
conn.execute(
text(
'SELECT * FROM sientia_data.simple_metrics '
'WHERE model_id = :m ORDER BY metric'
),
{'m': str(model_id)},
)
.mappings()
.all()
)
assert len(rows) == 4, f'Expected 4 metric rows, got {len(rows)}'
for column in EXPECTED_SIMPLE_METRICS_COLUMNS:
assert column in rows[0], f'Missing simple_metrics column: {column}'
for row in rows:
for column in EXPECTED_SIMPLE_METRICS_COLUMNS:
if column in NULLABLE_SIMPLE_METRICS_COLUMNS:
continue
assert row[column] is not None, f"Column '{column}' is NULL in {dict(row)}"
by_metric = {row['metric']: row for row in rows}
assert set(by_metric) == {'rmse', 'mse', 'mae', 'r2'}
def _decimal_close(actual, expected, places: int = 6) -> bool:
return abs(float(actual) - expected) < 10 ** (-places)
assert _decimal_close(by_metric['rmse']['value'], expected_rmse)
assert _decimal_close(by_metric['mse']['value'], expected_mse)
assert _decimal_close(by_metric['mae']['value'], expected_mae)
assert _decimal_close(by_metric['r2']['value'], expected_r2)
assert all(row['data_size'] == n for row in rows), 'data_size must equal target row count'
assert all(row['interval_minutes'] == 60 for row in rows)
# ``model_id`` is now ``text`` in the new DDL, so we compare with the
# stringified test id rather than the numeric value.
assert all(row['model_id'] == str(model_id) for row in rows)
@pytest.mark.asyncio
@pytest.mark.integration
async def test_simple_metrics_subset_metrics_writes_only_requested_rows(
temporal_test_env: WorkflowEnvironment,
temporal_worker_simple_metrics: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario S.1.2: Requesting ``metrics=['rmse']`` must persist exactly one row
with metric ``rmse`` and skip mse/mae/r2.
"""
client = temporal_test_env.client
model_id = 512
pairs = [(1.0, 2.0), (2.0, 4.0), (3.0, 6.0)]
_seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs)
input_data = _simple_metrics_input(model_id, metrics=['rmse'])
await start_and_await_workflow(
client,
SimpleMetrics.run,
input_data,
make_workflow_id('test-simple-metrics-subset'),
)
with postgres_engine.connect() as conn:
metrics = [
r[0]
for r in conn.execute(
text(
'SELECT metric FROM sientia_data.simple_metrics '
'WHERE model_id = :m'
),
{'m': str(model_id)},
).all()
]
assert metrics == ['rmse']
@pytest.mark.asyncio
@pytest.mark.integration
async def test_simple_metrics_zero_variance_target_returns_zero_r2(
temporal_test_env: WorkflowEnvironment,
temporal_worker_simple_metrics: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario S.2.1: When the target column has zero variance the activity must
return ``r2 = 0`` (division-by-zero guard) and still persist all four metrics.
"""
client = temporal_test_env.client
model_id = 521
pairs = [(0.0, 5.0), (1.0, 5.0), (2.0, 5.0), (3.0, 5.0)]
_seed_predictions_and_targets(postgres_engine, model_id=model_id, pairs=pairs)
input_data = _simple_metrics_input(model_id)
await start_and_await_workflow(
client,
SimpleMetrics.run,
input_data,
make_workflow_id('test-simple-metrics-zero-variance'),
)
with postgres_engine.connect() as conn:
r2_value = conn.execute(
text(
"SELECT value FROM sientia_data.simple_metrics "
"WHERE model_id = :m AND metric = 'r2'"
),
{'m': str(model_id)},
).scalar()
assert r2_value is not None
assert Decimal(str(r2_value)) == Decimal('0'), f'expected r2=0, got {r2_value!r}'
@pytest.mark.asyncio
@pytest.mark.integration
async def test_simple_metrics_no_overlapping_data_short_circuits(
temporal_test_env: WorkflowEnvironment,
temporal_worker_simple_metrics: Worker,
test_activities: Activities,
postgres_engine,
):
"""
Scenario S.3.1: When the join produces no rows (no matching laborious_data
row for the configured ``target``), the workflow returns early without
invoking ``calculate_simple_metrics`` and writes nothing.
"""
client = temporal_test_env.client
model_id = 531
# Insert predictions but no matching target rows for the configured variable.
_seed_predictions_and_targets(
postgres_engine,
model_id=model_id,
pairs=[(1.0, 1.0)],
target_name='wrong_variable_name',
)
input_data = _simple_metrics_input(model_id)
await start_and_await_workflow(
client,
SimpleMetrics.run,
input_data,
make_workflow_id('test-simple-metrics-empty-join'),
)
with postgres_engine.connect() as conn:
count = conn.execute(
text(
'SELECT COUNT(*) FROM sientia_data.simple_metrics '
'WHERE model_id = :m'
),
{'m': str(model_id)},
).scalar()
assert count == 0, 'Empty target data must short-circuit and skip persistence'

View File

@@ -1,2 +1,2 @@
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git:sientia_do
git+ssh://git@github.com/Aignosi/sientia-model-library.git:sientia_model
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git:sientia-do
git+ssh://git@github.com/Aignosi/sientia-mlops-library.git:sientia

View File

@@ -1,143 +0,0 @@
{
"models": [
{
"id": "1001",
"name": "test-runtime",
"active": false,
"model_config": {
"alias": "production",
"retention_minutes": 60,
"target": "Square"
}
}
],
"pipelines": [
{
"schedule_name": "laborious-test-runtime",
"model_id": "1001",
"workflow_type": "predictions_batch",
"frequency": "60s",
"max_retry_policy": 1,
"query": "select * from sientia_data.laborious_data where model_id = 1 and \"timestamp\" > NOW() - INTERVAL '5 minutes' order by \"timestamp\" desc limit 30;",
"retention_time": 60,
"write_tags": [],
"input_filters": [
{
"filter_name": "EMPTY_DATA",
"policy": "STOP"
},
{
"filter_name": "SPECIFIC_VARIABLES_NULL_VALUES",
"policy": "CONTINUE",
"config": {
"variables": [
"Counter"
]
}
}
],
"mlflow_transform_filters": [
{
"filter_name": "API_ERROR",
"policy": "REPEAT"
},
{
"filter_name": "NAN_VALUES",
"policy": "STOP"
}
],
"mlflow_predict_filters": [
{
"filter_name": "API_ERROR",
"policy": "CONTINUE"
}
],
"path_priority": [
"STOP",
"CONTINUE",
"REPEAT"
],
"active": true,
"updated_at": {
"$date": "2026-05-07T23:35:01.600Z"
},
"save_transform": false,
"pi_web_api_output_config": {},
"datetime_columns": [
"timestamp",
"created_at"
]
},
{
"schedule_name": "minimal-retrain-test-runtime",
"model_id": "1001",
"model_name": "test-runtime",
"workflow_type": "minimal_retrain",
"frequency": "1h",
"max_retry_policy": 1,
"query": "select * from sientia_data.laborious_data where model_id = 1 and \"timestamp\" > NOW() - INTERVAL '60 minutes' order by \"timestamp\" desc;",
"schema": "sientia_data",
"table_name": "log_retrain",
"datetime_columns": ["timestamp", "created_at"],
"model_config": {
"target": "Square"
},
"active": true,
"updated_at": {
"$date": "2026-05-07T23:35:01.600Z"
}
},
{
"schedule_name": "drift-test-runtime",
"model_id": "1001",
"model_name": "test-runtime",
"workflow_type": "drift",
"frequency": "5m",
"offset": "2m",
"max_retry_policy": 1,
"execution_timeout_seconds": 300,
"task_timeout_seconds": 300,
"interval": 5,
"drift_metrics": [
"kolmogorov_smirnov",
"jensen_shannon",
"wasserstein"
],
"chunk_period": "min",
"schema": "sientia_data",
"source_table_name": "laborious_data",
"target_table_name": "drift_metrics",
"model_config": {
"target": "Square"
},
"active": true,
"updated_at": {
"$date": "2026-05-07T23:35:01.600Z"
}
},
{
"schedule_name": "simple-metrics-test-runtime",
"model_id": "1001",
"model_name": "test-runtime",
"workflow_type": "simple_metrics",
"frequency": "5m",
"offset": "2m",
"max_retry_policy": 1,
"execution_timeout_seconds": 300,
"task_timeout_seconds": 300,
"interval_minutes": 5,
"metrics": ["rmse", "mse", "mae", "r2"],
"schema": "sientia_data",
"predictions_table_name": "predictions",
"data_table_name": "laborious_data",
"target_table_name": "simple_metrics",
"model_config": {
"target": "Square"
},
"active": true,
"updated_at": {
"$date": "2026-05-18T23:35:01.600Z"
}
}
]
}

View File

@@ -1,305 +0,0 @@
---
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 wrappers `_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 | Wrapperbased training using public API | Retraining must call the public `retrain(...)` method of `SientiaModel` | `fit_models`, `retrain_model` paths |
| FR-03 | Wrapperbased 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 wrappers 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 wrappers 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 wrappers `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 (T1T2) 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]]

View File

@@ -6,9 +6,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore
from sientia_do.repository.minio_repository import MinioRepository
from laborious.activities.api import API
from laborious.activities.gates import Gates
@@ -16,78 +14,65 @@ with workflow.unsafe.imports_passed_through():
from laborious.activities.model_metrics import ModelMetrics
from laborious.activities.opc import OPC
from laborious.activities.storage import Storage
from laborious.utils.connectors_config import build_mlflow_config
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
"""
Central orchestrator for all Temporal activities used by Laborious workflows.
Main activities orchestrator for the Laborious system.
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.
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.
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.
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
Attributes:
Inherits and exposes behaviour from mixins; the MLFlow mixin holds ``mlflow_repository``
and ``plugin_store`` after ``__init__``.
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
"""
def __init__(
self,
postgres_config: dict[str, Any],
plugin_store: PluginStore,
mlflow_config: dict[str, Any],
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,
):
"""
Wire Postgres, MinIO, MLflow, OPC, gates, metrics, and PI Web API into a single object.
Initialize the Activities orchestrator with all required configurations.
A single ``MetricsController`` instance is created (or reused) and passed to MinIO,
MLflow repository, and all mixins so Prometheus and SDK metrics stay consistent.
This constructor initializes all parent classes with their respective
configurations and sets up the foundation for all activity operations.
Args:
- 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()``.
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
Raises:
Exception: If any parent ``__init__`` fails (e.g. invalid config keys).
Return:
None
Exception: If any parent class initialization fails
"""
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,
)
metrics_controller = MetricsController(logger=logger)
minio_repository = MinioRepository(
endpoint=minio_config['endpoint_url'],
@@ -96,10 +81,11 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
bucket=minio_config['default_bucket'],
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
secure=minio_config['secure'],
)
# Initialize parent classes
Storage.__init__(
self,
host=postgres_config['host'],
@@ -113,17 +99,19 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
minio_repository=minio_repository,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
MLFlow.__init__(
self,
mlflow_repository=mlflow_repository,
plugin_store=plugin_store,
mlflow_host=mlflow_config['host'],
mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'],
minio_repository=minio_repository,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
Gates.__init__(
@@ -131,7 +119,7 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
minio_repository=minio_repository,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
OPC.__init__(
@@ -139,14 +127,14 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
opc_servers=opc_config,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
ModelMetrics.__init__(
self,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
API.__init__(
@@ -156,22 +144,26 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
auth_token=pi_web_api_config['auth_token'],
logger=logger,
notification_handler=notification_handler,
metrics_controller=mc,
metrics_controller=metrics_controller,
)
def shutdown(self) -> None:
async def shutdown(self):
"""
Close database pools, sync clients, and OPC sessions in a defined order.
Gracefully shutdown all activities and clean up resources.
Should be invoked on worker exit so connection pools and OPC sessions are released
cleanly before process termination.
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
Return:
None
The method should be called before the application terminates to ensure
proper resource cleanup and prevent resource leaks.
"""
Storage.close(self)
MLFlow.close(self)
Gates.close(self)
OPC.close(self)
await OPC.aclose(self)
ModelMetrics.close(self)
API.close(self)

View File

@@ -11,7 +11,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.pi_web_api_client_sync import PIWebAPIClient
from sientia_do.repository.pi_web_api_client import PIWebAPIClient
from laborious import metrics
@@ -105,7 +105,7 @@ class API(SientiaMonitoring):
self.pi_web_api_client.close()
SientiaMonitoring.shutdown(self)
def process_pi_web_api_response(
async def process_pi_web_api_response(
self,
response_data: list[dict[str, Any]],
tags: dict[str, str],
@@ -156,7 +156,7 @@ class API(SientiaMonitoring):
self.error(
f'Error writing tag {tag_name}:{web_id} to PI Web API: {errors}', metadata
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PI_WEB_API_PREDICTION_WRITTEN_ERROR_COUNT,
tags={
**core_labels,
@@ -165,7 +165,7 @@ class API(SientiaMonitoring):
)
confidence = PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
else:
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PI_WEB_API_PREDICTION_WRITTEN_COUNT,
tags={
**core_labels,
@@ -182,7 +182,7 @@ class API(SientiaMonitoring):
metadata,
)
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR',
message=f'The number of written tags does not match the number of tag names: Expected {tag_names} tags, but {written_tags} tags were written.\nResponse:\n {json.dumps(response_data, indent=4)}\nTags:\n {json.dumps(tags, indent=4)}',
@@ -194,7 +194,7 @@ class API(SientiaMonitoring):
return confidence, message
@activity.defn(name='write_pi_web_api_data')
def write_pi_web_api_data(self, input_data: dict[str, Any]) -> dict[Any, Any]:
async def write_pi_web_api_data(self, input_data: dict[str, Any]) -> dict[Any, Any]:
"""
Write prediction and confidence data to PI Web API.
@@ -232,7 +232,7 @@ class API(SientiaMonitoring):
confidence_value = data.head(1)['prediction_confidence'].values[0]
try:
prediction_response = self.pi_web_api_client.write_value(
prediction_response = await self.pi_web_api_client.write_value(
web_ids=prediction_tags,
value={
'Timestamp': data.head(1)['timestamp'].values[0],
@@ -241,7 +241,7 @@ class API(SientiaMonitoring):
metadata=metadata,
)
confidence, message = self.process_pi_web_api_response(
confidence, message = await self.process_pi_web_api_response(
response_data=prediction_response,
tags=raw_prediction_tags,
core_labels=core_labels,
@@ -258,7 +258,7 @@ class API(SientiaMonitoring):
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR',
message=f'Error writing prediction data to PI Web API: {e}\n Tags: {raw_prediction_tags}',
@@ -275,7 +275,7 @@ class API(SientiaMonitoring):
return data.to_dict()
try:
confidence_response = self.pi_web_api_client.write_value(
confidence_response = await self.pi_web_api_client.write_value(
web_ids=confidence_tags,
value={
'Timestamp': data.head(1)['timestamp'].values[0],
@@ -284,7 +284,7 @@ class API(SientiaMonitoring):
metadata=metadata,
)
self.process_pi_web_api_response(
await self.process_pi_web_api_response(
response_data=confidence_response,
tags=raw_confidence_tags,
core_labels=core_labels,
@@ -293,7 +293,7 @@ class API(SientiaMonitoring):
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR',
message=f'Error writing confidence data to PI Web API: {e}\n Tags: {raw_confidence_tags}',

View File

@@ -1,5 +1,8 @@
from sientia_do.repository.minio_repository import MinioRepository
from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through():
import traceback
from collections.abc import Callable, Mapping
@@ -10,8 +13,6 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_do.utils.formatters import create_sample_dict
from laborious import metrics
@@ -65,7 +66,7 @@ mlflow_content_path_confidence: Mapping[str, int] = {
}
class Gates(SientiaMonitoring):
class Gates(MinioManager):
"""
Data quality gates and filtering activities for the Laborious system.
@@ -105,12 +106,8 @@ class Gates(SientiaMonitoring):
Raises:
Exception: If BaseActivity initialization fails
"""
self.minio_repository = minio_repository
SientiaMonitoring.__init__(
self,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
def close(self) -> None:
@@ -118,12 +115,7 @@ class Gates(SientiaMonitoring):
Close the gates activity and clean up resources.
"""
if self.minio_repository is not None:
try:
self.minio_repository.close()
finally:
self.minio_repository = None
SientiaMonitoring.shutdown(self)
MinioManager.close(self)
def __del__(self):
self.close()
@@ -163,7 +155,7 @@ class Gates(SientiaMonitoring):
return policy, filter_config
@activity.defn(name='input_gate')
def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
"""
Apply input data quality filters and validation.
@@ -202,7 +194,7 @@ class Gates(SientiaMonitoring):
filters = input_data['filters']
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
path_priority = input_data['path_priority']
filter_output = []
@@ -222,7 +214,7 @@ class Gates(SientiaMonitoring):
filter_output.append(policy)
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'INTPUT_GATE_ERROR__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}',
@@ -243,7 +235,7 @@ class Gates(SientiaMonitoring):
return None, 0, ''
@activity.defn(name='mlflow_response_gate')
def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
async def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
"""
Validate MLFlow API response quality and integrity.
@@ -287,7 +279,7 @@ class Gates(SientiaMonitoring):
self.debug(f'Filters: {filters}', metadata)
payload = MinioDataFramePayload.from_dict(raw_data)
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
gate_type = input_data['type']
path_priority = input_data['path_priority']
@@ -306,7 +298,7 @@ class Gates(SientiaMonitoring):
if mlflow_response_filter_functions[fil](status, filter_config):
filter_output.append(policy)
comments.append(status.get('message', 'Unknown MLFlow API error'))
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}',
message=status.get('message', 'Unknown MLFlow API error'),
@@ -316,7 +308,7 @@ class Gates(SientiaMonitoring):
)
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}',
@@ -337,7 +329,7 @@ class Gates(SientiaMonitoring):
return None, 0, ''
@activity.defn(name='mlflow_content_gate')
def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
async def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
"""
Validate MLFlow prediction content quality and integrity.
@@ -376,7 +368,7 @@ class Gates(SientiaMonitoring):
filters = input_data['filters']
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
gate_type = input_data['type']
path_priority = input_data['path_priority']
@@ -393,7 +385,7 @@ class Gates(SientiaMonitoring):
try:
if mlflow_content_filter_functions[fil](data, filter_config):
filter_output.append(policy)
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}',
message=f'Data not passed the content filter {fil}:{config}',
@@ -403,7 +395,7 @@ class Gates(SientiaMonitoring):
)
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}',
@@ -479,7 +471,7 @@ class Gates(SientiaMonitoring):
return policy_type, int(policy_value)
@activity.defn(name='format_transformed_data')
def format_transformed_data(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
async def format_transformed_data(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
Format transformed data for storage and export operations.
@@ -515,7 +507,7 @@ class Gates(SientiaMonitoring):
self.info('Formatting transformed data...', metadata)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
data['timestamp'] = data.index
data = data.reset_index(drop=True)
@@ -523,7 +515,7 @@ class Gates(SientiaMonitoring):
data = data.melt(id_vars='timestamp', var_name='variable', value_name='value')
data['model_id'] = model_id
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe=data,
minio_repo=self.minio_repository,
model_name=input_data['model_name'],
@@ -534,7 +526,7 @@ class Gates(SientiaMonitoring):
)
@activity.defn(name='format_prediction')
def format_prediction(self, input_data: dict[str, Any]) -> dict:
async def format_prediction(self, input_data: dict[str, Any]) -> dict:
"""
Format prediction data according to configured storage policies.
@@ -566,7 +558,7 @@ class Gates(SientiaMonitoring):
self.info('Formatting prediction...', metadata)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
# Create timestamp column from index and reset index
data['timestamp'] = data.index
@@ -615,7 +607,7 @@ class Gates(SientiaMonitoring):
return data.to_dict()
@activity.defn(name='format_default_prediction')
def format_default_prediction(self, input_data: dict[str, Any]) -> dict:
async def format_default_prediction(self, input_data: dict[str, Any]) -> dict:
"""
Create and format default prediction data for error conditions.
@@ -660,7 +652,7 @@ class Gates(SientiaMonitoring):
return data.to_dict()
@activity.defn(name='format_retrain_report')
def format_retrain_report(self, input_data: dict[str, Any]) -> dict:
async def format_retrain_report(self, input_data: dict[str, Any]) -> dict:
"""
Format retrain report data for storage and audit trail maintenance.
@@ -730,7 +722,7 @@ class Gates(SientiaMonitoring):
return report.to_dict()
@activity.defn(name='write_metrics')
def write_metrics(self, input_data: dict[str, Any]):
async def write_metrics(self, input_data: dict[str, Any]):
"""
Write prediction performance metrics to Prometheus monitoring system.
@@ -767,19 +759,19 @@ class Gates(SientiaMonitoring):
'model_name': metadata['model_name'],
'workflow_name': metadata['workflow_name'],
}
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PREDICTIONS_WRITTEN_COUNT,
tags=core_tags,
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PREDICTION_CONFIDENCE_MONITOR,
method='set',
tags=core_tags,
value=prediction_confidence,
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PREDICTION_RESPONSE_TIME_MONITOR,
method='observe',
tags=core_tags,
@@ -789,7 +781,7 @@ class Gates(SientiaMonitoring):
for server_id, tags in opc_metrics.items():
for tag, response_time in tags.items():
if response_time is not None:
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
method='observe',
tags={
@@ -800,7 +792,7 @@ class Gates(SientiaMonitoring):
value=response_time,
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.PREDICTION_OPC_WRITING_COUNT,
tags={
**core_tags,

View File

@@ -1,23 +1,16 @@
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
import pandas as pd
from pandas import DataFrame, to_datetime
from pandas import to_datetime
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.constants import (
DATETIME_FORMAT,
DATETIME_FORMAT_MS_WITH_TZ,
@@ -25,83 +18,81 @@ 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(SientiaMonitoring):
class MLFlow(MinioManager):
"""
Temporal activities that talk to MLflow through ``SientiaMLflowRepository`` and ``SientiaModel`` wrappers.
MLFlow integration activities for model inference operations.
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``.
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.
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.
The class implements comprehensive error handling and logging for all
MLFlow operations, ensuring reliable model inference in production environments.
Attributes:
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).
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
"""
_MAX_DEBUG_DATAFRAME_ROWS = 100
_DEFAULT_MODEL_ALIAS = 'production'
_REFERENCE_ARTIFACT_CANDIDATES = ('evaluation_data.csv', 'test_data.csv')
def __init__(
self,
mlflow_repository: SientiaMLflowRepository,
plugin_store: PluginStore,
mlflow_host: str,
mlflow_port: int,
mlflow_username: str,
mlflow_password: str,
minio_repository: MinioRepository | None = None,
logger: Logger | None = None,
notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
):
"""
Attach shared MLflow and MinIO clients used by all ML activities in this mixin.
Initialize MLFlow activities with server configuration.
Args:
- 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.
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
Return:
None
Raises:
Exception: If MLFlowRepository initialization fails
"""
self.minio_repository = minio_repository
SientiaMonitoring.__init__(
self,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
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:
"""
Release MinIO manager resources held by the mixin.
Return:
None
Close the MLFlow activity and clean up resources.
"""
if self.minio_repository is not None:
try:
self.minio_repository.close()
finally:
self.minio_repository = None
SientiaMonitoring.shutdown(self)
MinioManager.close(self)
def __del__(self):
self.close()
@@ -124,133 +115,53 @@ class MLFlow(SientiaMonitoring):
metadata,
)
def _detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame:
@activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
Ensure the transform output index is homogeneous and encoded as ``DATETIME_FORMAT_WITH_TZ`` strings.
Transform input data using MLFlow models.
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.
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
Args:
- data: DataFrame whose index carries the time dimension after transform.
- metadata: Workflow metadata for log correlation.
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
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``.
Returns:
dict: Transformed data from MLFlow model
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)
def _resolve_model_alias(self, model_config: dict[str, Any] | None = None) -> str:
"""
Resolve which MLflow alias should be used for model lookup/promotion.
Args:
- model_config: Optional model configuration that may include ``alias``.
Return:
str: Alias name trimmed and normalized; defaults to ``production``.
"""
if not model_config:
return self._DEFAULT_MODEL_ALIAS
alias = str(model_config.get('alias', self._DEFAULT_MODEL_ALIAS)).strip()
return alias or self._DEFAULT_MODEL_ALIAS
@activity.defn(name='request_transform')
def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
Pivot long-format sensor rows, load the production wrapper, and run ``wrapper.transform``.
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: Dict with ``metadata``, ``model_name``, ``data`` (``MinioDataFramePayload``
dict or inline dataframe dict), and optional ``model_config``.
Return:
``MinioDataFramePayload`` with transformed frame and ``success: True``, or a payload
with ``success: False`` and exception details in ``status`` if transform fails.
Exception: If transformation fails or MLFlow model is unavailable
"""
metadata = input_data['metadata']
self.info('Transforming data...', metadata)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Raw input data:', data, metadata)
# Long → wide: keep newest row per (variable, timestamp), then pivot for the wrapper API.
# Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair
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)
@@ -261,25 +172,10 @@ class MLFlow(SientiaMonitoring):
self._debug_dataframe('Processed input data:', data, metadata)
try:
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias=model_alias,
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()},
}
# Request transformation from MLFlow model
response_data = await self.model_monitoring_repository.transform(
model_name, data, model_config, metadata
)
self.debug(
f'Transform raw response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}',
@@ -294,7 +190,7 @@ class MLFlow(SientiaMonitoring):
self.info('Data transformed successfully', metadata)
if not response_data.get('success', False):
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
@@ -305,7 +201,7 @@ class MLFlow(SientiaMonitoring):
logger=self.logger,
)
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
@@ -319,78 +215,58 @@ class MLFlow(SientiaMonitoring):
)
@activity.defn(name='request_predict')
def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
async def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
Load the production wrapper and call ``wrapper.predict`` on the prepared feature frame.
Execute predictions using MLFlow models.
The activity normalizes ``NaN`` to ``None`` for JSON-friendly columns, sets the row index
the same way as ``retrain_model`` (UTC ``DatetimeIndex`` from ``DATETIME_FORMAT_WITH_TZ``),
restores that index on the prediction frame, normalizes the prediction index to
``DATETIME_FORMAT_WITH_TZ`` strings like ``request_transform``, and records ``response_time``.
Non-DataFrame predictions are coerced to a single ``prediction`` column.
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
Args:
- input_data: Same envelope as ``request_transform`` (``metadata``, ``model_name``,
``data``, optional ``model_config`` with ``retention_minutes``).
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
Return:
``MinioDataFramePayload`` with predictions or error status mirroring transform behaviour.
Returns:
dict: Prediction results from MLFlow model
Raises:
Exception: If prediction fails or MLFlow model is unavailable
"""
metadata = input_data['metadata']
self.info('Predicting data...', metadata)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Input data for prediction:', data, metadata)
# Convert numpy.nan to None for model compatibility
data.replace(np.nan, None, inplace=True)
data.index = pd.DatetimeIndex(
to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ, utc=True)
data['timestamp'] = data.index
data['timestamp'] = to_datetime(
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
)
input_index = data.index
try:
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias=model_alias,
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()
predict_data = self._detect_and_parse_datetime_index(predict_data, metadata)
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)}',
@@ -400,7 +276,7 @@ class MLFlow(SientiaMonitoring):
self.info('Data predicted successfully', metadata)
if not response_data.get('success', False):
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
@@ -411,7 +287,7 @@ class MLFlow(SientiaMonitoring):
logger=self.logger,
)
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
@@ -425,23 +301,36 @@ class MLFlow(SientiaMonitoring):
)
@activity.defn(name='retrain_model')
def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Fit an updated wrapper from historical data, then log and register in MLflow.
Retrain MLFlow models with updated training data.
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 →
run ``wrapper.retrain`` outside run timing → ``start_run`` with retrain tags → log input
CSV artifact → ``store_model`` and ``log_params``. Does not promote; the workflow calls
``update_production_model`` after validation.
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
Args:
- input_data: Must include ``metadata``, ``model_name``, ``data`` (payload), and
``model_config`` with at least ``target``.
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
Return:
On success: ``success``, ``experiment`` (``run_id``, ``experiment_id``, ``experiment_name``),
``message``, ``timestamp``. On failure: ``success: False``, error fields, and optional trace.
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
"""
if self.minio_repository is None:
@@ -450,12 +339,13 @@ class MLFlow(SientiaMonitoring):
metadata = input_data['metadata']
try:
# Payload-based retrain input (inline dict or MinIO offloaded).
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='ERROR_LOADING_RETRAIN_DATA',
message=f'Error loading retrain data: {e}',
@@ -481,6 +371,7 @@ class MLFlow(SientiaMonitoring):
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'
@@ -491,130 +382,74 @@ class MLFlow(SientiaMonitoring):
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.index.name = None
data.index = pd.DatetimeIndex(
to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ, utc=True)
data['timestamp'] = data.index
data['timestamp'] = to_datetime(
data['timestamp'], format=DATETIME_FORMAT_WITH_TZ
).dt.strftime(DATETIME_FORMAT)
data['timestamp'] = to_datetime(data['timestamp'], format=DATETIME_FORMAT)
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),
}
try:
model_alias = self._resolve_model_alias(model_config)
mv_src = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name,
alias=model_alias,
)
source_run_id = mv_src.run_id
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias=model_alias,
retention_minutes=0,
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,
)
self.error(trace, metadata=metadata)
# Keep heavy model fitting outside MLflow run timing.
prediction_data = wrapper.retrain(data)
prediction_data.rename(columns={target: 'prediction'}, inplace=True)
# Merge prediction data with retrain data
evaluation_data = pd.merge(
data, prediction_data, left_index=True, right_index=True, how='left'
)
# Rename target column to "target"
evaluation_data.rename(columns={target: 'target'}, inplace=True)
# Reset index and put as column "timestamp"
evaluation_data['timestamp'] = evaluation_data.index
evaluation_data.reset_index(drop=True, inplace=True)
evaluation_data.sort_values(by='timestamp', inplace=True, ascending=True)
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:
tmp_dir = tempfile.mkdtemp(prefix='laborious_retrain_')
try:
raw_csv = Path(tmp_dir) / 'retrain_input.csv'
evaluation_csv = Path(tmp_dir) / 'evaluation_data.csv'
data.to_csv(raw_csv, index=False)
evaluation_data.to_csv(evaluation_csv, index=False)
mlflow.log_artifact(str(raw_csv))
mlflow.log_artifact(str(evaluation_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),
}
return {**retrain_output, 'timestamp': timestamp}
@activity.defn(name='update_production_model')
def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]:
async def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]:
"""
Point the ``production`` alias at the model version registered for the retrain run.
Update production model with newly trained model version.
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.
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
Args:
- input_data: ``metadata``, ``model_name``, and ``experiment`` with ``run_id`` and
``experiment_id`` (as returned from ``retrain_model``).
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
Return:
Dict with ``model_name``, promoted ``version``, ``mlflow_run_id``, ``mlflow_experiment_id``.
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
"""
metadata = input_data['metadata']
model_name = input_data['model_name']
@@ -624,30 +459,16 @@ class MLFlow(SientiaMonitoring):
)
try:
run_id = experiment['run_id']
experiment_id = experiment['experiment_id']
version = self._resolve_model_version_for_run(run_id)
promote_alias = self._resolve_model_alias(input_data.get('model_config'))
self.mlflow_repository.promote_to_alias(
model_name=model_name,
version=version,
alias=promote_alias,
metadata=metadata,
response = await self.model_monitoring_repository.update_production_model(
experiment=experiment, model_name=model_name, metadata=metadata
)
self.info(f'Production model {model_name} updated successfully', metadata)
return {
'model_name': model_name,
'version': version,
'mlflow_run_id': run_id,
'mlflow_experiment_id': experiment_id,
}
return response
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
message=f'Error updating production model {model_name}: {e}',
@@ -658,98 +479,50 @@ class MLFlow(SientiaMonitoring):
self.error(trace, metadata=metadata)
raise e
def _resolve_reference_artifact_name(self, run_id: str) -> str | None:
"""
Pick the first available reference CSV artifact path from the MLflow run.
Candidates are checked in priority order: ``retrain_input.csv``, then ``train_data.csv``.
A path matches when it equals the candidate or ends with ``/<candidate>`` for nested layouts.
Args:
- run_id: MLflow run UUID linked to the production model version.
Return:
Artifact path string for ``download_artifacts``, or ``None`` if no candidate exists.
"""
listed = self.mlflow_repository._client.list_artifacts(run_id)
paths = [file_info.path for file_info in listed]
for candidate in self._REFERENCE_ARTIFACT_CANDIDATES:
for path in paths:
if path == candidate or path.endswith(f'/{candidate}'):
return path
return None
def _find_downloaded_csv(self, tmpdir: str, artifact_name: str) -> Path | None:
"""
Locate a downloaded reference CSV in the temp directory.
Args:
- tmpdir: Directory where ``download_artifacts`` wrote files.
- artifact_name: Basename of the resolved artifact (e.g. ``retrain_input.csv``).
Return:
``Path`` to the CSV file if found, else ``None``.
"""
direct = Path(tmpdir) / artifact_name
if direct.exists():
return direct
matches = list(Path(tmpdir).rglob(artifact_name))
return matches[0] if matches else None
@activity.defn(name='get_reference_data')
def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None:
async def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None:
"""
Download reference training CSV from the MLflow run linked to the production alias.
Get reference data from the MLflow Model Registry.
Resolves ``retrain_input.csv`` or ``train_data.csv`` via artifact listing before download.
``retrain_input.csv`` is preferred when both exist (most recent retrain snapshot). Used by
drift workflows to compare live data against the reference distribution logged with the model.
Timestamps are normalized to ``DATETIME_FORMAT`` string columns before returning records.
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
Args:
- input_data: ``metadata``, ``model_name``, and optional ``model_config`` with ``alias``.
input_data (dict): Input data containing:
- metadata (dict): Workflow execution metadata
- model_name (str): Name of the MLFlow model to get reference data from
Return:
List of row dicts with normalized timestamps, or ``None`` if resolution or load fails.
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
"""
metadata = input_data['metadata']
model_name = input_data['model_name']
artifact = 'evaluation_data.csv'
try:
model_alias = self._resolve_model_alias(input_data.get('model_config'))
mv = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name,
alias=model_alias,
)
run_id = mv.run_id
reference_data = await self.model_monitoring_repository.load_artifact_dataframe(
model_name=model_name, artifact_path=artifact, metadata=metadata
)
artifact_path = self._resolve_reference_artifact_name(run_id)
if artifact_path is None:
self.warning(f'Reference data not found for model {model_name}', metadata)
return None
artifact_name = Path(artifact_path).name
tmpdir = tempfile.mkdtemp(prefix='laborious_eval_')
try:
self.mlflow_repository.download_artifacts(
run_id=run_id,
artifact_path=artifact_path,
dst_path=tmpdir,
metadata=metadata,
)
csv_path = self._find_downloaded_csv(tmpdir, artifact_name)
if csv_path is None:
self.warning(f'Reference data not found for model {model_name}', metadata)
return None
reference_data = pd.read_csv(csv_path)
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)
if reference_data is None:
self.warning(f'Reference data not found for model {model_name}', 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')

View File

@@ -7,15 +7,14 @@ with workflow.unsafe.imports_passed_through():
from typing import Any
import numpy as np
import pandas as pd
from pandas import DataFrame, Index, Series, to_datetime
from pandas import DataFrame, Index, to_datetime
from sientia.ModelAnalysis import ModelAnalysis
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
from sientia_model.analytics.drift_analysis import DriftAnalysis, DriftInsufficientDataError
from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ
from laborious import metrics
from laborious.utils.dataframe_debug import build_dataframe_debug_message
@@ -28,12 +27,9 @@ warnings.filterwarnings(
class ModelMetrics(SientiaMonitoring):
"""
Metrics and statistical analysis activities for the Laborious pipeline.
Metrics activities for the Laborious system.
This class centralizes drift/statistical computations and model-quality
aggregates used by scheduled workflows. Besides producing tabular outputs
for persistence, it also emits operational metrics (count, lag, error)
through ``SientiaMonitoring`` so execution health is observable in runtime.
This class provides activities for writing metrics to the Prometheus monitoring system.
"""
_MAX_DEBUG_DATAFRAME_ROWS = 100
@@ -48,10 +44,7 @@ class ModelMetrics(SientiaMonitoring):
def close(self) -> None:
"""
Shutdown monitoring resources associated with model metrics activities.
This is invoked during worker teardown to flush/close metric controller
internals and prevent dangling telemetry tasks.
Close the model metrics activity and clean up resources.
"""
SientiaMonitoring.shutdown(self)
@@ -76,26 +69,7 @@ class ModelMetrics(SientiaMonitoring):
metadata,
)
def _drift_analyze_stage_error(
self,
exc: Exception,
context: str,
metadata: dict[str, Any],
core_labels: dict[str, Any],
) -> None:
"""
Log analyzer failure for a drift stage and increment the analyze error metric.
Args:
- exc (Exception): Failure raised by ``sientia_model``.
- context (str): Short label for the log line (e.g. univariate detection).
- metadata (dict[str, Any]): Workflow metadata for logging.
- core_labels (dict[str, Any]): Tags from ``get_core_labels`` for metrics.
"""
self.error(f'{context}: {exc}', metadata)
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels)
def get_drift_metrics(
async def get_drift_metrics(
self,
reference_data: DataFrame,
target_data: DataFrame,
@@ -106,38 +80,24 @@ class ModelMetrics(SientiaMonitoring):
metadata: dict[str, Any],
) -> DataFrame:
"""
Compute univariate and multivariate drift outputs and merge them into one dataframe.
The method orchestrates three analysis stages (univariate drift,
multivariate drift, and dataframe projection), emitting lag/count/error
metrics for each stage independently so failures are attributable.
Calculate univariate drift metrics for a model.
Args:
- reference_data (DataFrame): Baseline dataset representing expected behavior.
- target_data (DataFrame): Current analysis dataset to compare against reference.
- target_name (str): Target column name used by ``DriftAnalysis`` config.
- reference_columns (Index): Feature columns evaluated for drift.
- drift_metrics (list[str]): Enabled univariate methods.
- chunk_period (str): Time bucket granularity used by analysis methods.
- metadata (dict[str, Any]): Workflow metadata for logs and notifications.
Return:
DataFrame: Consolidated drift dataframe from ``get_drift_metrics_dataframe`` using
``method`` / ``value`` (and optional ``threshold``, ``drift_type``), ready for
activity-level formatting before Postgres export.
model_analysis (ModelAnalysis): Model analysis object
reference_data (DataFrame): Reference data
target_data (DataFrame): Target data
reference_columns (list[str]): Reference columns
drift_metrics (list[str]): Drift metrics
metadata (dict[str, Any]): Workflow execution metadata
"""
# ``DriftAnalysis`` uses truthiness checks on ``features`` (e.g. ``if not features``);
# a pandas ``Index`` is ambiguous in boolean context — normalize to a list.
feature_names: list[str] = list(reference_columns)
config = {
'target': target_name,
'prediction': 'prediction',
'timestamp': 'timestamp',
'features': feature_names,
'features': reference_columns,
}
drift_analysis = DriftAnalysis(config=config)
model_analysis = ModelAnalysis(config=config)
self._debug_dataframe(
f'Reference data: Size {reference_data.shape}', reference_data, metadata
@@ -148,88 +108,64 @@ class ModelMetrics(SientiaMonitoring):
core_labels = self.get_core_labels(metadata, operation_type='detect_univariate_drift')
start_time = time.time()
try:
univariate_drift = drift_analysis.detect_univariate_drift(
univariate_drift = model_analysis.detect_univariate_drift(
reference_df=reference_data,
analysis_df=target_data,
features=feature_names,
features=reference_columns,
timestamp_col=config['timestamp'],
methods=drift_metrics,
chunk_period=chunk_period,
)
except Exception as e:
if isinstance(e, DriftInsufficientDataError):
raise
self._drift_analyze_stage_error(
e, 'Error detecting univariate drift', metadata, core_labels
self.error(f'Error detecting univariate drift: {e}', metadata)
await self.emit_metric(
metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels
)
raise
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
core_labels = self.get_core_labels(metadata, operation_type='detect_multivariate_drift')
start_time = time.time()
try:
multivariate_drift = drift_analysis.detect_multivariate_drift(
multivariate_drift = model_analysis.detect_multivariate_drift(
reference_df=reference_data,
analysis_df=target_data,
features=feature_names,
features=reference_columns,
timestamp_col=config['timestamp'],
chunk_period=chunk_period,
)
except Exception as e:
if isinstance(e, DriftInsufficientDataError):
raise
self._drift_analyze_stage_error(
e, 'Error detecting multivariate drift', metadata, core_labels
self.error(f'Error detecting multivariate drift: {e}', metadata)
await self.emit_metric(
metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels
)
raise
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
start_time = time.time()
core_labels = self.get_core_labels(metadata, operation_type='get_drift_metrics_dataframe')
try:
drift_df = drift_analysis.get_drift_metrics_dataframe(
drift_df = model_analysis.get_drift_metrics_dataframe(
univariate_drift=univariate_drift,
multivariate_drift=multivariate_drift,
)
except Exception as e:
if isinstance(e, DriftInsufficientDataError):
raise
self._drift_analyze_stage_error(
e, 'Error building drift metrics dataframe', metadata, core_labels
self.error(f'Error getting drift metrics: {e}', metadata)
await self.emit_metric(
metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels
)
raise
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata)
return drift_df
@staticmethod
def _to_naive_utc(series: Series) -> Series:
"""
Parse ``series`` as datetime and return a TZ-naive UTC copy.
``sientia_model.analytics.drift_analysis.DriftAnalysis`` preserves the
timezone of the input dataframe in its outputs, while target rows
loaded from PostgreSQL come in with ``+00:00``. Forcing both sides of
a comparison to TZ-naive UTC keeps ``isin`` / ``floor`` operations
deterministic regardless of how the analyzer constructs its
timestamps.
Args:
- series (Series): Input series containing datetime-parseable values.
Return:
Series: Datetime64 series with ``tz=None`` representing UTC instants.
"""
parsed = to_datetime(series)
if getattr(parsed.dt, 'tz', None) is not None:
parsed = parsed.dt.tz_convert('UTC').dt.tz_localize(None)
return parsed
@activity.defn(name='calculate_drift')
def calculate_drift(self, input_data: dict[str, Any]) -> list[dict]:
async def calculate_drift(self, input_data: dict[str, Any]) -> list[dict]:
"""
Calculate drift metrics for a model.
@@ -259,17 +195,14 @@ class ModelMetrics(SientiaMonitoring):
target_data = target_data.pivot(index='timestamp', columns='variable', values='value')
target_data['timestamp'] = target_data.index
# Keep timestamps as datetime: DriftAnalysis._chunk_dataframe relies on
# ``pd.Grouper(freq=...)`` which rejects string timestamp columns.
target_data['timestamp'] = to_datetime(target_data['timestamp'])
target_data['timestamp'] = target_data['timestamp'].dt.strftime(DATETIME_FORMAT)
target_data = target_data.reset_index(drop=True)
target_data.dropna(inplace=True)
if reference_raw_data is not None:
self.info('Using reference data', metadata)
reference_data = DataFrame(reference_raw_data)
if 'timestamp' in reference_data.columns:
reference_data['timestamp'] = to_datetime(reference_data['timestamp'])
accurate = True
else:
# Get 30% first rows of target_data
@@ -278,7 +211,7 @@ class ModelMetrics(SientiaMonitoring):
reference_data = target_data.head(int(len(target_data) * 0.3))
accurate = False
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING',
message='Using 30% first rows of target data as reference data',
@@ -292,7 +225,7 @@ class ModelMetrics(SientiaMonitoring):
).columns
try:
drift_df = self.get_drift_metrics(
drift_df = await self.get_drift_metrics(
reference_data=reference_data,
target_data=target_data,
target_name=target_name,
@@ -302,30 +235,32 @@ class ModelMetrics(SientiaMonitoring):
metadata=metadata,
)
except Exception as e:
if isinstance(e, DriftInsufficientDataError):
self.error(str(e), metadata)
notification_id = e.notification_id
notification_message = str(e)
else:
self.error(f'Error getting drift metrics: {e}', metadata)
notification_id = 'MODEL_METRICS_GET_DRIFT_METRICS_ERROR'
notification_message = f'Error getting drift metrics: {e}'
self.send_notification(
self.error(f'Error getting drift metrics: {e}', metadata)
await self.send_notification_async(
metadata=metadata,
notification_id=notification_id,
message=notification_message,
notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR',
message=f'Error getting drift metrics: {e}',
block='model_metrics',
level=NotificationLevel.ERROR,
attachment_content=traceback.format_exc(),
)
raise
return []
# Drop chunks whose floored timestamp does not appear in the analysis window.
# ``DriftAnalysis`` chunks over ``analysis_df``; this only excludes rows that
# do not belong to the current target window (e.g. stray merged reference rows).
target_floor = self._to_naive_utc(target_data['timestamp']).dt.floor(chunk_period)
drift_floor = self._to_naive_utc(drift_df['timestamp']).dt.floor(chunk_period)
drift_df = drift_df[drift_floor.isin(target_floor)]
if drift_df.empty:
self.warning('No drift metrics found', metadata)
return []
# Drop unnecessary columns
drift_df.drop(columns=['p_value'], inplace=True)
# Extract timestamps only until minutes
if chunk_period == 'min':
target_timestamps = target_data['timestamp'].apply(lambda x: x[:16])
else:
target_timestamps = target_data['timestamp']
# Drop rows where timestamp is not in target data, to avoid save drift from reference
drift_df = drift_df[drift_df['timestamp'].isin(target_timestamps)]
if drift_df.empty:
self.warning(
@@ -334,36 +269,33 @@ class ModelMetrics(SientiaMonitoring):
)
return []
# Analyzer emits diagnostic columns that are not stored in ``sientia_data.drift_metrics``.
drift_df = drift_df.drop(columns=['threshold', 'drift_type'], errors='ignore')
drift_df['model_id'] = str(model_id)
drift_df['accurate'] = accurate
# ``timestamp`` is overridden with the most recent target instant so
# every persisted row shares a single business timestamp (the run's
# logical "now"), matching what downstream consumers expect.
latest_target_timestamp = self._to_naive_utc(target_data['timestamp']).max()
drift_df['timestamp'] = (
pd.Timestamp(latest_target_timestamp)
.tz_localize('UTC')
.strftime(DATETIME_FORMAT_WITH_TZ)
# Rename columns to match database columns
drift_df.rename(
columns={
'metric': 'method',
'statistic': 'value',
},
inplace=True,
)
# ``chunk_start_date`` / ``chunk_end_date`` may carry nanosecond
# precision (beyond ``timestamptz`` microseconds), so serialize as ISO
# text for the ``text`` Postgres columns.
for column in ('chunk_start_date', 'chunk_end_date'):
drift_df[column] = drift_df[column].apply(
lambda value: pd.Timestamp(value).isoformat() if pd.notna(value) else None
)
# Drop duplicates
drift_df.drop_duplicates(
subset=['timestamp', 'method', 'feature'], keep='first', inplace=True
)
drift_df['model_id'] = model_id
drift_df['accurate'] = accurate
drift_df['timestamp'] = to_datetime(drift_df['timestamp'])
drift_df['timestamp'] = drift_df['timestamp'].dt.tz_localize('UTC')
drift_df['timestamp'] = drift_df['timestamp'].dt.strftime(DATETIME_FORMAT_WITH_TZ)
self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata)
return drift_df.to_dict(orient='records')
@activity.defn(name='calculate_simple_metrics')
def calculate_simple_metrics(self, input_data: dict[str, Any]) -> list[dict]:
async def calculate_simple_metrics(self, input_data: dict[str, Any]) -> list[dict]:
"""
Calculate simple metrics for a model. Metrics available are:
- rmse
@@ -387,7 +319,7 @@ class ModelMetrics(SientiaMonitoring):
metadata = input_data['metadata']
model_id = input_data['model_id']
target_data = DataFrame(input_data['target_data'])
metric_names = input_data['metrics']
metrics = input_data['metrics']
interval_minutes = input_data['interval_minutes']
data_size = target_data.shape[0]
@@ -397,9 +329,9 @@ class ModelMetrics(SientiaMonitoring):
diff = target_data['target'] - target_data['prediction']
diff_squared = diff**2
self.info(f'Calculating simple metrics for model {model_id}: {metric_names}', metadata)
self.info(f'Calculating simple metrics for model {model_id}: {metrics}', metadata)
for metric in metric_names:
for metric in metrics:
if metric == 'rmse':
output_data.append({'metric': 'rmse', 'value': np.sqrt(np.mean(diff_squared))})
elif metric == 'mse':

View File

@@ -82,13 +82,15 @@ class OPC(SientiaMonitoring):
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
self.logger = logger
self.notification_handler = notification_handler
self.opc_servers = opc_servers
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
self.opc_repository: dict[str, OpcRepository] = {}
def init_opc(self):
async def init_opc(self):
"""
Initialize OPC server connections and establish communication channels.
@@ -112,10 +114,10 @@ class OPC(SientiaMonitoring):
the initialization of other OPC servers. Each server is handled
independently to ensure maximum availability.
"""
self.info('Initializing OPC servers...')
self.logger.info('Initializing OPC servers...')
for opc_id, server in self.opc_servers.items():
self.opc_repository[opc_id] = OpcRepository(
opc_id=opc_id,
opc_id=server['id'],
server_name=server['server_name'],
url=server['url'],
logger=self.logger,
@@ -124,12 +126,12 @@ class OPC(SientiaMonitoring):
private_key_path=server['private_key_path'],
server_cert_path=server['server_cert_path'],
notification_handler=self.notification_handler,
reconnection_interval=server.get('reconnection_interval', 60),
reconnection_interval=server['reconnection_interval'],
metrics_controller=self.metrics_controller,
)
is_connected, error_data = self.opc_repository[opc_id].connect()
is_connected, error_data = await self.opc_repository[opc_id].connect()
if not is_connected:
self.send_notification(
await self.send_notification_async(
metadata={
'model_id': '-',
'model_name': '-',
@@ -143,9 +145,11 @@ class OPC(SientiaMonitoring):
attachment_content=error_data.get('attachment_content', None),
)
else:
self.info(f'OPC server {opc_id}:{server["server_name"]} connected successfully.')
self.logger.info(
f'OPC server {opc_id}:{server["server_name"]} connected successfully.'
)
def write_data(
async def write_data(
self,
server_id: str,
tag: str,
@@ -163,11 +167,11 @@ class OPC(SientiaMonitoring):
"""
try:
is_success, info_data = self.opc_repository[server_id].write_data(
is_success, info_data = await self.opc_repository[server_id].write_data(
tag, data, data_type, metadata
)
if not is_success:
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=info_data['notification_id'],
message=info_data['message'],
@@ -179,7 +183,7 @@ class OPC(SientiaMonitoring):
return info_data['response_time'], None
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id=f'WRITE_OPC_{tag_type.upper()}_ERROR',
message=f'Error writing data to OPC server: {e}',
@@ -187,9 +191,9 @@ class OPC(SientiaMonitoring):
level=NotificationLevel.ERROR,
attachment_content=trace,
)
raise
raise e
def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool:
async def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool:
"""
Validate that an OPC server is available and configured for write operations.
@@ -212,7 +216,7 @@ class OPC(SientiaMonitoring):
"""
if self.opc_repository.get(server_id) is None:
message = f'OPC server {server_id} not found to perform write operation.'
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='OPC_SERVER_NOT_FOUND',
message=message,
@@ -223,7 +227,7 @@ class OPC(SientiaMonitoring):
return False
return True
def _write_tags_from_config(
async def _write_tags_from_config(
self,
server_id: str,
tags_config: dict[str, dict[str, Any]],
@@ -254,7 +258,7 @@ class OPC(SientiaMonitoring):
reconnect_in_progress_seen = False
for tag, tag_config in tags_config.items():
response_time, error_info = self.write_data(
response_time, error_info = await self.write_data(
server_id=server_id,
tag=tag,
data=data.head(1)[data_column].values[0],
@@ -279,7 +283,7 @@ class OPC(SientiaMonitoring):
return response_times, session_bad_seen, session_bad_status, reconnect_in_progress_seen
def manage_output_tags(
async def manage_output_tags(
self,
server_id: str,
config: dict[str, Any],
@@ -329,7 +333,7 @@ class OPC(SientiaMonitoring):
group_session_bad,
group_status,
group_reconnect,
) = self._write_tags_from_config(
) = await self._write_tags_from_config(
server_id=server_id,
tags_config=config[config_key],
data=data,
@@ -355,7 +359,7 @@ class OPC(SientiaMonitoring):
)
@activity.defn(name='write_opc_data')
def write_opc_data(
async def write_opc_data(
self, input_data: dict[str, Any]
) -> tuple[dict[Hashable, Any], dict[str, dict[str, float | None]]]:
"""
@@ -386,10 +390,10 @@ class OPC(SientiaMonitoring):
session_bad_status: str | None = None
reconnect_in_progress_seen = False
opc_metrics: dict[str, dict[str, float | None]] = {}
metrics: dict[str, dict[str, float | None]] = {}
for server_id, config in opc_output_config.items():
if not self.validate_server(server_id, metadata):
if not await self.validate_server(server_id, metadata):
success = False
continue
@@ -399,8 +403,8 @@ class OPC(SientiaMonitoring):
local_session_bad,
local_status,
local_reconnect_in_progress,
) = self.manage_output_tags(server_id, config, data, metadata)
opc_metrics[server_id] = local_response_times
) = await self.manage_output_tags(server_id, config, data, metadata)
metrics[server_id] = local_response_times
local_count = len(local_response_times)
success = success and local_success
if local_session_bad:
@@ -409,10 +413,8 @@ class OPC(SientiaMonitoring):
if local_reconnect_in_progress:
reconnect_in_progress_seen = True
n_pred = len(config.get('prediction_tags') or {})
n_conf = len(config.get('confidence_tags') or {})
self.info(
f'Process completed for OPC server {server_id}: {local_count} of {n_pred} prediction tags and {n_conf} confidence tags',
f'Process completed for OPC server {server_id}: {local_count} of {len(config.get("prediction_tags", []))} prediction tags and {len(config.get("confidence_tags", []))} confidence tags',
metadata,
)
@@ -425,7 +427,7 @@ class OPC(SientiaMonitoring):
opc_status=session_bad_status,
reconnect_in_progress=reconnect_in_progress_seen,
),
opc_metrics,
metrics,
)
def process_confidence(
@@ -489,7 +491,7 @@ class OPC(SientiaMonitoring):
return data.to_dict()
def close(self):
async def aclose(self):
"""
Gracefully shutdown all OPC server connections and cleanup resources.
@@ -510,5 +512,4 @@ class OPC(SientiaMonitoring):
their current state and provides a clean shutdown experience.
"""
for opc in self.opc_repository.values():
opc.disconnect()
self.opc_repository.clear()
await opc.disconnect()

View File

@@ -1,5 +1,7 @@
from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through():
# Extend the Temporal Postgres activities for convenient query -> MinIO export
import traceback
@@ -11,9 +13,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_do.temporal.activities.postgres_sync import Postgres
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.activities.postgres import Postgres
from sientia_do.temporal.constants import now
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
@@ -21,7 +22,7 @@ with workflow.unsafe.imports_passed_through():
_LOAD_QUERY_OFFLOAD_SKIP_KEYS = frozenset({'model_name', 'key_prefix', 'size_threshold_bytes'})
class Storage(Postgres, SientiaMonitoring):
class Storage(Postgres, MinioManager):
"""
Extensions for Postgres activities with a helper to export query results
directly to MinIO as Parquet and return the object name.
@@ -59,16 +60,14 @@ class Storage(Postgres, SientiaMonitoring):
metrics_controller=metrics_controller,
)
self.minio_repository = minio_repository
SientiaMonitoring.__init__(
self,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
@activity.defn(name='load_query_with_minio_offload')
def load_query_with_minio_offload(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
async def load_query_with_minio_offload(
self, input_data: dict[str, Any]
) -> MinioDataFramePayload:
"""
Run the custom SQL load, then return a MinIO-aware dataframe wire dict.
@@ -89,7 +88,7 @@ class Storage(Postgres, SientiaMonitoring):
metadata: dict = input_data.get('metadata', {})
model_name = input_data['model_name']
rows = self.load_custom_query(
rows = await self.load_custom_query(
input_data,
)
if not rows:
@@ -100,7 +99,7 @@ class Storage(Postgres, SientiaMonitoring):
else:
dataframe = pd.DataFrame(rows)
return MinioDataFramePayload.from_dataframe(
return await MinioDataFramePayload.from_dataframe(
dataframe,
minio_repo=self.minio_repository,
workflow_metadata=metadata,
@@ -110,29 +109,15 @@ class Storage(Postgres, SientiaMonitoring):
)
@activity.defn(name='export_payload_to_postgres')
def export_payload_to_postgres(self, input_data: dict[str, Any]) -> dict:
async def export_payload_to_postgres(self, input_data: dict[str, Any]) -> dict:
"""
Resolve a MinIO-aware payload into a DataFrame and persist it into PostgreSQL.
This activity accepts the serialized payload produced by previous steps
(inline dict or MinIO object reference), reconstructs the tabular data,
and delegates the final write to ``export_data_to_postgres`` using the
same input contract expected by the Postgres activity mixin.
Args:
- input_data (dict[str, Any]): Activity input containing ``data`` as a
``MinioDataFramePayload``-compatible dict plus database write options
(schema/table/on_conflict/metadata and related fields).
Return:
dict: Result dictionary returned by ``export_data_to_postgres``, including
success status and optional write diagnostics.
Export a payload to PostgreSQL.
"""
metadata = input_data.get('metadata')
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = payload.retrieve(self.minio_repository, metadata)
data = await payload.retrieve(self.minio_repository, metadata)
return self.export_data_to_postgres(
return await self.export_data_to_postgres(
{
**input_data,
'data': data,
@@ -140,7 +125,7 @@ class Storage(Postgres, SientiaMonitoring):
)
@activity.defn(name='cleanup_minio_objects_expired')
def cleanup_minio_objects_expired(self, input_data: dict[str, Any]) -> dict[str, Any]:
async def cleanup_minio_objects_expired(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Delete objects under the given prefixes that are older than the retention window.
@@ -169,7 +154,7 @@ class Storage(Postgres, SientiaMonitoring):
'deleted_count': 0,
}
try:
keys = self.minio_repository.list_objects(
keys = await self.minio_repository.list_objects(
prefix=prefix,
recursive=True,
metadata=metadata,
@@ -181,7 +166,7 @@ class Storage(Postgres, SientiaMonitoring):
continue
if ts >= cutoff:
continue
self.minio_repository.delete_file(
await self.minio_repository.delete_file(
object_name=key,
metadata=metadata,
)
@@ -199,7 +184,7 @@ class Storage(Postgres, SientiaMonitoring):
report['deleted_count'] += 1
except Exception as e:
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
message=f'Error cleaning up MinIO objects: {e}',
@@ -216,16 +201,10 @@ class Storage(Postgres, SientiaMonitoring):
return report
def close(self) -> None:
"""
Shutdown Storage resources in deterministic order.
"""Close Storage resources (MinIO client and Postgres engine)."""
if hasattr(self, 'engine'):
Postgres.close(self)
MinioManager.close(self)
The method first closes Postgres resources via ``Postgres.close`` (engine,
sessions, and monitoring hooks), then closes the optional MinIO repository
and clears the local reference to avoid accidental reuse after shutdown.
"""
Postgres.close(self)
if self.minio_repository is not None:
try:
self.minio_repository.close()
finally:
self.minio_repository = None
def __del__(self):
self.close()

View File

@@ -5,63 +5,29 @@ from typing import Any
def build_mlflow_config() -> dict[str, Any]:
"""
Read MLflow tracking and registry credentials from the environment.
Build MLFlow server configuration from environment variables.
Used by ``Activities`` when constructing ``SientiaMLflowRepository``. The ``url`` value is the
same string workers and notebooks should use for ``MLFLOW_TRACKING_URI``-style clients.
This function constructs an MLFlow configuration dictionary from
environment variables with sensible defaults for local development.
It handles server connection and authentication parameters.
Environment Variables:
MLFLOW_URL: Host with scheme
MLFLOW_USERNAME: Basic-auth or service user (default: aignosi)
MLFLOW_PASSWORD: Password or token (default: aignosi)
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)
Return:
dict[str, Any]: ``url``, ``username``, ``password``.
Returns:
dict: MLFlow configuration dictionary with all required parameters
"""
return {
'url': getenv('MLFLOW_URL', 'http://localhost:5080'),
'host': getenv('MLFLOW_HOST', 'http://localhost'),
'port': int(getenv('MLFLOW_PORT', '5080')),
'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.
@@ -107,15 +73,14 @@ def build_minio_config() -> dict[str, Any]:
Build MinIO (S3-compatible) configuration from environment variables.
Environment Variables:
MINIO_ENDPOINT_URL: Host:port or URL for the S3 API (default: http://localhost:9000)
MINIO_ENDPOINT: MinIO endpoint including scheme (default: http://localhost:9000)
MINIO_ACCESS_KEY: Access key (default: minioadmin)
MINIO_SECRET_KEY: Secret key (default: minioadmin)
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``.
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
"""
return {
'endpoint_url': getenv('MINIO_ENDPOINT_URL', 'http://localhost:9000'),

View File

@@ -21,7 +21,7 @@ from typing import Any, Literal
from pandas import DataFrame, read_parquet
from sientia_do.observability.logger import Logger
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, DATETIME_FORMAT_WITH_TZ, now
# Keys that are part of the serialized wire format (not arbitrary metadata).
@@ -89,7 +89,7 @@ class MinioDataFramePayload:
metadata: dict[str, Any] | None = None,
) -> None:
"""
Emit a debug message only when a logger instance is available.
Emit debug logs only when logger is provided
Args:
- logger (Logger | None): Logger instance used for debug messages
@@ -182,14 +182,7 @@ class MinioDataFramePayload:
def cleanup_prefix(self) -> str | None:
"""
Return the MinIO prefix eligible for retention cleanup.
Cleanup is only applicable when payload data was offloaded to MinIO
(``object_key`` present and inline ``data`` absent). Inline-only payloads
return ``None`` because there is no object tree to prune.
Return:
str | None: Prefix used by cleanup listing, or ``None`` when cleanup does not apply.
Return True if cleanup is enabled for this payload.
"""
if self.object_key is not None and self.data is None:
return self.object_prefix
@@ -197,19 +190,12 @@ class MinioDataFramePayload:
def has_data(self) -> bool:
"""
Indicate whether the payload contains retrievable tabular content.
A payload is considered non-empty when either inline ``data`` exists
(and is not an empty dict) or an ``object_key`` is available for MinIO
download.
Return:
bool: ``True`` when data can be retrieved, ``False`` otherwise.
Return True if the payload has some data internally or in MinIO.
"""
return (self.data is not None and self.data != {}) or self.object_key is not None
@classmethod
def from_dataframe(
async def from_dataframe(
cls,
dataframe: DataFrame | None,
minio_repo: MinioRepository,
@@ -288,7 +274,7 @@ class MinioDataFramePayload:
dataframe.to_parquet(parquet_buffer, engine='pyarrow', index=True)
file_bytes = parquet_buffer.getvalue()
upload_result = minio_repo.upload_file(
upload_result = await minio_repo.upload_file(
file_bytes=file_bytes,
relative_key=object_key,
metadata=workflow_metadata,
@@ -313,7 +299,7 @@ class MinioDataFramePayload:
status=status,
)
def retrieve(
async def retrieve(
self,
minio_repo: MinioRepository,
workflow_metadata: dict[str, Any] | None = None,
@@ -350,7 +336,7 @@ class MinioDataFramePayload:
f'MinioDataFramePayload.retrieve downloading object from MinIO: {self.object_key}',
workflow_metadata,
)
file_bytes = minio_repo.download_file(
file_bytes = await minio_repo.download_file(
object_name=self.object_key, metadata=workflow_metadata
)
df = read_parquet(BytesIO(file_bytes))

View File

@@ -0,0 +1,32 @@
from sientia_do.notifications.handlers import NotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository import MinioRepository
class MinioManager(SientiaMonitoring):
minio_repository: MinioRepository | None = None
def __init__(
self,
minio_repository: MinioRepository | None = None,
logger: Logger | None = None,
notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
):
if self.minio_repository is None:
self.minio_repository = minio_repository
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
def close(self) -> None:
"""
Close the MinioManager and clean up resources.
"""
if self.minio_repository is not None:
try:
self.minio_repository.close()
finally:
self.minio_repository = None
SientiaMonitoring.shutdown(self)

File diff suppressed because it is too large Load Diff

View File

@@ -1,23 +1,14 @@
"""
Synchronous OPC UA client repository using asyncua ``sync`` API.
``asyncua.sync.Client`` runs the asyncio stack on a background thread so Temporal
activities and other callers stay blocking while preserving the same session
lifecycle, security policy, reconnect semantics, and write error classification
as the async ``main`` implementation at ``fcc8920a8be4`` (async → sync/thread conversion).
"""
import asyncio
import json
import threading
import time
import traceback
from datetime import datetime
from pathlib import Path
from typing import Any
from asyncua import ua
from asyncua.crypto import security_policies
from asyncua.sync import Client
from asyncua import Client
from asyncua.crypto.security_policies import SecurityPolicyBasic256
from asyncua.ua import DataValue, Variant, VariantType
from asyncua.ua.uaerrors import UaStatusCodeError
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
@@ -80,7 +71,7 @@ def _opc_authentication_token_str(client: Client | None) -> str:
if client is None:
return 'unknown'
try:
proto = client.aio_obj.uaclient.protocol
proto = client.uaclient.protocol
if proto is None:
return 'unknown'
tok = getattr(proto, 'authentication_token', None)
@@ -143,35 +134,28 @@ def _model_labels_from_write_metadata(metadata: dict[str, Any] | None) -> dict[s
data_type_map = {
'float': {
'converter': float,
'opc_type': ua.VariantType.Float,
'opc_type': VariantType.Float,
},
'double': {
'converter': float,
'opc_type': ua.VariantType.Double,
'opc_type': VariantType.Double,
},
'int': {
'converter': int,
'opc_type': ua.VariantType.Int32,
'opc_type': VariantType.Int32,
},
'bool': {
'converter': bool,
'opc_type': ua.VariantType.Boolean,
'opc_type': VariantType.Boolean,
},
'str': {
'converter': str,
'opc_type': ua.VariantType.String,
'opc_type': VariantType.String,
},
}
class OpcRepository(SientiaMonitoring):
"""
Synchronous OPC UA repository for connect/disconnect and typed writes.
Uses ``asyncua.sync.Client`` with the same session metrics, Tier-1 Bad* reconnect,
and structured write error payloads as the async repository on ``main``.
"""
def __init__(
self,
opc_id: str,
@@ -208,9 +192,9 @@ class OpcRepository(SientiaMonitoring):
'schedule_name': '-',
}
self._last_write_mono: float | None = None
self._connection_lock = threading.Lock()
self._session_ready = threading.Event()
self._reconnect_thread: threading.Thread | None = None
self._connection_lock = asyncio.Lock()
self._session_ready = asyncio.Event()
self._reconnect_task: asyncio.Task[None] | None = None
self._allow_reconnect = True
def _opc_debug_tags(self, session_id: str) -> dict[str, str]:
@@ -241,7 +225,7 @@ class OpcRepository(SientiaMonitoring):
if self.client is None:
return False
try:
proto = self.client.aio_obj.uaclient.protocol
proto = self.client.uaclient.protocol
return proto is not None and proto.state != 'closed'
except Exception:
return False
@@ -273,9 +257,9 @@ class OpcRepository(SientiaMonitoring):
'level': NotificationLevel.WARNING,
}
def set_security(self) -> None:
async def set_security(self) -> None:
"""
Configure certificates and timeouts on the sync asyncua client.
Configure certificates and timeouts on the asyncua client.
Raises:
ValueError: If cert paths or client are missing.
@@ -294,20 +278,18 @@ class OpcRepository(SientiaMonitoring):
self.client.application_uri = self.server_uri
self.info('Setting security...', self.metadata)
self.client.set_security(
security_policies.SecurityPolicyBasic256,
str(cert),
str(private_key),
None,
str(server_cert) if server_cert else None,
await self.client.set_security(
SecurityPolicyBasic256,
certificate=str(cert),
private_key=str(private_key),
server_certificate=str(server_cert) if server_cert else None,
)
aio = self.client.aio_obj
aio.secure_channel_timeout = OPC_UA_SESSION_AND_CHANNEL_TIMEOUT_MS
aio.session_timeout = OPC_UA_SESSION_AND_CHANNEL_TIMEOUT_MS
self.client.secure_channel_timeout = OPC_UA_SESSION_AND_CHANNEL_TIMEOUT_MS
self.client.session_timeout = OPC_UA_SESSION_AND_CHANNEL_TIMEOUT_MS
def _create_client(self) -> None:
async def _create_client(self) -> None:
"""
Instantiate the sync Client and apply security when configured.
Instantiate the asyncua Client and apply security when configured.
Caller must hold _connection_lock. Does not open a UA session.
@@ -320,19 +302,16 @@ class OpcRepository(SientiaMonitoring):
'call disconnect() before creating a new client'
)
self.client = Client(self.url, timeout=10)
aio = self.client.aio_obj
if hasattr(aio, 'watchdog_intervall'):
aio.watchdog_intervall = 50
aio.name = self.pod_id
aio.description = self.pod_id
self.client = Client(self.url, timeout=10, watchdog_intervall=50) # type: ignore[attr-defined]
self.client.name = self.pod_id
self.client.application_name = self.pod_id
pod_uri = self.pod_id.replace('-', ':')
self.client.application_uri = pod_uri
aio.product_uri = pod_uri
self.client.product_uri = pod_uri
if self.cert_path:
self.set_security()
await self.set_security()
def _open_session(self) -> tuple[bool, dict[str, Any]]:
async def _open_session(self) -> tuple[bool, dict[str, Any]]:
"""
Open the OPC UA session on the existing client.
@@ -360,31 +339,30 @@ class OpcRepository(SientiaMonitoring):
'pod_id': self.pod_id,
'server_name': self.server_name,
}
self.emit_metric_sync(metrics.OPC_CONNECTIONS_TOTAL, tags)
await self.emit_metric(metrics.OPC_CONNECTIONS_TOTAL, tags)
try:
self.client.connect()
await self.client.connect()
aio = self.client.aio_obj
session_id = _opc_authentication_token_str(self.client)
revised_session_timeout_ms = int(aio.session_timeout)
revised_secure_channel_timeout_ms = int(aio.secure_channel_timeout)
revised_session_timeout_ms = int(self.client.session_timeout)
revised_secure_channel_timeout_ms = int(self.client.secure_channel_timeout)
self.info(
f'OPC new session connected opc_server_id={self.id} session_id={session_id} '
f'revised_session_timeout_ms={revised_session_timeout_ms} '
f'revised_secure_channel_timeout_ms={revised_secure_channel_timeout_ms}',
self.metadata,
)
self.emit_metric_sync(
await self.emit_metric(
metrics.OPC_SESSION_CREATED_TOTAL, self._opc_debug_tags(session_id)
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.OPC_SESSION_REVISED_TIMEOUT_MS,
method='set',
tags=self._opc_debug_tags(session_id),
value=revised_session_timeout_ms,
)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.OPC_CONNECTION_STATUS,
method='set',
tags={**tags, 'server_url': self.url},
@@ -396,10 +374,10 @@ class OpcRepository(SientiaMonitoring):
return True, {}
except Exception as e:
self._disconnect_locked()
await self._disconnect_locked()
trace = traceback.format_exc()
self.error(trace, self.metadata)
self.emit_metric_sync(metrics.OPC_CONNECTIONS_FAILED, tags)
await self.emit_metric(metrics.OPC_CONNECTIONS_FAILED, tags)
return False, {
'notification_id': f'OPC_CONNECTION_ERROR_{self.id}',
'message': f'Failed to connect to OPC server: {e}',
@@ -408,7 +386,7 @@ class OpcRepository(SientiaMonitoring):
'attachment_content': trace,
}
def _connect_locked(self) -> tuple[bool, dict[str, Any]]:
async def _connect_locked(self) -> tuple[bool, dict[str, Any]]:
"""
Create the client when absent, then open a UA session.
@@ -426,10 +404,10 @@ class OpcRepository(SientiaMonitoring):
'call disconnect() before connecting again'
)
if self.client is None:
self._create_client()
return self._open_session()
await self._create_client()
return await self._open_session()
def _disconnection_fallback(self) -> list[dict[str, Any]]:
async def _disconnection_fallback(self) -> list[dict[str, Any]]:
"""
Try up to five times to disconnect from the OPC UA server.
"""
@@ -441,7 +419,7 @@ class OpcRepository(SientiaMonitoring):
f'Disconnecting from OPC UA server, attempt {i + 1} of 5',
self.metadata,
)
self.client.disconnect()
await self.client.disconnect()
return []
except Exception as e:
self.error(
@@ -455,10 +433,10 @@ class OpcRepository(SientiaMonitoring):
'traceback': traceback.format_exc(),
}
)
time.sleep(self.disconnection_interval * i)
await asyncio.sleep(self.disconnection_interval * i)
return error_stack
def _disconnect_locked(self) -> None:
async def _disconnect_locked(self) -> None:
"""
Tear down the current session and client.
@@ -475,11 +453,11 @@ class OpcRepository(SientiaMonitoring):
f'OPC disconnecting opc_server_id={self.id} session_id={session_id}',
self.metadata,
)
self.emit_metric_sync(metrics.OPC_SESSION_CLOSED_TOTAL, self._opc_debug_tags(session_id))
await self.emit_metric(metrics.OPC_SESSION_CLOSED_TOTAL, self._opc_debug_tags(session_id))
errors = self._disconnection_fallback()
errors = await self._disconnection_fallback()
if errors:
self.send_notification(
await self.send_notification_async(
metadata=self.metadata,
notification_id=f'OPC_DISCONNECTION_ERROR_{self.id}',
message='Failed to disconnect from OPC server in 5 attempts.',
@@ -490,7 +468,7 @@ class OpcRepository(SientiaMonitoring):
else:
self.warning(f'Disconnected from OPC server {self.id} successfully', self.metadata)
self.emit_metric_sync(
await self.emit_metric(
metric_object=metrics.OPC_CONNECTION_STATUS,
method='set',
tags={
@@ -502,7 +480,7 @@ class OpcRepository(SientiaMonitoring):
)
self.client = None
def _reconnect_locked(self) -> tuple[bool, dict[str, Any]]:
async def _reconnect_locked(self) -> tuple[bool, dict[str, Any]]:
"""
Close the current session and open a new one.
@@ -512,32 +490,32 @@ class OpcRepository(SientiaMonitoring):
tuple[bool, dict[str, Any]]: Result from _connect_locked after teardown.
"""
self.last_reconnection_time = datetime.now()
self._disconnect_locked()
return self._connect_locked()
await self._disconnect_locked()
return await self._connect_locked()
def connect(self) -> tuple[bool, dict[str, Any]]:
async def connect(self) -> tuple[bool, dict[str, Any]]:
"""
Open an OPC UA session under the connection lock (worker initialization).
"""
with self._connection_lock:
async with self._connection_lock:
self.info(
f'Starting connection to OPC server {self.id}:{self.server_name}...',
self.metadata,
)
return self._connect_locked()
return await self._connect_locked()
def disconnect(self) -> None:
async def disconnect(self) -> None:
"""
Gracefully disconnect from the OPC server under the connection lock.
Disables background reconnect so late writes during worker shutdown do not
respawn sessions.
"""
with self._connection_lock:
async with self._connection_lock:
self._allow_reconnect = False
self._disconnect_locked()
await self._disconnect_locked()
def validate_connection(self) -> tuple[bool, dict[str, Any]]:
async def validate_connection(self) -> tuple[bool, dict[str, Any]]:
"""
Read-only check that the asyncua protocol is open.
@@ -551,20 +529,20 @@ class OpcRepository(SientiaMonitoring):
self.error(f'OPC server {self.id} is not connected', self.metadata)
return False, self._not_connected_error()
def _reconnect_thread_in_progress(self) -> bool:
def _reconnect_task_in_progress(self) -> bool:
"""
Return whether a background reconnect thread is currently running.
Return whether a background reconnect task is currently running.
Return:
bool: True when a reconnect thread exists and is alive.
bool: True when a reconnect task exists and has not finished.
"""
return self._reconnect_thread is not None and self._reconnect_thread.is_alive()
return self._reconnect_task is not None and not self._reconnect_task.done()
def _start_reconnect(self, reason: str, session_id: str) -> None:
async def _start_reconnect(self, reason: str, session_id: str) -> None:
"""
Schedule a background reconnect when allowed by interval and thread state.
Schedule a background reconnect when allowed by interval and task state.
Clears _session_ready before starting the thread. No-op when _allow_reconnect is
Clears _session_ready before starting the task. No-op when _allow_reconnect is
False, the reconnection window has not elapsed, or a reconnect is already running.
Args:
@@ -580,7 +558,7 @@ class OpcRepository(SientiaMonitoring):
self.metadata,
)
return
if self._reconnect_thread_in_progress():
if self._reconnect_task_in_progress():
self.warning(
f'OPC reconnect skipped reason=in_progress opc_server_id={self.id} '
f'reconnect_reason={reason}',
@@ -594,29 +572,24 @@ class OpcRepository(SientiaMonitoring):
f'old_session_id={session_id}',
self.metadata,
)
self._reconnect_thread = threading.Thread(
target=self._run_reconnect,
args=(reason, session_id),
daemon=True,
)
self._reconnect_thread.start()
self._reconnect_task = asyncio.create_task(self._run_reconnect(reason, session_id))
def _run_reconnect(self, reason: str, session_id: str) -> None:
async def _run_reconnect(self, reason: str, session_id: str) -> None:
"""
Tear down and re-establish the OPC UA session under the connection lock.
Background task that tears down and re-establishes the OPC UA session.
Args:
reason (str): Trigger for reconnect (OPC status or ProtocolClosed).
session_id (str): Previous session token string.
session_id (str): Previous session token string for logging.
"""
try:
with self._connection_lock:
async with self._connection_lock:
self.info(
f'OPC reconnect started reconnect_reason={reason} opc_server_id={self.id} '
f'old_session_id={session_id}',
self.metadata,
)
success, error = self._reconnect_locked()
success, error = await self._reconnect_locked()
if not success:
self.error(
f'OPC reconnect failed reconnect_reason={reason} opc_server_id={self.id}',
@@ -631,7 +604,7 @@ class OpcRepository(SientiaMonitoring):
)
self.error(traceback.format_exc(), self.metadata)
def _log_write_inter_arrival(self, session_id: str, node: str) -> None:
async def _log_write_inter_arrival(self, session_id: str, node: str) -> None:
"""
Log elapsed wall time since the previous successful OPC write on this repository.
@@ -648,18 +621,26 @@ class OpcRepository(SientiaMonitoring):
self.metadata,
)
if self.client is not None:
session_timeout_ms = float(self.client.aio_obj.session_timeout)
session_timeout_ms = float(self.client.session_timeout)
if session_timeout_ms > 0 and delta_s > (session_timeout_ms / 1000.0):
self.emit_metric_sync(
await self.emit_metric(
metrics.OPC_WRITE_INTER_ARRIVAL_OVER_SESSION_TIMEOUT_TOTAL,
self._opc_debug_tags(session_id),
)
self._last_write_mono = now
def _emit_opc_write_metric(
async def _emit_opc_write_metric(
self, session_id: str, result: str, metadata: dict[str, Any] | None
) -> None:
self.emit_metric_sync(
"""
Emit opc_write_attempts_total for a single write attempt outcome.
Args:
session_id (str): OPC UA session token string, or "unknown".
result (str): Outcome label (OK, OPC status name, ProtocolClosed, etc.).
metadata (dict[str, Any] | None): Write context for model_id/model_name labels.
"""
await self.emit_metric(
metrics.OPC_WRITE_ATTEMPTS_TOTAL,
{
**self._opc_debug_tags(session_id),
@@ -677,6 +658,20 @@ class OpcRepository(SientiaMonitoring):
opc_error_kind: str | None = None,
opc_status: str | None = None,
) -> dict[str, Any]:
"""
Build a structured error dict returned from failed write_data paths.
Args:
notification_id (str): Stable notification identifier.
message (str): Human-readable failure message.
level (NotificationLevel): Severity for downstream notifications.
attachment_content (str | None): Optional traceback or diagnostic text.
opc_error_kind (str | None): Classifier (session_bad, connection_lost, etc.).
opc_status (str | None): OPC UA status name or synthetic reason.
Return:
dict[str, Any]: Error payload consumed by the OPC activity layer.
"""
payload: dict[str, Any] = {
'notification_id': notification_id,
'message': message,
@@ -691,7 +686,7 @@ class OpcRepository(SientiaMonitoring):
payload['opc_status'] = opc_status
return payload
def _handle_tier1_bad(
async def _handle_tier1_bad(
self,
exc: BaseException,
session_id: str,
@@ -715,14 +710,14 @@ class OpcRepository(SientiaMonitoring):
opc_status = _opc_status_from_exception(exc)
trace = traceback.format_exc()
self.error(trace, metadata)
self._emit_opc_write_metric(session_id, opc_status, metadata)
await self._emit_opc_write_metric(session_id, opc_status, metadata)
self.error(
f'OPC write failed opc_status={opc_status} opc_server_id={self.id} '
f'session_id={session_id} model_id={metadata.get("model_id", "unknown")} '
f'model_name={metadata.get("model_name", "unknown")} node={node} phase={phase}',
metadata,
)
self._start_reconnect(opc_status, session_id)
await self._start_reconnect(opc_status, session_id)
return False, self._write_failure_payload(
notification_id=f'OPC_WRITE_DATA_ERROR_{self.id}',
message=f'Failed to {phase} on OPC server: {exc} | metadata: {metadata}',
@@ -731,33 +726,11 @@ class OpcRepository(SientiaMonitoring):
opc_status=opc_status,
)
@staticmethod
def _write_node_value(
node_obj: Any,
ua_data: ua.DataValue,
data: Any,
variant_type: ua.VariantType,
) -> None:
async def _write_reconnect_in_progress(
self, metadata: dict[str, Any]
) -> tuple[bool, dict[str, Any]]:
"""
Write a DataValue to a node, falling back to set_value when write_value is unavailable.
Args:
node_obj: Sync or async node wrapper from asyncua.
ua_data (ua.DataValue): Encoded value for write_value.
data: Scalar converted value for set_value fallback.
variant_type (ua.VariantType): OPC UA type for set_value fallback.
"""
if hasattr(node_obj, 'write_value'):
try:
node_obj.write_value(ua_data)
return
except (AttributeError, TypeError):
pass
node_obj.set_value(data, variant_type)
def _write_reconnect_in_progress(self, metadata: dict[str, Any]) -> tuple[bool, dict[str, Any]]:
"""
Fail a write because a background reconnect thread is already running.
Fail a write because a background reconnect task is already running.
Args:
metadata (dict[str, Any]): Write context passed through to the activity.
@@ -765,7 +738,7 @@ class OpcRepository(SientiaMonitoring):
Return:
tuple[bool, dict[str, Any]]: (False, error info with opc_error_kind reconnect_in_progress).
"""
self._emit_opc_write_metric('unknown', 'ReconnectInProgress', metadata)
await self._emit_opc_write_metric('unknown', 'ReconnectInProgress', metadata)
self.warning(
f'OPC write rejected reconnect_in_progress opc_server_id={self.id} '
f'model_id={metadata.get("model_id", "unknown")} '
@@ -780,7 +753,7 @@ class OpcRepository(SientiaMonitoring):
'opc_error_kind': 'reconnect_in_progress',
}
def _write_connection_lost(
async def _write_connection_lost(
self, metadata: dict[str, Any], opc_status: str
) -> tuple[bool, dict[str, Any]]:
"""
@@ -793,7 +766,7 @@ class OpcRepository(SientiaMonitoring):
Return:
tuple[bool, dict[str, Any]]: (False, error info with opc_error_kind connection_lost).
"""
self._emit_opc_write_metric('unknown', opc_status, metadata)
await self._emit_opc_write_metric('unknown', opc_status, metadata)
return False, self._write_failure_payload(
notification_id=f'OPC_WRITE_CONNECTION_LOST_{self.id}',
message=f'OPC write skipped: connection lost ({opc_status}) | metadata: {metadata}',
@@ -802,7 +775,7 @@ class OpcRepository(SientiaMonitoring):
opc_status=opc_status,
)
def write_data(
async def write_data(
self, node: str, value: Any, data_type: str, metadata: dict[str, Any]
) -> tuple[bool, dict[str, Any]]:
"""
@@ -821,34 +794,35 @@ class OpcRepository(SientiaMonitoring):
tuple[bool, dict[str, Any]]: (True, {response_time}) on success, or
(False, structured error info) on failure.
"""
if self._reconnect_thread_in_progress():
return self._write_reconnect_in_progress(metadata)
if self._reconnect_task_in_progress():
return await self._write_reconnect_in_progress(metadata)
if not self._session_ready.is_set():
session_id = _opc_authentication_token_str(self.client)
self._start_reconnect('SessionNotReady', session_id)
if self._reconnect_thread_in_progress():
return self._write_reconnect_in_progress(metadata)
return self._write_connection_lost(metadata, 'SessionNotReady')
await self._start_reconnect('SessionNotReady', session_id)
if self._reconnect_task_in_progress():
return await self._write_reconnect_in_progress(metadata)
return await self._write_connection_lost(metadata, 'SessionNotReady')
is_connected, _error = self.validate_connection()
is_connected, _error = await self.validate_connection()
if not is_connected:
session_id = _opc_authentication_token_str(self.client)
self._start_reconnect('ProtocolClosed', session_id)
return self._write_connection_lost(metadata, 'ProtocolClosed')
await self._start_reconnect('ProtocolClosed', session_id)
return await self._write_connection_lost(metadata, 'ProtocolClosed')
start_time = time.time()
session_id = _opc_authentication_token_str(self.client)
try:
assert self.client is not None
node_obj = self.client.get_node(node)
node_obj = self.client.get_node(node) # type: ignore[union-attr]
except Exception as e:
if is_reconnectable_opcua_bad(e):
return self._handle_tier1_bad(e, session_id, node, metadata, 'get_node')
return await self._handle_tier1_bad(e, session_id, node, metadata, 'get_node')
trace = traceback.format_exc()
self.error(trace, metadata)
self._emit_opc_write_metric(session_id, f'GetNodeError:{type(e).__name__}', metadata)
await self._emit_opc_write_metric(
session_id, f'GetNodeError:{type(e).__name__}', metadata
)
return False, self._write_failure_payload(
notification_id=f'OPC_WRITE_GET_NODE_ERROR_{self.id}',
message=f'Failed to get node from OPC server: {e} | metadata: {metadata}',
@@ -856,7 +830,7 @@ class OpcRepository(SientiaMonitoring):
)
if data_type not in data_type_map:
self._emit_opc_write_metric(session_id, 'UnsupportedDataType', metadata)
await self._emit_opc_write_metric(session_id, 'UnsupportedDataType', metadata)
return False, self._write_failure_payload(
notification_id=f'OPC_WRITE_DATA_TYPE_ERROR_{self.id}',
message=f'Unsupported data type: {data_type} | metadata: {metadata}',
@@ -864,29 +838,28 @@ class OpcRepository(SientiaMonitoring):
data = data_type_map[data_type]['converter'](value)
self.info(f'Writing {data} - {type(data)} to {node}', metadata)
variant_type = data_type_map[data_type]['opc_type']
ua_data = ua.DataValue(
ua.Variant(data, variant_type),
ua_data = DataValue(
Variant(data, data_type_map[data_type]['opc_type']),
)
try:
self._write_node_value(node_obj, ua_data, data, variant_type)
await node_obj.write_value(ua_data)
end_time = time.time()
response_time = end_time - start_time
except Exception as e:
if is_reconnectable_opcua_bad(e):
return self._handle_tier1_bad(e, session_id, node, metadata, 'write_value')
return await self._handle_tier1_bad(e, session_id, node, metadata, 'write_value')
trace = traceback.format_exc()
self.error(trace, metadata)
self._emit_opc_write_metric(session_id, type(e).__name__, metadata)
await self._emit_opc_write_metric(session_id, type(e).__name__, metadata)
return False, self._write_failure_payload(
notification_id=f'OPC_WRITE_DATA_ERROR_{self.id}',
message=f'Failed to write data to OPC server: {e} | metadata: {metadata}',
attachment_content=trace,
)
self._emit_opc_write_metric(session_id, 'OK', metadata)
self._log_write_inter_arrival(session_id, node)
await self._emit_opc_write_metric(session_id, 'OK', metadata)
await self._log_write_inter_arrival(session_id, node)
return True, {
'response_time': response_time,

View File

@@ -1,32 +1,35 @@
"""
Laborious Worker Module
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.
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.
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.
The worker supports multiple runtime-scoped task queues (via ``sientia_do.temporal.worker.prepare_worker``):
- predictions_batch-{runtime}-queue: Batch prediction workflows (heavy workload)
- minimal_retrain-{runtime}-queue: Model retraining workflows
- drift-{runtime}-queue: Drift detection workflows
- simple_metrics-{runtime}-queue: Simple metrics workflows
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.
``RUNTIME`` must be set; it is passed to every ``prepare_worker`` call. Schedulers must use the
same queue names (breaking change vs legacy ``drift-queue`` / ``simple_metrics-queue``).
Shutdown closes workers, notifications, activities (pools + OPC), and clears ``app_up``.
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
Environment Variables:
- 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.
- RUNTIME: Required non-empty string; suffix for all task queue names
- 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)
"""
from temporalio import client, workflow
@@ -40,21 +43,19 @@ 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.workflows.drift import Drift
from laborious.workflows.minimal_retrain import MinimalRetrain
@@ -71,18 +72,23 @@ SDK_METRICS_PORT = int(os.getenv('HTTP_SDK_METRICS_PORT', '9091'))
async def main():
"""
Run the full worker lifecycle: metrics, notifications, runtime install, workers, gather.
Main entry point for the Laborious worker application.
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``.
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.
Raises:
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).
Exception: Any unhandled exception during worker execution
SystemExit: On graceful shutdown or error conditions
"""
host = os.getenv('TEMPORAL_HOST', 'localhost:7233')
logger = get_logger(__name__)
@@ -97,21 +103,6 @@ async def main():
logger.custom_info(f'Starting Worker with POD_ID: {POD_ID}', metadata)
logger.custom_info('Starting prometheus client...', metadata)
start_prometheus_server()
logger.custom_info('Starting Notification Handler...', metadata)
mongo_config = build_mongodb_config()
notification_handler = NotificationHandler(
connection_string=mongo_config['connection_string'],
database=mongo_config['database_name'],
logger=logger,
project_name=os.getenv('PROJECT_NAME', 'laborious'),
)
metrics_controller = MetricsController(logger=logger)
runtime = os.getenv('RUNTIME', '').strip()
if not runtime:
logger.custom_critical(
@@ -122,58 +113,39 @@ async def main():
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.custom_info('Starting prometheus client...', metadata_runtime)
start_prometheus_server()
logger.custom_info('Starting Notification Handler...', metadata_runtime)
mongo_config = build_mongodb_config()
notification_handler = NotificationHandler(
connection_string=mongo_config['connection_string'],
database=mongo_config['database_name'],
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
project_name=os.getenv('PROJECT_NAME', 'laborious'),
)
if runtime == 'legacy':
to_install_runtime = 'single'
else:
to_install_runtime = runtime
try:
await plugin_store.install_runtime(
runtime_name=to_install_runtime, metadata=metadata_runtime
)
except Exception as exc:
logger.custom_critical(
f'Failed to install runtime {to_install_runtime}: {exc}', metadata_runtime
)
metrics.APP_UP.labels(pod_id=POD_ID).set(0)
sys.exit(1)
logger.custom_info('Starting Activities...', metadata)
logger.custom_info('Starting Activities...', metadata_runtime)
activities = Activities(
postgres_config=build_postgres_config(),
plugin_store=plugin_store,
mlflow_config=build_mlflow_config(),
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)
activities.init_opc()
logger.custom_info('Initializing OPC...', metadata_runtime)
await activities.init_opc()
logger.custom_info(f'Starting SDK Metrics Server on port {SDK_METRICS_PORT}...', metadata)
logger.custom_info(
f'Starting SDK Metrics Server on port {SDK_METRICS_PORT}...',
metadata_runtime,
)
new_runtime = Runtime(
telemetry=TelemetryConfig(
@@ -181,7 +153,7 @@ async def main():
)
)
logger.custom_info(f'Starting Temporal Client at {host}...', metadata)
logger.custom_info(f'Starting Temporal Client at {host}...', metadata_runtime)
temporal_client = await client.Client.connect(
target_host=host,
@@ -189,7 +161,7 @@ async def main():
runtime=new_runtime,
)
logger.custom_info('Starting Workers...', metadata)
logger.custom_info(f'Starting Workers (runtime={runtime})...', metadata_runtime)
workers = [
prepare_worker(
@@ -216,7 +188,7 @@ async def main():
activities.export_data_to_postgres,
],
logger=logger,
runtime='core',
runtime=runtime,
),
prepare_worker(
temporal_client=temporal_client,
@@ -229,7 +201,7 @@ async def main():
activities.export_data_to_postgres,
],
logger=logger,
runtime='core',
runtime=runtime,
),
prepare_worker(
temporal_client=temporal_client,
@@ -267,17 +239,21 @@ async def main():
for w in workers:
handlers.append(w.run())
logger.custom_info('Workers started successfully', metadata)
logger.custom_info('Workers started successfully', metadata_runtime)
exit_code = 0
try:
# This will run the workers and wait for them to complete.
# If an exception occurs in any of the worker handlers, it will be propagated here.
await asyncio.gather(*handlers)
except BaseException as e: # NOSONAR
logger.custom_error(f'An unhandled exception occurred: {e}', metadata)
exit_code = 1
finally:
notification_handler.shutdown()
activities.shutdown()
if notification_handler:
notification_handler.shutdown()
if activities:
await activities.shutdown()
metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN
sys.exit(exit_code)

View File

@@ -161,7 +161,7 @@ class FormatAndExportPrediction:
write_transformed_handler = None
opc_metrics: dict[str, dict[str, float | None]] = {}
opc_metrics = {}
# write to pi web api
if pi_web_api_output_config:

View File

@@ -250,7 +250,7 @@ class PredictionProcess:
async def path_flag_handler(
self,
data: dict[str, Any],
path_flag: str | None,
path_flag: str,
input_data: dict,
confidence: int,
last_timestamp: str,

View File

@@ -116,12 +116,18 @@ python_functions = ["test_*"]
addopts = [
"-v",
"--strict-markers",
# pytest>=9.1 has a known bug where its unraisableexception plugin crashes
# (tracemalloc partially-initialized AttributeError) when 2+ unraisable
# exceptions land close together — e.g. "coroutine was never awaited" from
# AsyncMock-mocked sync methods (metrics_controller, minio_repository) being
# GC'd. Harmless mock artifacts turned into a hard ERROR by the plugin itself.
"-p", "no:unraisableexception",
]
markers = [
"asyncio: marks tests as async",
"integration: marks tests as integration tests",
"unit: marks tests as unit tests",
"opc: marks E2E tests that use in-process asyncua + real OpcRepository",
"opc: marks tests that use the in-process OPC UA server (OpcRepository E2E)",
]
[tool.coverage.run]

18
requirements-light.txt Normal file
View File

@@ -0,0 +1,18 @@
temporalio
psycopg2-binary
sqlalchemy
asyncua==1.0.6
redis
sientia_do>=1.12.2
mlflow
prometheus-client
botocore
boto3
s3fs
pyarrow
kaleido
hyperopt
shap
pycurl
scipy<1.14.0
scikit-learn==1.5.2

View File

@@ -4,7 +4,7 @@ sqlalchemy
asyncua==1.0.6
redis
sientia_do>=1.12.2
sientia_model>=0.8.2
sientia>0.40.0
prometheus-client
botocore
boto3
@@ -15,4 +15,4 @@ hyperopt
shap
pycurl
scipy<1.14.0
scikit-learn==1.5.2
scikit-learn==1.5.2

View File

@@ -1,4 +1,4 @@
sonar.projectKey=Aignosi_sientia-dataops-laborious_temporal_beaec423-6c42-4f26-8134-b676287b499d
sonar.projectKey=Aignosi_sientia-dataops-laborious_temporal_ca1a7039-6db9-49e5-be78-54d29bc93e4f
sonar.projectName=sientia-dataops-laborious_temporal
sonar.sources=laborious
sonar.tests=tests

View File

@@ -1,17 +1,6 @@
import os
from sientia_do.temporal.activities.postgres_sync import Postgres
def _noop_postgres_del(_self):
"""
Unit tests use MagicMock metrics controllers; postgres_sync.Postgres.__del__ calls
close() during GC and triggers async shutdown. Explicit ``close()`` is covered in tests.
"""
return None
Postgres.__del__ = _noop_postgres_del # type: ignore[method-assign]
import sys
from unittest.mock import MagicMock
# The production code converts SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES to int at import-time.
# Tests must set it to a valid integer string to avoid import errors.
@@ -58,9 +47,13 @@ class DummyMinioDataFramePayload:
"""
Pytest configuration file with global mocks for external dependencies.
The historical ``sientia`` package is no longer imported by the codebase;
drift analysis lives in ``sientia_model.analytics.drift_analysis`` and is
imported lazily inside Temporal activities. No global module-level mock is
required here — unit tests that need to control ``DriftAnalysis`` outputs
should patch ``laborious.activities.model_metrics.DriftAnalysis`` directly.
This module mocks the 'sientia' module to avoid requiring its installation
during unit tests. The mock is registered in sys.modules before any test
imports are executed.
"""
# Mock sientia module
sientia_mock = MagicMock()
sientia_mock.ModelAnalysis = MagicMock
sys.modules['sientia'] = sientia_mock
sys.modules['sientia.ModelAnalysis'] = MagicMock()

View File

@@ -1,4 +1,6 @@
from unittest.mock import ANY, MagicMock, patch
from unittest.mock import ANY, AsyncMock, MagicMock, patch
from pytest import mark
from laborious.activities.activities import Activities
from laborious.activities.api import API
@@ -46,8 +48,7 @@ def test___init__(
'secure': False,
}
mlflow_repository = MagicMock()
plugin_store = MagicMock()
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
opc_config = {
'bootstrap_servers': 'localhost:9092',
@@ -66,13 +67,12 @@ def test___init__(
activities = Activities(
postgres_config=postgres_config,
plugin_store=plugin_store,
mlflow_config=mlflow_config,
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,8 +101,10 @@ def test___init__(
mock_mlflow_init.assert_called_once_with(
ANY,
mlflow_repository=mlflow_repository,
plugin_store=plugin_store,
mlflow_host=mlflow_config['host'],
mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'],
minio_repository=mock_minio_repository.return_value,
logger=logger,
notification_handler=notification_handler,
@@ -154,6 +156,7 @@ def test___init__(
)
@mark.asyncio
@patch('laborious.activities.activities.Storage')
@patch('laborious.activities.activities.MLFlow')
@patch('laborious.activities.activities.OPC')
@@ -161,7 +164,7 @@ def test___init__(
@patch('laborious.activities.activities.ModelMetrics')
@patch('laborious.activities.activities.API')
@patch('laborious.activities.activities.MinioRepository')
def test_shutdown(
async def test_shutdown(
_mock_minio_repository,
mock_api_init,
mock_model_metrics_init,
@@ -170,7 +173,7 @@ def test_shutdown(
mock_mlflow_init,
mock_storage_init,
):
mock_opc_init.close = MagicMock()
mock_opc_init.aclose = AsyncMock()
postgres_config = {
'host': 'localhost',
'port': 5432,
@@ -190,8 +193,7 @@ def test_shutdown(
'secure': False,
}
mlflow_repository = MagicMock()
plugin_store = MagicMock()
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
opc_config = {
'bootstrap_servers': 'localhost:9092',
@@ -210,90 +212,18 @@ def test_shutdown(
activities = Activities(
postgres_config=postgres_config,
plugin_store=plugin_store,
mlflow_config=mlflow_config,
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,
)
activities.shutdown()
mock_opc_init.close.assert_called_once()
await activities.shutdown()
mock_opc_init.aclose.assert_called_once()
mock_storage_init.close.assert_called_once()
mock_mlflow_init.close.assert_called_once()
mock_gates_init.close.assert_called_once()
mock_model_metrics_init.close.assert_called_once()
mock_api_init.close.assert_called_once()
@patch('laborious.activities.activities.SientiaMLflowRepository')
@patch('laborious.activities.activities.build_mlflow_config')
@patch('laborious.activities.activities.Storage.__init__')
@patch('laborious.activities.activities.MLFlow.__init__')
@patch('laborious.activities.activities.OPC.__init__')
@patch('laborious.activities.activities.Gates.__init__')
@patch('laborious.activities.activities.ModelMetrics.__init__')
@patch('laborious.activities.activities.API.__init__')
@patch('laborious.activities.activities.MinioRepository')
@patch('laborious.activities.activities.MetricsController')
def test___init___builds_mlflow_repository_when_not_provided(
mock_metrics_controller,
mock_minio_repository,
_mock_api_init,
_mock_model_metrics_init,
_mock_gates_init,
_mock_opc_init,
_mock_mlflow_init,
_mock_storage_init,
mock_build_mlflow_config,
mock_mlflow_repository_cls,
):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
minio_config = {
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'default_bucket': 'test',
'retention_hours': 24,
'secure': False,
}
opc_config = {'bootstrap_servers': 'localhost:9092', 'polling_time': 1000, 'group_id': 'test'}
pi_web_api_config = {'base_url': 'https://pi', 'auth_type': 'bearer', 'auth_token': 'token'}
logger = MagicMock()
notification_handler = MagicMock()
plugin_store = MagicMock()
mock_build_mlflow_config.return_value = {
'url': 'http://mlflow:80',
'username': 'u',
'password': 'p',
}
Activities(
postgres_config=postgres_config,
plugin_store=plugin_store,
minio_config=minio_config,
opc_config=opc_config,
pi_web_api_config=pi_web_api_config,
logger=logger,
notification_handler=notification_handler,
)
mock_build_mlflow_config.assert_called_once()
mock_mlflow_repository_cls.assert_called_once_with(
host='http://mlflow:80',
username='u',
password='p',
logger=logger,
notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
)

View File

@@ -1,6 +1,7 @@
from unittest.mock import ANY, MagicMock, call, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pytest import fixture
import pytest_asyncio
from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel
from laborious.activities.api import API, PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
@@ -70,7 +71,7 @@ def test_get_pi_web_api_core_labels_without_operation_type(mock_pi_web_api_clien
auth_token='test_token',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
with patch.object(
SientiaMonitoring,
@@ -104,7 +105,7 @@ def test_get_pi_web_api_core_labels_with_operation_type(mock_pi_web_api_client):
auth_token='test_token',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
with patch.object(
SientiaMonitoring,
@@ -131,17 +132,17 @@ def test__init__():
auth_token='test_token',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
assert api.pi_web_api_client is not None
@fixture
@pytest_asyncio.fixture
@patch('laborious.activities.api.PIWebAPIClient')
def api(mock_pi_web_api_client):
mock_client = MagicMock()
mock_client.write_value = MagicMock()
mock_client.write_value = AsyncMock()
mock_client.close = MagicMock()
mock_client.base_url = 'https://test-pi-server.com'
mock_pi_web_api_client.return_value = mock_client
@@ -152,12 +153,12 @@ def api(mock_pi_web_api_client):
auth_token='test_token',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
api_instance.send_notification = MagicMock()
api_instance.send_notification_async = AsyncMock()
api_instance.info = MagicMock()
api_instance.error = MagicMock()
api_instance.emit_metric_sync = MagicMock()
api_instance.emit_metric = AsyncMock()
api_instance.get_core_labels = MagicMock(
return_value={
'pod_id': 'test_pod',
@@ -169,8 +170,9 @@ def api(mock_pi_web_api_client):
return api_instance
@mark.asyncio
@patch('laborious.activities.api.DataFrame')
def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
input_data = {
**base_input_data,
'pi_web_api_output_config': {
@@ -188,7 +190,7 @@ def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
[{'WebId': 'web_id_3', 'Errors': []}, {'WebId': 'web_id_4', 'Errors': []}],
]
result = api.write_pi_web_api_data(input_data)
result = await api.write_pi_web_api_data(input_data)
api.pi_web_api_client.write_value.assert_has_calls(
[
@@ -218,8 +220,9 @@ def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
}
@mark.asyncio
@patch('laborious.activities.api.DataFrame')
def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data):
async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data):
mock_dataframe.return_value = _create_mock_dataframe(
{
'prediction': [0.75],
@@ -230,9 +233,9 @@ def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_
api.pi_web_api_client.write_value.side_effect = Exception('Prediction write failed')
result = api.write_pi_web_api_data(base_input_data)
result = await api.write_pi_web_api_data(base_input_data)
api.send_notification.assert_called_once_with(
api.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR',
message="Error writing prediction data to PI Web API: Prediction write failed\n Tags: {'tag1': 'web_id_1'}",
@@ -245,8 +248,9 @@ def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_
assert api.pi_web_api_client.write_value.call_count == 1
@mark.asyncio
@patch('laborious.activities.api.DataFrame')
def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data):
async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data):
mock_dataframe.return_value = _create_mock_dataframe()
# First call succeeds, second fails
@@ -255,9 +259,9 @@ def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_
Exception('Confidence write failed'),
]
result = api.write_pi_web_api_data(base_input_data)
result = await api.write_pi_web_api_data(base_input_data)
api.send_notification.assert_called_once_with(
api.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR',
message="Error writing confidence data to PI Web API: Confidence write failed\n Tags: {'tag2': 'web_id_2'}",
@@ -274,8 +278,9 @@ def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_
assert api.pi_web_api_client.write_value.call_count == 2
@mark.asyncio
@patch('laborious.activities.api.DataFrame')
def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
input_data = {
**base_input_data,
'pi_web_api_output_config': {
@@ -293,7 +298,7 @@ def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
[],
]
result = api.write_pi_web_api_data(input_data)
result = await api.write_pi_web_api_data(input_data)
api.pi_web_api_client.write_value.assert_has_calls(
[
@@ -323,35 +328,15 @@ def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
}
@patch('laborious.activities.api.DataFrame')
def test_write_pi_web_api_data_updates_confidence_and_comments(
mock_dataframe, api, base_input_data
):
mock_dataframe.return_value = _create_mock_dataframe()
api.pi_web_api_client.write_value.side_effect = [
[{'WebId': 'web_id_1', 'Errors': []}],
[{'WebId': 'web_id_2', 'Errors': []}],
]
with patch.object(
api,
'process_pi_web_api_response',
new=MagicMock(side_effect=[(0.33, 'PI warning'), (0, '')]),
) as process_mock:
result = api.write_pi_web_api_data(base_input_data)
assert process_mock.call_count == 2
assert result is not None
@patch('laborious.activities.api.SientiaMonitoring.shutdown')
def test_close(mock_shutdown, api):
@mark.asyncio
async def test_close(api):
api.close()
api.pi_web_api_client.close.assert_called_once()
mock_shutdown.assert_called_once_with(api)
def test_process_pi_web_api_response_success(api):
@mark.asyncio
async def test_process_pi_web_api_response_success(api):
"""Test successful processing of PI Web API response with all tags written."""
response_data = [
{'WebId': 'web_id_1', 'Errors': []},
@@ -365,7 +350,7 @@ def test_process_pi_web_api_response_success(api):
'workflow_name': 'test_workflow',
}
confidence, message = api.process_pi_web_api_response(
confidence, message = await api.process_pi_web_api_response(
response_data=response_data,
tags=tags,
core_labels=core_labels,
@@ -374,9 +359,9 @@ def test_process_pi_web_api_response_success(api):
assert confidence == 0
assert message == ''
assert api.emit_metric_sync.call_count == 2
# Verify that emit_metric_sync was called with correct tags structure
call_args_list = api.emit_metric_sync.call_args_list
assert api.emit_metric.call_count == 2
# Verify that emit_metric was called with correct tags structure
call_args_list = api.emit_metric.call_args_list
assert len(call_args_list) == 2
# Check that all calls include core_labels and tag_name
for call_args in call_args_list:
@@ -384,7 +369,8 @@ def test_process_pi_web_api_response_success(api):
assert call_args.kwargs['tags']['tag_name'] in ['tag1', 'tag2']
def test_process_pi_web_api_response_with_errors(api):
@mark.asyncio
async def test_process_pi_web_api_response_with_errors(api):
"""Test processing response with errors in some tags."""
response_data = [
{'WebId': 'web_id_1', 'Errors': ['Error writing tag']},
@@ -398,7 +384,7 @@ def test_process_pi_web_api_response_with_errors(api):
'workflow_name': 'test_workflow',
}
confidence, message = api.process_pi_web_api_response(
confidence, message = await api.process_pi_web_api_response(
response_data=response_data,
tags=tags,
core_labels=core_labels,
@@ -410,10 +396,11 @@ def test_process_pi_web_api_response_with_errors(api):
message
== "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag2'] tags were written."
)
assert api.emit_metric_sync.call_count == 2
assert api.emit_metric.call_count == 2
def test_process_pi_web_api_response_missing_tags(api):
@mark.asyncio
async def test_process_pi_web_api_response_missing_tags(api):
"""Test processing response when number of written tags doesn't match expected."""
response_data = [
{'WebId': 'web_id_1', 'Errors': []},
@@ -426,7 +413,7 @@ def test_process_pi_web_api_response_missing_tags(api):
'workflow_name': 'test_workflow',
}
confidence, message = api.process_pi_web_api_response(
confidence, message = await api.process_pi_web_api_response(
response_data=response_data,
tags=tags,
core_labels=core_labels,
@@ -438,13 +425,14 @@ def test_process_pi_web_api_response_missing_tags(api):
message
== "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag1'] tags were written."
)
api.send_notification.assert_called_once()
call_args = api.send_notification.call_args
api.send_notification_async.assert_called_once()
call_args = api.send_notification_async.call_args
assert call_args.kwargs['notification_id'] == 'WRITE_PI_WEB_API_PREDICTION_ERROR'
assert call_args.kwargs['level'] == NotificationLevel.ERROR
def test_process_pi_web_api_response_missing_webid(api):
@mark.asyncio
async def test_process_pi_web_api_response_missing_webid(api):
"""Test processing response when WebId is missing in response item."""
response_data = [
{'Errors': []},
@@ -458,7 +446,7 @@ def test_process_pi_web_api_response_missing_webid(api):
'workflow_name': 'test_workflow',
}
confidence, message = api.process_pi_web_api_response(
confidence, message = await api.process_pi_web_api_response(
response_data=response_data,
tags=tags,
core_labels=core_labels,
@@ -473,7 +461,8 @@ def test_process_pi_web_api_response_missing_webid(api):
api.error.assert_any_call('The response did not contain some WebIds', metadata['metadata'])
def test_process_pi_web_api_response_missing_tag_name(api):
@mark.asyncio
async def test_process_pi_web_api_response_missing_tag_name(api):
"""Test processing response when tag name is not found for WebId."""
response_data = [
{'WebId': 'unknown_web_id', 'Errors': []},
@@ -486,7 +475,7 @@ def test_process_pi_web_api_response_missing_tag_name(api):
'workflow_name': 'test_workflow',
}
confidence, message = api.process_pi_web_api_response(
confidence, message = await api.process_pi_web_api_response(
response_data=response_data,
tags=tags,
core_labels=core_labels,

View File

@@ -1,9 +1,8 @@
from unittest.mock import ANY, MagicMock, call, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pandas import DataFrame
from pytest import fixture
from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from laborious.activities.gates import Gates
@@ -18,17 +17,17 @@ def _passthrough_from_dict():
def _minio_payload(retrieve_return, status=None):
"""
Build a MinioDataFramePayload-like test double with retrieve.
Build a MinioDataFramePayload-like test double with async retrieve.
Args:
retrieve_return: Value returned from retrieve(minio_repo, metadata).
retrieve_return: Value returned from await retrieve(minio_repo, metadata).
status: Optional status dict for MLflow response gate (payload.status).
Return:
MagicMock: Object with async retrieve and optional status.
"""
p = MagicMock()
p.retrieve = MagicMock(return_value=retrieve_return)
p.retrieve = AsyncMock(return_value=retrieve_return)
p.status = status
return p
@@ -38,7 +37,7 @@ def gates_activity():
gates = Gates(
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
gates.error = MagicMock()
gates.debug = MagicMock()
@@ -46,7 +45,8 @@ def gates_activity():
gates.warning = MagicMock()
gates.critical = MagicMock()
gates.send_notification = MagicMock()
gates.emit_metric_sync = MagicMock()
gates.send_notification_async = AsyncMock()
gates.emit_metric = AsyncMock()
return gates
@@ -60,7 +60,8 @@ metadata = {
}
def test_input_gate_invalid_filter(gates_activity):
@mark.asyncio
async def test_input_gate_invalid_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -70,7 +71,7 @@ def test_input_gate_invalid_filter(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == (None, 0, '')
@@ -79,8 +80,9 @@ def test_input_gate_invalid_filter(gates_activity):
)
@mark.asyncio
@patch('laborious.activities.gates.input_filter_functions')
def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity):
async def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity):
# Arrange
mock_input_filter_functions.__contains__.return_value = True
mock_input_filter_functions.__getitem__.return_value = MagicMock(
@@ -94,11 +96,11 @@ def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.send_notification.assert_called_once_with(
gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='INTPUT_GATE_ERROR__EMPTY_DATA',
message="Error in filter EMPTY_DATA:{'POLICY': 'STOP', 'CONFIG': {}}: \n Test error",
@@ -108,7 +110,8 @@ def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity
)
def test_input_gate_no_filters(gates_activity):
@mark.asyncio
async def test_input_gate_no_filters(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -118,14 +121,15 @@ def test_input_gate_no_filters(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_input_gate_with_filter(gates_activity):
@mark.asyncio
async def test_input_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -135,14 +139,15 @@ def test_input_gate_with_filter(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == ('STOP', -1, 'Input data with bad quality')
gates_activity.debug.assert_called()
def test_input_gate_with_filter_lowercase_keys(gates_activity):
@mark.asyncio
async def test_input_gate_with_filter_lowercase_keys(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -152,13 +157,14 @@ def test_input_gate_with_filter_lowercase_keys(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == ('STOP', -1, 'Input data with bad quality')
def test_input_gate_with_filter_capitalized_keys(gates_activity):
@mark.asyncio
async def test_input_gate_with_filter_capitalized_keys(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -168,13 +174,14 @@ def test_input_gate_with_filter_capitalized_keys(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == ('STOP', -1, 'Input data with bad quality')
def test_input_gate_with_filter_not_caught(gates_activity):
@mark.asyncio
async def test_input_gate_with_filter_not_caught(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -184,14 +191,15 @@ def test_input_gate_with_filter_not_caught(gates_activity):
}
# Act
result = gates_activity.input_gate(input_data)
result = await gates_activity.input_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_mlflow_response_gate_invalid_filter(gates_activity):
@mark.asyncio
async def test_mlflow_response_gate_invalid_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -205,14 +213,15 @@ def test_mlflow_response_gate_invalid_filter(gates_activity):
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == (None, 0, '')
@mark.asyncio
@patch('laborious.activities.gates.mlflow_response_filter_functions')
def test_mlflow_response_gate_filter_exception(
async def test_mlflow_response_gate_filter_exception(
mock_mlflow_response_filter_functions, gates_activity
):
# Arrange
@@ -232,11 +241,11 @@ def test_mlflow_response_gate_filter_exception(
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.send_notification.assert_called_once_with(
gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER',
message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error",
@@ -246,7 +255,8 @@ def test_mlflow_response_gate_filter_exception(
)
def test_mlflow_response_gate_no_filters(gates_activity):
@mark.asyncio
async def test_mlflow_response_gate_no_filters(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -260,14 +270,15 @@ def test_mlflow_response_gate_no_filters(gates_activity):
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_mlflow_response_gate_with_filter(gates_activity):
@mark.asyncio
async def test_mlflow_response_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -281,15 +292,16 @@ def test_mlflow_response_gate_with_filter(gates_activity):
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == ('STOP', -1, 'API error occurred')
gates_activity.debug.assert_called()
gates_activity.send_notification.assert_called()
gates_activity.send_notification_async.assert_called()
def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
@mark.asyncio
async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -303,13 +315,14 @@ def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == ('STOP', -1, 'API error occurred')
def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
@mark.asyncio
async def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -323,14 +336,15 @@ def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
}
# Act
result = gates_activity.mlflow_response_gate(input_data)
result = await gates_activity.mlflow_response_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_mlflow_content_gate_invalid_filter(gates_activity):
@mark.asyncio
async def test_mlflow_content_gate_invalid_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -341,14 +355,17 @@ def test_mlflow_content_gate_invalid_filter(gates_activity):
}
# Act
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
# Assert
assert result == (None, 0, '')
@mark.asyncio
@patch('laborious.activities.gates.mlflow_content_filter_functions')
def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functions, gates_activity):
async def test_mlflow_content_gate_filter_exception(
mock_mlflow_content_filter_functions, gates_activity
):
# Arrange
mock_mlflow_content_filter_functions.__contains__.return_value = True
mock_mlflow_content_filter_functions.__getitem__.return_value = MagicMock(
@@ -363,12 +380,12 @@ def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functio
}
# Act
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
gates_activity.send_notification.assert_called_once_with(
gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR',
message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test error",
@@ -378,7 +395,8 @@ def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functio
)
def test_mlflow_content_gate_no_filters(gates_activity):
@mark.asyncio
async def test_mlflow_content_gate_no_filters(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -389,14 +407,15 @@ def test_mlflow_content_gate_no_filters(gates_activity):
}
# Act
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_mlflow_content_gate_with_filter(gates_activity):
@mark.asyncio
async def test_mlflow_content_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -407,15 +426,16 @@ def test_mlflow_content_gate_with_filter(gates_activity):
}
# Act
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
# Assert
assert result == ('STOP', -1, 'Transformed data not passed the content filter')
gates_activity.debug.assert_called()
gates_activity.send_notification.assert_called()
gates_activity.send_notification_async.assert_called()
def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
@mark.asyncio
async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -426,14 +446,15 @@ def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
}
# Act
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
# Assert
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_mlflow_content_gate_filter_returns_false(gates_activity):
@mark.asyncio
async def test_mlflow_content_gate_filter_returns_false(gates_activity):
input_data = {
**metadata,
'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}},
@@ -442,7 +463,7 @@ def test_mlflow_content_gate_filter_returns_false(gates_activity):
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
result = gates_activity.mlflow_content_gate(input_data)
result = await gates_activity.mlflow_content_gate(input_data)
assert result == (None, 0, '')
gates_activity.debug.assert_called()
@@ -504,7 +525,8 @@ def test_get_prediction_store_policy_valid_policy(gates_activity):
assert policy_value == 1
def test_format_prediction_no_timestamp(gates_activity):
@mark.asyncio
async def test_format_prediction_no_timestamp(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -523,7 +545,7 @@ def test_format_prediction_no_timestamp(gates_activity):
}
# Act
result = gates_activity.format_prediction(input_data)
result = await gates_activity.format_prediction(input_data)
# Assert
assert result['prediction'] == {0: 1}
@@ -535,7 +557,8 @@ def test_format_prediction_no_timestamp(gates_activity):
assert result['comments'] == {0: ''}
def test_format_prediction_with_timestamp_erl(gates_activity):
@mark.asyncio
async def test_format_prediction_with_timestamp_erl(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -562,7 +585,7 @@ def test_format_prediction_with_timestamp_erl(gates_activity):
}
# Act
result = gates_activity.format_prediction(input_data)
result = await gates_activity.format_prediction(input_data)
# Assert
assert result['prediction'] == {0: 2, 1: 1}
@@ -574,7 +597,8 @@ def test_format_prediction_with_timestamp_erl(gates_activity):
assert result['comments'] == {0: '', 1: ''}
def test_format_prediction_with_timestamp_lts(gates_activity):
@mark.asyncio
async def test_format_prediction_with_timestamp_lts(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -601,7 +625,7 @@ def test_format_prediction_with_timestamp_lts(gates_activity):
}
# Act
result = gates_activity.format_prediction(input_data)
result = await gates_activity.format_prediction(input_data)
# Assert
assert result['prediction'] == {0: 3, 1: 2}
@@ -613,7 +637,8 @@ def test_format_prediction_with_timestamp_lts(gates_activity):
assert result['comments'] == {0: '', 1: ''}
def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
@mark.asyncio
async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -638,15 +663,16 @@ def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
gates_activity.get_prediction_store_policy = MagicMock(return_value=('invalid', 1))
try:
gates_activity.format_prediction(input_data)
await gates_activity.format_prediction(input_data)
except ValueError as e:
assert str(e) == 'Invalid policy type: invalid'
else:
raise AssertionError('Expected ValueError')
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity):
@mark.asyncio
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
@@ -665,7 +691,7 @@ def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity)
}
# Act
result = gates_activity.format_transformed_data(input_data)
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert result is payload_result
@@ -679,8 +705,9 @@ def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity)
gates_activity.info.assert_called()
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity):
@mark.asyncio
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
@@ -705,7 +732,7 @@ def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activi
}
# Act
result = gates_activity.format_transformed_data(input_data)
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert result is payload_result
@@ -719,8 +746,9 @@ def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activi
gates_activity.info.assert_called()
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity):
@mark.asyncio
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
@@ -732,7 +760,7 @@ def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity)
}
# Act
result = gates_activity.format_transformed_data(input_data)
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert result is payload_result
@@ -746,7 +774,8 @@ def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity)
gates_activity.info.assert_called()
def test_format_default_prediction(gates_activity):
@mark.asyncio
async def test_format_default_prediction(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -757,7 +786,7 @@ def test_format_default_prediction(gates_activity):
}
# Act
result = gates_activity.format_default_prediction(input_data)
result = await gates_activity.format_default_prediction(input_data)
# Assert
assert result['prediction'] == {0: 0}
@@ -770,7 +799,8 @@ def test_format_default_prediction(gates_activity):
gates_activity.debug.assert_called()
def test_format_retrain_report(gates_activity):
@mark.asyncio
async def test_format_retrain_report(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -789,7 +819,7 @@ def test_format_retrain_report(gates_activity):
}
# Act
result = gates_activity.format_retrain_report(input_data)
result = await gates_activity.format_retrain_report(input_data)
# Assert
assert result['model_id'] == {0: 'test_model'}
@@ -801,7 +831,8 @@ def test_format_retrain_report(gates_activity):
assert result['mlflow_experiment_id'] == {0: 'test_mlflow_experiment_id'}
def test_format_retrain_report_failure(gates_activity):
@mark.asyncio
async def test_format_retrain_report_failure(gates_activity):
# Arrange
input_data = {
**metadata,
@@ -820,7 +851,7 @@ def test_format_retrain_report_failure(gates_activity):
}
# Act
result = gates_activity.format_retrain_report(input_data)
result = await gates_activity.format_retrain_report(input_data)
# Assert
assert result['model_id'] == {0: 'test_model'}
@@ -834,8 +865,9 @@ def test_format_retrain_report_failure(gates_activity):
gates_activity.debug.assert_called()
@mark.asyncio
@patch('laborious.activities.gates.metrics')
def test_write_metrics(mock_metrics, gates_activity):
async def test_write_metrics(mock_metrics, gates_activity):
"""Test write_metrics method."""
input_data = {
**metadata,
@@ -846,7 +878,7 @@ def test_write_metrics(mock_metrics, gates_activity):
},
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}},
}
gates_activity.write_metrics(input_data)
await gates_activity.write_metrics(input_data)
core_tags = {
'pod_id': gates_activity.pod_id,
'runtime': gates_activity.runtime,
@@ -854,7 +886,7 @@ def test_write_metrics(mock_metrics, gates_activity):
'model_name': metadata['metadata']['model_name'],
'workflow_name': metadata['metadata']['workflow_name'],
}
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT,
@@ -862,7 +894,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR,
@@ -872,7 +904,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR,
@@ -882,7 +914,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
@@ -894,7 +926,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
@@ -908,7 +940,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
@@ -920,7 +952,7 @@ def test_write_metrics(mock_metrics, gates_activity):
),
]
)
gates_activity.emit_metric_sync.assert_has_calls(
gates_activity.emit_metric.assert_has_calls(
[
call(
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
@@ -936,8 +968,9 @@ def test_write_metrics(mock_metrics, gates_activity):
)
@mark.asyncio
@patch('laborious.activities.gates.metrics')
def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity):
async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity):
"""Test write_metrics method with None response_time in opc_metrics."""
input_data = {
**metadata,
@@ -948,10 +981,10 @@ def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity)
},
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': None}},
}
gates_activity.write_metrics(input_data)
await gates_activity.write_metrics(input_data)
# Verify that metrics for tag1 are emitted
gates_activity.emit_metric_sync.assert_any_call(
gates_activity.emit_metric.assert_any_call(
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
method='observe',
tags={
@@ -969,38 +1002,7 @@ def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity)
# Verify that metrics for tag2 (with None response_time) are NOT emitted
calls = [
c
for c in gates_activity.emit_metric_sync.call_args_list
for c in gates_activity.emit_metric.call_args_list
if len(c[1].get('tags', {})) > 0 and c[1]['tags'].get('tag') == 'tag2'
]
assert len(calls) == 0, 'Metrics should not be emitted for None response_time'
@patch.object(SientiaMonitoring, 'shutdown')
def test_close_disposes_minio_repository(mock_shutdown):
"""
``Gates.close`` should close the optional MinIO client and clear the repository reference.
"""
minio = MagicMock()
gates = Gates(
minio_repository=minio,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
)
gates.close()
minio.close.assert_called_once()
assert gates.minio_repository is None
mock_shutdown.assert_called_once_with(gates)
@patch.object(SientiaMonitoring, 'shutdown')
def test_close_without_minio_repository(mock_shutdown):
"""When no MinIO repository is configured, ``close`` only shuts down monitoring."""
gates = Gates(
minio_repository=None,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
)
gates.close()
mock_shutdown.assert_called_once_with(gates)

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,6 +1,6 @@
from unittest.mock import ANY, MagicMock, call, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
import pytest
import pytest_asyncio
from pandas import DataFrame
from pytest import mark
from sientia_do.notifications.models import NotificationLevel
@@ -32,26 +32,27 @@ def test__init__():
opc_servers=servers,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
assert opc.opc_servers == servers
assert opc.opc_repository == {}
@mark.asyncio
@patch('laborious.activities.opc.OpcRepository')
@patch('laborious.activities.opc.OPC.send_notification')
def test_init_opc(mock_send_notification, mock_opc_repository):
@patch('laborious.activities.opc.OPC.send_notification_async')
async def test_init_opc(mock_send_notification, mock_opc_repository):
mock_logger = MagicMock()
mock_metrics_controller = MagicMock()
mock_metrics_controller = AsyncMock()
server1 = MagicMock(
connect=MagicMock(return_value=(True, {})), write_data=MagicMock(return_value=(True, {}))
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
)
server2 = MagicMock(
connect=MagicMock(return_value=(True, {})), write_data=MagicMock(return_value=(True, {}))
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
)
server3 = MagicMock(
connect=MagicMock(
connect=AsyncMock(
return_value=(
False,
{
@@ -63,7 +64,7 @@ def test_init_opc(mock_send_notification, mock_opc_repository):
},
)
),
write_data=MagicMock(return_value=(True, {})),
write_data=AsyncMock(return_value=(True, {})),
)
mock_opc_repository.side_effect = [server1, server2, server3]
mock_notification_handler = MagicMock()
@@ -105,7 +106,7 @@ def test_init_opc(mock_send_notification, mock_opc_repository):
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
opc.init_opc()
await opc.init_opc()
assert opc.opc_servers == servers
assert opc.logger == mock_logger
@@ -170,9 +171,9 @@ def test_init_opc(mock_send_notification, mock_opc_repository):
)
@pytest.fixture
@pytest_asyncio.fixture
@patch('laborious.activities.opc.OpcRepository')
def opc(mock_opc_repository):
async def opc(mock_opc_repository):
servers = {
'server1': {
'id': 'server1',
@@ -186,17 +187,18 @@ def opc(mock_opc_repository):
}
}
mock_opc_repository.return_value.write_data = MagicMock(return_value=(True, {}))
mock_opc_repository.return_value.connect = MagicMock(return_value=(True, {}))
mock_opc_repository.return_value.write_data = AsyncMock(return_value=(True, {}))
mock_opc_repository.return_value.connect = AsyncMock(return_value=(True, {}))
opc = OPC(
opc_servers=servers,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
opc.init_opc()
await opc.init_opc()
opc.send_notification = MagicMock()
opc.emit_metric_sync = MagicMock()
opc.send_notification_async = AsyncMock()
opc.emit_metric = AsyncMock()
return opc
@@ -209,10 +211,11 @@ WRITE_DATA_CASES = [
@mark.parametrize('tag,data_type,data', WRITE_DATA_CASES)
def test_write_data_success(opc, tag, data_type, data):
@mark.asyncio
async def test_write_data_success(opc, tag, data_type, data):
opc.opc_repository['server1'].write_data.return_value = (True, {'response_time': 0.1})
response_time, error_info = opc.write_data(
response_time, error_info = await opc.write_data(
server_id='server1',
tag=tag,
data=data,
@@ -225,7 +228,8 @@ def test_write_data_success(opc, tag, data_type, data):
opc.opc_repository['server1'].write_data.assert_called_once_with(tag, data, data_type, metadata)
def test_write_data_failed(opc):
@mark.asyncio
async def test_write_data_failed(opc):
opc.opc_repository['server1'].write_data.return_value = (
False,
{
@@ -237,7 +241,7 @@ def test_write_data_failed(opc):
},
)
response_time, error_info = opc.write_data(
response_time, error_info = await opc.write_data(
server_id='server1',
tag='tag1',
data=50,
@@ -248,7 +252,7 @@ def test_write_data_failed(opc):
assert response_time is None
assert error_info is not None
opc.send_notification.assert_called_once_with(
opc.send_notification_async.assert_called_once_with(
metadata=metadata,
notification_id='OPC_WRITE_DATA_ERROR_server1',
message='Failed to write data to OPC server: Test error',
@@ -258,11 +262,12 @@ def test_write_data_failed(opc):
)
def test_write_data_exception(opc):
@mark.asyncio
async def test_write_data_exception(opc):
opc.opc_repository['server1'].write_data.side_effect = Exception('Test error')
try:
opc.write_data(
await opc.write_data(
server_id='server1',
tag='tag1',
data=50,
@@ -272,7 +277,7 @@ def test_write_data_exception(opc):
)
except Exception:
opc.send_notification.assert_called_once_with(
opc.send_notification_async.assert_called_once_with(
metadata=metadata,
notification_id='WRITE_OPC_PREDICTION_ERROR',
message='Error writing data to OPC server: Test error',
@@ -339,12 +344,13 @@ def test_apply_opc_write_error(
assert result == expected
def test_write_tags_from_config_prediction_success(opc):
opc.write_data = MagicMock(return_value=(0.1, None))
@mark.asyncio
async def test_write_tags_from_config_prediction_success(opc):
opc.write_data = AsyncMock(return_value=(0.1, None))
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
tags_config = {'tag1': {'data_type': 'float'}}
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
response_times, session_bad, opc_status, reconnect = await opc._write_tags_from_config(
server_id='server1',
tags_config=tags_config,
data=data,
@@ -368,12 +374,13 @@ def test_write_tags_from_config_prediction_success(opc):
)
def test_write_tags_from_config_confidence_success(opc):
opc.write_data = MagicMock(return_value=(0.2, None))
@mark.asyncio
async def test_write_tags_from_config_confidence_success(opc):
opc.write_data = AsyncMock(return_value=(0.2, None))
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
tags_config = {'tag2': {'data_type': 'float'}}
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
response_times, session_bad, opc_status, reconnect = await opc._write_tags_from_config(
server_id='server1',
tags_config=tags_config,
data=data,
@@ -397,11 +404,12 @@ def test_write_tags_from_config_confidence_success(opc):
)
def test_write_tags_from_config_write_failure(opc):
opc.write_data = MagicMock(return_value=(None, {}))
@mark.asyncio
async def test_write_tags_from_config_write_failure(opc):
opc.write_data = AsyncMock(return_value=(None, {}))
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
response_times, session_bad, opc_status, reconnect = await opc._write_tags_from_config(
server_id='server1',
tags_config={'tag1': {'data_type': 'float'}},
data=data,
@@ -417,8 +425,9 @@ def test_write_tags_from_config_write_failure(opc):
assert reconnect is False
def test_write_tags_from_config_session_bad(opc):
opc.write_data = MagicMock(
@mark.asyncio
async def test_write_tags_from_config_session_bad(opc):
opc.write_data = AsyncMock(
return_value=(
None,
{
@@ -429,7 +438,7 @@ def test_write_tags_from_config_session_bad(opc):
)
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
response_times, session_bad, opc_status, reconnect = await opc._write_tags_from_config(
server_id='server1',
tags_config={'tag1': {'data_type': 'float'}},
data=data,
@@ -445,8 +454,9 @@ def test_write_tags_from_config_session_bad(opc):
assert reconnect is False
def test_write_tags_from_config_reconnect_in_progress(opc):
opc.write_data = MagicMock(
@mark.asyncio
async def test_write_tags_from_config_reconnect_in_progress(opc):
opc.write_data = AsyncMock(
return_value=(
None,
{'opc_error_kind': 'reconnect_in_progress'},
@@ -454,7 +464,7 @@ def test_write_tags_from_config_reconnect_in_progress(opc):
)
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
response_times, session_bad, opc_status, reconnect = await opc._write_tags_from_config(
server_id='server1',
tags_config={'tag1': {'data_type': 'float'}},
data=data,
@@ -470,8 +480,9 @@ def test_write_tags_from_config_reconnect_in_progress(opc):
assert reconnect is True
def test_manage_output_tags_success(opc):
opc._write_tags_from_config = MagicMock(
@mark.asyncio
async def test_manage_output_tags_success(opc):
opc._write_tags_from_config = AsyncMock(
side_effect=[
({'tag1': 0.1}, False, None, False),
({'tag2': 0.1}, False, None, False),
@@ -483,7 +494,7 @@ def test_manage_output_tags_success(opc):
'confidence_tags': {'tag2': {'data_type': 'float'}},
}
output_data, opc_metrics, session_bad, opc_status, reconnect = opc.manage_output_tags(
output_data, opc_metrics, session_bad, opc_status, reconnect = await opc.manage_output_tags(
server_id='server1',
config=config,
data=data,
@@ -495,11 +506,12 @@ def test_manage_output_tags_success(opc):
assert session_bad is False
assert opc_status is None
assert reconnect is False
assert opc._write_tags_from_config.call_count == 2
assert opc._write_tags_from_config.await_count == 2
def test_manage_output_tags_failed(opc):
opc._write_tags_from_config = MagicMock(
@mark.asyncio
async def test_manage_output_tags_failed(opc):
opc._write_tags_from_config = AsyncMock(
side_effect=[
({'tag1': 0.1}, False, None, False),
({'tag2': None}, False, None, False),
@@ -511,7 +523,7 @@ def test_manage_output_tags_failed(opc):
'confidence_tags': {'tag2': {'data_type': 'float'}},
}
output_data, opc_metrics, _, _, _ = opc.manage_output_tags(
output_data, opc_metrics, _, _, _ = await opc.manage_output_tags(
server_id='server1',
config=config,
data=data,
@@ -522,12 +534,13 @@ def test_manage_output_tags_failed(opc):
assert opc_metrics == {'tag1': 0.1, 'tag2': None}
def test_manage_output_tags_do_nothing(opc):
opc._write_tags_from_config = MagicMock()
@mark.asyncio
async def test_manage_output_tags_do_nothing(opc):
opc._write_tags_from_config = AsyncMock()
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
config = {'_invalid_key': {'tag1': {'data_type': 'float'}}}
output_data, opc_metrics, _, _, _ = opc.manage_output_tags(
output_data, opc_metrics, _, _, _ = await opc.manage_output_tags(
server_id='server1',
config=config,
data=data,
@@ -539,8 +552,9 @@ def test_manage_output_tags_do_nothing(opc):
opc._write_tags_from_config.assert_not_called()
@mark.asyncio
@patch('laborious.activities.opc.DataFrame')
def test_write_opc_data_success(mock_dataframe, opc):
async def test_write_opc_data_success(mock_dataframe, opc):
# Arrange
input_data = {
**metadata,
@@ -554,12 +568,12 @@ def test_write_opc_data_success(mock_dataframe, opc):
}
# Act
opc.manage_output_tags = MagicMock(
opc.manage_output_tags = AsyncMock(
return_value=(True, {'tag1': 0.1, 'tag2': 0.2}, False, None, False)
)
opc.process_confidence = MagicMock(return_value={'data': 'data'})
output_data, opc_metrics = opc.write_opc_data(input_data)
output_data, opc_metrics = await opc.write_opc_data(input_data)
# Assert
assert output_data == {'data': 'data'}
@@ -580,7 +594,8 @@ def test_write_opc_data_success(mock_dataframe, opc):
)
def test_write_opc_data_empty_config(opc):
@mark.asyncio
async def test_write_opc_data_empty_config(opc):
# Arrange
input_data = {
**metadata,
@@ -590,14 +605,15 @@ def test_write_opc_data_empty_config(opc):
}
# Act
opc.write_opc_data(input_data)
await opc.write_opc_data(input_data)
# Assert
opc.opc_repository['server1'].write_data.assert_not_called()
def test_write_opc_data_no_validate_server(opc):
opc.validate_server = MagicMock(return_value=False)
@mark.asyncio
async def test_write_opc_data_no_validate_server(opc):
opc.validate_server = AsyncMock(return_value=False)
input_data = {
**metadata,
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
@@ -610,7 +626,7 @@ def test_write_opc_data_no_validate_server(opc):
}
# Act
opc.write_opc_data(input_data)
await opc.write_opc_data(input_data)
# Assert
opc.opc_repository['server1'].write_data.assert_not_called()
@@ -649,8 +665,9 @@ def test_process_confidence_generic_failure(opc):
assert result['comments'][0] == OPC_WRITTING_ERROR_MESSAGE
def test_manage_output_tags_merges_error_flags(opc):
opc._write_tags_from_config = MagicMock(
@mark.asyncio
async def test_manage_output_tags_merges_error_flags(opc):
opc._write_tags_from_config = AsyncMock(
side_effect=[
({'tag1': None}, True, 'BadSessionIdInvalid', False),
({'tag2': 0.2}, False, None, True),
@@ -668,7 +685,7 @@ def test_manage_output_tags_merges_error_flags(opc):
session_bad_seen,
opc_status,
reconnect_in_progress,
) = opc.manage_output_tags('server1', config, data, metadata['metadata'])
) = await opc.manage_output_tags('server1', config, data, metadata['metadata'])
assert success is False
assert session_bad_seen is True
@@ -708,13 +725,14 @@ def test_process_confidence_concatenates_multiple_comments(opc):
)
def test_validate_server(opc):
assert opc.validate_server('server1', metadata) is True
assert opc.validate_server('server2', metadata) is False
@mark.asyncio
async def test_validate_server(opc):
assert await opc.validate_server('server1', metadata) is True
assert await opc.validate_server('server2', metadata) is False
def test_close(opc):
repo = opc.opc_repository['server1']
repo.disconnect = MagicMock(return_value=True)
opc.close()
repo.disconnect.assert_called_once()
@mark.asyncio
async def test_close(opc):
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
await opc.aclose()
opc.opc_repository['server1'].disconnect.assert_called_once()

View File

@@ -1,11 +1,10 @@
import datetime
import os
from unittest.mock import ANY, MagicMock, patch
from unittest.mock import ANY, AsyncMock, MagicMock, patch
from pytest import fixture, raises
from pytest import fixture, mark, raises
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.temporal.activities.postgres_sync import Postgres
from sientia_do.temporal.activities.postgres import Postgres
from laborious.activities.storage import Storage
@@ -18,15 +17,6 @@ def _passthrough_from_dict():
yield
@fixture(autouse=True)
def _patch_monitoring_shutdown():
"""
Avoid running real async SientiaMonitoring.shutdown when Storage.close runs inside tests.
"""
with patch.object(SientiaMonitoring, 'shutdown') as mock_shutdown:
yield mock_shutdown
metadata = {
'metadata': {
'model_id': 'test_model_id',
@@ -52,7 +42,7 @@ def storage(mock_minio_repository):
minio_repository=mock_minio_repository.return_value,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
@@ -60,7 +50,7 @@ def storage(mock_minio_repository):
def test___init___not_hasattr(mock_minio_repository):
logger = MagicMock()
notification_handler = MagicMock()
metrics_controller = MagicMock()
metrics_controller = AsyncMock()
minio_repo = mock_minio_repository.return_value
storage = Storage(
host='localhost',
@@ -87,7 +77,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
storage.minio_repository = None
logger = MagicMock()
notification_handler = MagicMock()
metrics_controller = MagicMock()
metrics_controller = AsyncMock()
storage.__init__(
host='localhost',
port=5432,
@@ -121,55 +111,54 @@ def test___init___done_repository(mock_minio_repository, storage):
minio_repository=mock_minio_repository.return_value,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
mock_minio_repository.assert_not_called()
assert storage.minio_repository is not None
def test_close(storage, _patch_monitoring_shutdown):
def test_close(storage):
storage.minio_repository = MagicMock()
storage.close()
assert storage.minio_repository is None
_patch_monitoring_shutdown.assert_called_once_with(storage)
def test_close_when_minio_repository_already_none(storage, _patch_monitoring_shutdown):
"""Closing without an initialized MinIO repository skips MinIO teardown."""
storage.minio_repository = None
def test___del__(storage):
storage.close = MagicMock()
storage.close()
storage.__del__()
assert storage.minio_repository is None
_patch_monitoring_shutdown.assert_called_once_with(storage)
storage.close.assert_called_once()
def test_load_query_with_minio_offload_no_rows(storage):
storage.load_custom_query = MagicMock(return_value=None)
@mark.asyncio
async def test_load_query_with_minio_offload_no_rows(storage):
storage.load_custom_query = AsyncMock(return_value=None)
storage_result = {'success': False}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=MagicMock,
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = storage.load_query_with_minio_offload(
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result == storage_result
mock_from_dataframe.assert_called_once()
mock_from_dataframe.assert_awaited_once()
def test_load_query_with_minio_offload_inline(storage):
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
@mark.asyncio
async def test_load_query_with_minio_offload_inline(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=MagicMock,
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = storage.load_query_with_minio_offload(
result = await storage.load_query_with_minio_offload(
{
**metadata,
'query': 'SELECT 1',
@@ -178,42 +167,44 @@ def test_load_query_with_minio_offload_inline(storage):
}
)
assert result == storage_result
mock_from_dataframe.assert_called_once()
mock_from_dataframe.assert_awaited_once()
def test_load_query_with_minio_offload_minio(storage):
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
@mark.asyncio
async def test_load_query_with_minio_offload_minio(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage_result = {'success': True, 'data': None, 'object_key': 'object-key'}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=MagicMock,
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = storage.load_query_with_minio_offload(
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result == storage_result
mock_from_dataframe.assert_called_once()
mock_from_dataframe.assert_awaited_once()
@mark.asyncio
@patch.dict(os.environ, {'SIENTIA_MINIO_RETENTION_HOURS': '1'})
@patch('laborious.activities.storage.now')
def test_cleanup_minio_objects_expired(mock_now, storage):
async def test_cleanup_minio_objects_expired(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = MagicMock(
storage.minio_repository.list_objects = AsyncMock(
return_value=[
'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
'sientia/streamlit-connectors/training_datasets/m/m-initial-2025-01-10_12-00-00.parquet',
]
)
storage.minio_repository.delete_file = MagicMock()
storage.send_notification = MagicMock()
storage.minio_repository.delete_file = AsyncMock()
storage.send_notification_async = AsyncMock()
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 1
assert result['failed_count'] == 0
@@ -233,65 +224,72 @@ def test_cleanup_minio_objects_expired(mock_now, storage):
)
def test_load_query_with_minio_offload_minio_not_initialized(storage):
@mark.asyncio
async def test_load_query_with_minio_offload_minio_not_initialized(storage):
storage.minio_repository = None
with raises(ValueError, match='Minio repository not initialized'):
storage.load_query_with_minio_offload({**metadata, 'query': 'SELECT 1', 'model_name': 'm'})
await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm'}
)
def test_export_payload_to_postgres(storage):
payload = MagicMock()
payload.retrieve = MagicMock(return_value=MagicMock())
storage.export_data_to_postgres = MagicMock(return_value={'success': True})
@mark.asyncio
async def test_export_payload_to_postgres(storage):
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=MagicMock())
storage.export_data_to_postgres = AsyncMock(return_value={'success': True})
result = storage.export_payload_to_postgres(
result = await storage.export_payload_to_postgres(
{**metadata, 'data': payload, 'schema': 'public', 'table': 't'}
)
payload.retrieve.assert_called_once_with(storage.minio_repository, metadata['metadata'])
storage.export_data_to_postgres.assert_called_once()
payload.retrieve.assert_awaited_once_with(storage.minio_repository, metadata['metadata'])
storage.export_data_to_postgres.assert_awaited_once()
assert result == {'success': True}
def test_cleanup_minio_objects_expired_minio_not_initialized(storage):
@mark.asyncio
async def test_cleanup_minio_objects_expired_minio_not_initialized(storage):
storage.minio_repository = None
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'test'
with raises(ValueError, match='Minio repository not initialized'):
storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
@mark.asyncio
@patch('laborious.activities.storage.now')
def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage):
async def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = MagicMock(
storage.minio_repository.list_objects = AsyncMock(
return_value=['some/random/key-without-timestamp.parquet']
)
storage.minio_repository.delete_file = MagicMock()
storage.send_notification = MagicMock()
storage.minio_repository.delete_file = AsyncMock()
storage.send_notification_async = AsyncMock()
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'test'
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 0
assert result['failed_count'] == 0
storage.minio_repository.delete_file.assert_not_called()
@mark.asyncio
@patch('laborious.activities.storage.now')
def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
old_key = 'training_datasets/m/m-initial-2024-12-01_00-00-00.parquet'
storage.minio_repository.list_objects = MagicMock(return_value=[old_key])
storage.minio_repository.delete_file = MagicMock(side_effect=Exception('delete error'))
storage.send_notification = MagicMock()
storage.minio_repository.list_objects = AsyncMock(return_value=[old_key])
storage.minio_repository.delete_file = AsyncMock(side_effect=Exception('delete error'))
storage.send_notification_async = AsyncMock()
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 0
assert result['failed_count'] == 1
@@ -300,20 +298,21 @@ def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
assert result['failed'][old_key]['message'] == 'delete error'
@mark.asyncio
@patch('laborious.activities.storage.now')
def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage):
async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = MagicMock(side_effect=Exception('list error'))
storage.send_notification = MagicMock()
storage.minio_repository.list_objects = AsyncMock(side_effect=Exception('list error'))
storage.send_notification_async = AsyncMock()
storage.error = MagicMock()
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 0
assert result['failed_count'] == 0
storage.send_notification.assert_called_once_with(
storage.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
message='Error cleaning up MinIO objects: list error',

View File

@@ -1,7 +1,8 @@
from datetime import datetime
from io import BytesIO
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from pandas import DataFrame
from laborious.utils.models.minio_dataframe_payload import (
@@ -53,15 +54,17 @@ def test_has_data_true_when_object_key_set():
assert payload.has_data() is True
def test_retrieve_inline_dict_as_dataframe():
@pytest.mark.asyncio
async def test_retrieve_inline_dict_as_dataframe():
payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]})
minio = MagicMock()
out = payload.retrieve(minio, {'metadata': {}})
minio = AsyncMock()
out = await payload.retrieve(minio, {'metadata': {}})
assert list(out.columns) == ['a']
minio.download_file.assert_not_called()
def test_retrieve_downloads_parquet_when_offloaded():
@pytest.mark.asyncio
async def test_retrieve_downloads_parquet_when_offloaded():
source = DataFrame({'a': [1, 2]})
buf = BytesIO()
source.to_parquet(buf, engine='pyarrow', index=True)
@@ -73,12 +76,12 @@ def test_retrieve_downloads_parquet_when_offloaded():
object_key='training_datasets/m/f.parquet',
object_prefix='training_datasets/m',
)
minio = MagicMock()
minio.download_file = MagicMock(return_value=file_bytes)
minio = AsyncMock()
minio.download_file = AsyncMock(return_value=file_bytes)
out = payload.retrieve(minio, {'metadata': {}})
out = await payload.retrieve(minio, {'metadata': {}})
minio.download_file.assert_called_once_with(
minio.download_file.assert_awaited_once_with(
object_name='training_datasets/m/f.parquet',
metadata={'metadata': {}},
)
@@ -110,19 +113,21 @@ def test_parse_object_timestamp_bad_datetime():
assert MinioDataFramePayload.parse_object_timestamp(key) is None
def test_retrieve_empty_when_no_data():
@pytest.mark.asyncio
async def test_retrieve_empty_when_no_data():
payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key=None)
minio = MagicMock()
out = payload.retrieve(minio, {})
minio = AsyncMock()
out = await payload.retrieve(minio, {})
assert out.empty
minio.download_file.assert_not_called()
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
def test_from_dataframe_none(mock_now):
async def test_from_dataframe_none(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = MagicMock()
result = MinioDataFramePayload.from_dataframe(
minio = AsyncMock()
result = await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=minio,
model_name='m',
@@ -134,14 +139,15 @@ def test_from_dataframe_none(mock_now):
assert result.object_key is None
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
def test_from_dataframe_empty(mock_now):
async def test_from_dataframe_empty(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = MagicMock()
minio = AsyncMock()
mock_df = MagicMock()
mock_df.__bool__ = MagicMock(return_value=True)
mock_df.empty = True
result = MinioDataFramePayload.from_dataframe(
result = await MinioDataFramePayload.from_dataframe(
dataframe=mock_df,
minio_repo=minio,
model_name='m',
@@ -168,11 +174,12 @@ def _mock_dataframe(data_dict, timestamp_values=None):
return mock_df
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
def test_from_dataframe_inline():
minio = MagicMock()
async def test_from_dataframe_inline():
minio = AsyncMock()
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = MinioDataFramePayload.from_dataframe(
result = await MinioDataFramePayload.from_dataframe(
dataframe=df,
minio_repo=minio,
model_name='m',
@@ -183,30 +190,17 @@ def test_from_dataframe_inline():
assert result.last_timestamp == '2024-01-01'
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
def test_from_dataframe_inline_uses_provided_last_timestamp():
minio = MagicMock()
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = MinioDataFramePayload.from_dataframe(
dataframe=df,
minio_repo=minio,
model_name='m',
operation='initial',
last_timestamp='2024-01-02',
)
assert result.last_timestamp == '2024-01-02'
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0)
def test_from_dataframe_offloaded(mock_now):
async def test_from_dataframe_offloaded(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = MagicMock()
minio.upload_file = MagicMock(return_value={'minio_object_name': 'full/key.parquet'})
minio = AsyncMock()
minio.upload_file = AsyncMock(return_value={'minio_object_name': 'full/key.parquet'})
minio.bucket = 'test-bucket'
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = MinioDataFramePayload.from_dataframe(
result = await MinioDataFramePayload.from_dataframe(
dataframe=df,
minio_repo=minio,
model_name='m',
@@ -217,7 +211,7 @@ def test_from_dataframe_offloaded(mock_now):
assert result.object_key == 'full/key.parquet'
assert result.bucket == 'test-bucket'
assert result.uri == 's3://test-bucket/full/key.parquet'
minio.upload_file.assert_called_once()
minio.upload_file.assert_awaited_once()
def test_from_dict_inline():
@@ -270,9 +264,3 @@ def test_from_dict_passthrough_existing_instance():
original = MinioDataFramePayload(last_timestamp='2024-01-01', data={'a': 1}, bucket='b')
result = MinioDataFramePayload.from_dict(original)
assert result is original
def test_debug_with_logger_calls_custom_debug():
logger = MagicMock()
MinioDataFramePayload._debug(logger, 'msg', {'a': 1})
logger.custom_debug.assert_called_once_with('msg', {'a': 1})

File diff suppressed because it is too large Load Diff

View File

@@ -1,10 +1,10 @@
import concurrent.futures
import asyncio
import json
from datetime import datetime
from unittest.mock import ANY, MagicMock, Mock, patch
from unittest.mock import ANY, AsyncMock, MagicMock, Mock, patch
import pytest
from asyncua.crypto import security_policies
from asyncua.crypto.security_policies import SecurityPolicyBasic256
from asyncua.ua.uaerrors import BadNodeIdUnknown, BadSessionIdInvalid
from sientia_do.notifications.models import NotificationLevel
@@ -35,12 +35,12 @@ def opc_repository(mock_logger):
cert_path='/path/to/cert.pem',
private_key_path='/path/to/key.pem',
server_cert_path='/path/to/server_cert.pem',
metrics_controller=MagicMock(),
metrics_controller=AsyncMock(),
)
repository.disconnection_interval = 0.1
repository.send_notification = MagicMock()
repository.send_notification = MagicMock()
repository.emit_metric_sync = MagicMock()
repository.send_notification_async = AsyncMock()
repository.emit_metric = AsyncMock()
repository.info = MagicMock()
repository.error = MagicMock()
repository.warning = MagicMock()
@@ -52,13 +52,7 @@ def opc_repository(mock_logger):
@pytest.fixture
def mock_client():
with patch('laborious.utils.repository.opc_repository.Client') as mock:
client_instance = MagicMock()
aio = MagicMock()
client_instance.aio_obj = aio
aio.uaclient = MagicMock()
aio.uaclient.protocol = MagicMock(state='closed')
aio.session_timeout = 600_000
aio.secure_channel_timeout = 600_000
client_instance = AsyncMock()
mock.return_value = client_instance
yield client_instance
@@ -86,56 +80,60 @@ def test_init(opc_repository):
assert opc_repository.last_reconnection_time is None
def test_set_security(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_set_security(opc_repository, mock_client):
opc_repository.client = mock_client
opc_repository.set_security()
await opc_repository.set_security()
mock_client.application_uri = 'urn:test:server'
mock_client.set_security.assert_called_once_with(
security_policies.SecurityPolicyBasic256,
'/path/to/cert.pem',
'/path/to/key.pem',
None,
'/path/to/server_cert.pem',
SecurityPolicyBasic256,
certificate='/path/to/cert.pem',
private_key='/path/to/key.pem',
server_certificate='/path/to/server_cert.pem',
)
assert mock_client.aio_obj.secure_channel_timeout == 600_000
assert mock_client.aio_obj.session_timeout == 600_000
assert mock_client.secure_channel_timeout == 600_000
assert mock_client.session_timeout == 600_000
def test_set_security_missing_certificates(opc_repository):
@pytest.mark.asyncio
async def test_set_security_missing_certificates(opc_repository):
opc_repository.cert_path = None
opc_repository.private_key_path = None
try:
opc_repository.set_security()
await opc_repository.set_security()
except ValueError as e:
assert str(e) == 'Certificate and private key paths must be provided for secure connection.'
def test_set_security_missing_client(opc_repository):
@pytest.mark.asyncio
async def test_set_security_missing_client(opc_repository):
opc_repository.client = None
try:
opc_repository.set_security()
await opc_repository.set_security()
except ValueError as e:
assert str(e) == 'Client must be initialized before setting security'
def test_connect_with_security(opc_repository, mock_client):
opc_repository._create_client = MagicMock()
opc_repository._open_session = MagicMock(return_value=(True, {}))
result = opc_repository.connect()
@pytest.mark.asyncio
async def test_connect_with_security(opc_repository, mock_client):
opc_repository._create_client = AsyncMock()
opc_repository._open_session = AsyncMock(return_value=(True, {}))
result = await opc_repository.connect()
opc_repository._create_client.assert_called_once()
opc_repository._open_session.assert_called_once()
assert result == (True, {})
def test_connect_without_security(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_connect_without_security(opc_repository, mock_client):
opc_repository.cert_path = None
opc_repository._create_client = MagicMock()
opc_repository._open_session = MagicMock(return_value=(True, {}))
opc_repository.set_security = MagicMock()
result = opc_repository.connect()
opc_repository._create_client = AsyncMock()
opc_repository._open_session = AsyncMock(return_value=(True, {}))
opc_repository.set_security = AsyncMock()
result = await opc_repository.connect()
opc_repository._create_client.assert_called_once()
opc_repository._open_session.assert_called_once()
@@ -143,68 +141,70 @@ def test_connect_without_security(opc_repository, mock_client):
assert result == (True, {})
def test_connect_raises_when_session_already_open(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_connect_raises_when_session_already_open(opc_repository, mock_client):
opc_repository.client = mock_client
proto = MagicMock()
proto.state = 'open'
mock_client.aio_obj.uaclient.protocol = proto
mock_client.uaclient = MagicMock(protocol=proto)
with pytest.raises(OpcSessionAlreadyConnectedError, match='disconnect'):
opc_repository.connect()
await opc_repository.connect()
def test_create_client_raises_when_client_exists(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_create_client_raises_when_client_exists(opc_repository, mock_client):
opc_repository.client = mock_client
with pytest.raises(OpcClientAlreadyExistsError, match='already exists'):
opc_repository._create_client()
await opc_repository._create_client()
def test_open_session_success(opc_repository):
@pytest.mark.asyncio
async def test_open_session_success(opc_repository):
closed_proto = MagicMock()
closed_proto.state = 'closed'
opc_repository.client = MagicMock()
aio = MagicMock()
opc_repository.client.aio_obj = aio
aio.uaclient = MagicMock(protocol=closed_proto)
aio.session_timeout = 600_000
aio.secure_channel_timeout = 600_000
opc_repository.client = AsyncMock()
opc_repository.client.uaclient = MagicMock(protocol=closed_proto)
opc_repository.client.session_timeout = 600_000
opc_repository.client.secure_channel_timeout = 600_000
open_proto = MagicMock()
open_proto.state = 'open'
open_proto.authentication_token = 'tok'
def connect_side_effect():
aio.uaclient.protocol = open_proto
async def connect_side_effect():
opc_repository.client.uaclient.protocol = open_proto
opc_repository.client.connect = MagicMock(side_effect=connect_side_effect)
opc_repository.client.connect = AsyncMock(side_effect=connect_side_effect)
result = opc_repository._open_session()
result = await opc_repository._open_session()
opc_repository.client.connect.assert_called_once()
assert opc_repository.last_reconnection_time is None
assert result == (True, {})
assert opc_repository._session_ready.is_set()
def test_open_session_raises_when_already_connected(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_open_session_raises_when_already_connected(opc_repository, mock_client):
opc_repository.client = mock_client
proto = MagicMock()
proto.state = 'open'
mock_client.aio_obj.uaclient.protocol = proto
mock_client.uaclient = MagicMock(protocol=proto)
with pytest.raises(OpcSessionAlreadyConnectedError, match='disconnect'):
opc_repository._open_session()
await opc_repository._open_session()
def test_open_session_fail(opc_repository):
opc_repository._disconnect_locked = MagicMock()
@pytest.mark.asyncio
async def test_open_session_fail(opc_repository):
opc_repository._disconnect_locked = AsyncMock()
opc_repository.client = MagicMock()
aio = MagicMock()
opc_repository.client.aio_obj = aio
aio.uaclient = MagicMock(protocol=MagicMock(state='closed'))
opc_repository.client.connect = MagicMock(side_effect=Exception('Test error'))
opc_repository.client.uaclient = MagicMock(protocol=MagicMock(state='closed'))
opc_repository.client.connect = AsyncMock(side_effect=Exception('Test error'))
is_connected, error_data = opc_repository._open_session()
is_connected, error_data = await opc_repository._open_session()
opc_repository._disconnect_locked.assert_called_once()
opc_repository.client.connect.assert_called_once()
@@ -216,26 +216,29 @@ def test_open_session_fail(opc_repository):
assert error_data['attachment_content'] is not None
def test_open_session_raises_when_no_client(opc_repository):
@pytest.mark.asyncio
async def test_open_session_raises_when_no_client(opc_repository):
opc_repository.client = None
with pytest.raises(OpcClientNotInitializedError, match='not initialized'):
opc_repository._open_session()
await opc_repository._open_session()
def test_disconnection_fallback_success(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_disconnection_fallback_success(opc_repository, mock_client):
opc_repository.client = mock_client
mock_client.disconnect.return_value = True
result = opc_repository._disconnection_fallback()
result = await opc_repository._disconnection_fallback()
mock_client.disconnect.assert_called_once()
assert result == []
def test_disconnection_fallback_fail(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_disconnection_fallback_fail(opc_repository, mock_client):
opc_repository.client = mock_client
mock_client.disconnect.side_effect = Exception('Test error')
result = opc_repository._disconnection_fallback()
result = await opc_repository._disconnection_fallback()
assert result == [
{'attempt': 1, 'error': 'Test error', 'traceback': ANY},
{'attempt': 2, 'error': 'Test error', 'traceback': ANY},
@@ -246,30 +249,33 @@ def test_disconnection_fallback_fail(opc_repository, mock_client):
assert mock_client.disconnect.call_count == 5
def test_disconnect(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_disconnect(opc_repository, mock_client):
opc_repository.client = mock_client
opc_repository._disconnection_fallback = MagicMock(return_value=[])
opc_repository.disconnect()
opc_repository._disconnection_fallback = AsyncMock(return_value=[])
await opc_repository.disconnect()
opc_repository._disconnection_fallback.assert_called_once()
assert opc_repository.client is None
assert opc_repository._allow_reconnect is False
def test_disconnect_no_client(opc_repository):
@pytest.mark.asyncio
async def test_disconnect_no_client(opc_repository):
opc_repository.client = None
assert opc_repository.disconnect() is None
assert await opc_repository.disconnect() is None
def test_disconnect_error(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_disconnect_error(opc_repository, mock_client):
opc_repository.client = mock_client
opc_repository._disconnection_fallback = MagicMock(
opc_repository._disconnection_fallback = AsyncMock(
return_value=[{'attempt': 1, 'error': 'Test error', 'traceback': 'text'}]
)
opc_repository.disconnect()
await opc_repository.disconnect()
opc_repository._disconnection_fallback.assert_called_once()
opc_repository.send_notification.assert_called_once_with(
opc_repository.send_notification_async.assert_called_once_with(
metadata=opc_repository.metadata,
notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}',
message='Failed to disconnect from OPC server in 5 attempts.',
@@ -282,56 +288,58 @@ def test_disconnect_error(opc_repository, mock_client):
assert opc_repository.client is None
def test_validate_connection_none_client(opc_repository):
@pytest.mark.asyncio
async def test_validate_connection_none_client(opc_repository):
opc_repository.client = None
response = opc_repository.validate_connection()
response = await opc_repository.validate_connection()
assert response == (False, opc_repository._not_connected_error())
def test_validate_connection_session_not_open(opc_repository):
@pytest.mark.asyncio
async def test_validate_connection_session_not_open(opc_repository):
opc_repository.client = MagicMock()
opc_repository.client.aio_obj.uaclient.protocol = None
opc_repository.client.uaclient.protocol = None
response = opc_repository.validate_connection()
response = await opc_repository.validate_connection()
assert response == (False, opc_repository._not_connected_error())
opc_repository.error.assert_called_once()
def test_validate_connection_success(opc_repository):
@pytest.mark.asyncio
async def test_validate_connection_success(opc_repository):
opc_repository.client = MagicMock()
proto = MagicMock()
proto.state = 'open'
opc_repository.client.aio_obj.uaclient.protocol = proto
opc_repository.client.uaclient.protocol = MagicMock()
opc_repository.client.uaclient.protocol.state = 'open'
output = opc_repository.validate_connection()
output = await opc_repository.validate_connection()
assert output == (True, {})
def test_write_data_validate_connection_do_nothing(opc_repository):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
opc_repository.client = MagicMock(get_node=MagicMock())
mock_node = MagicMock()
@pytest.mark.asyncio
async def test_write_data_validate_connection_do_nothing(opc_repository):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = AsyncMock(get_node=MagicMock())
mock_node = AsyncMock()
opc_repository.client.get_node.return_value = mock_node
result = opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
result = await opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
opc_repository.validate_connection.assert_called_once()
opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode')
assert result == (True, {'response_time': ANY})
def test_write_data_validate_connection_failed(opc_repository):
opc_repository.validate_connection = MagicMock(return_value=(False, {}))
@pytest.mark.asyncio
async def test_write_data_validate_connection_failed(opc_repository):
opc_repository.client = MagicMock()
opc_repository._start_reconnect = MagicMock()
opc_repository.client.uaclient.protocol = MagicMock(state='closed')
opc_repository._start_reconnect = AsyncMock()
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
opc_repository.validate_connection.assert_called_once()
opc_repository.client.get_node.assert_not_called()
opc_repository._start_reconnect.assert_called_once()
assert opc_repository._start_reconnect.call_args.args[0] == 'ProtocolClosed'
assert is_success is False
@@ -339,12 +347,13 @@ def test_write_data_validate_connection_failed(opc_repository):
assert error_data['opc_status'] == 'ProtocolClosed'
def test_write_data_get_node_failed(opc_repository):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
opc_repository.client = MagicMock()
@pytest.mark.asyncio
async def test_write_data_get_node_failed(opc_repository):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = AsyncMock()
opc_repository.client.get_node = MagicMock(side_effect=Exception('Test error'))
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
@@ -361,13 +370,14 @@ def test_write_data_get_node_failed(opc_repository):
assert error_data['attachment_content'] is not None
def test_write_data_invalid_data_type(opc_repository, mock_client):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
@pytest.mark.asyncio
async def test_write_data_invalid_data_type(opc_repository, mock_client):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = mock_client
mock_node = MagicMock()
mock_node = AsyncMock()
mock_client.get_node = MagicMock(return_value=mock_node)
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'invalid_type', metadata['metadata']
)
@@ -385,27 +395,29 @@ def test_write_data_invalid_data_type(opc_repository, mock_client):
assert error_data.get('attachment_content') is None
def test_write_data(opc_repository, mock_client):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
@pytest.mark.asyncio
async def test_write_data(opc_repository, mock_client):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = mock_client
mock_node = MagicMock()
mock_node = AsyncMock()
mock_client.get_node = MagicMock(return_value=mock_node)
result = opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
result = await opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
mock_node.write_value.assert_called_once()
assert result == (True, {'response_time': ANY})
def test_write_data_write_value_failed(opc_repository, mock_client):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
@pytest.mark.asyncio
async def test_write_data_write_value_failed(opc_repository, mock_client):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = mock_client
mock_node = MagicMock()
mock_node = AsyncMock()
mock_client.get_node = MagicMock(return_value=mock_node)
mock_node.write_value.side_effect = Exception('Test error')
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
@@ -429,15 +441,16 @@ def test_is_reconnectable_opcua_bad():
assert is_reconnectable_opcua_bad(Exception('other')) is False
def test_write_data_bad_session_id_invalid_schedules_reconnect(opc_repository, mock_client):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
@pytest.mark.asyncio
async def test_write_data_bad_session_id_invalid_schedules_reconnect(opc_repository, mock_client):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = mock_client
opc_repository._start_reconnect = MagicMock()
mock_node = MagicMock()
opc_repository._start_reconnect = AsyncMock()
mock_node = AsyncMock()
mock_client.get_node = MagicMock(return_value=mock_node)
mock_node.write_value.side_effect = BadSessionIdInvalid()
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
@@ -448,36 +461,43 @@ def test_write_data_bad_session_id_invalid_schedules_reconnect(opc_repository, m
assert error_data['opc_status'] == 'BadSessionIdInvalid'
def test_write_data_reconnect_in_progress_immediate(opc_repository):
@pytest.mark.asyncio
async def test_write_data_reconnect_in_progress_immediate(opc_repository):
opc_repository._session_ready.clear()
opc_repository._reconnect_thread = MagicMock()
opc_repository._reconnect_thread.is_alive.return_value = True
opc_repository.validate_connection = MagicMock()
opc_repository._reconnect_task = asyncio.create_task(asyncio.sleep(60))
opc_repository.validate_connection = AsyncMock()
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
opc_repository._reconnect_task.cancel()
with pytest.raises(asyncio.CancelledError):
await opc_repository._reconnect_task
opc_repository._reconnect_task = None
opc_repository.validate_connection.assert_not_called()
assert is_success is False
assert error_data['opc_error_kind'] == 'reconnect_in_progress'
def test_start_reconnect_skips_within_interval(opc_repository):
@pytest.mark.asyncio
async def test_start_reconnect_skips_within_interval(opc_repository):
opc_repository.last_reconnection_time = datetime.now()
opc_repository.reconnection_interval = 3600
opc_repository._start_reconnect('BadSessionIdInvalid', 'tok')
await opc_repository._start_reconnect('BadSessionIdInvalid', 'tok')
assert opc_repository._reconnect_thread is None
assert opc_repository._reconnect_task is None
def test_write_data_protocol_closed_schedules_reconnect(opc_repository):
@pytest.mark.asyncio
async def test_write_data_protocol_closed_schedules_reconnect(opc_repository):
opc_repository.client = MagicMock()
opc_repository.client.aio_obj.uaclient.protocol = MagicMock(state='closed')
opc_repository._start_reconnect = MagicMock()
opc_repository.client.uaclient.protocol = MagicMock(state='closed')
opc_repository._start_reconnect = AsyncMock()
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
@@ -488,95 +508,101 @@ def test_write_data_protocol_closed_schedules_reconnect(opc_repository):
assert error_data['opc_status'] == 'ProtocolClosed'
def test_write_data_protocol_closed_skips_reconnect_within_interval(opc_repository):
@pytest.mark.asyncio
async def test_write_data_protocol_closed_skips_reconnect_within_interval(opc_repository):
opc_repository.client = MagicMock()
opc_repository.client.aio_obj.uaclient.protocol = MagicMock(state='closed')
opc_repository.client.uaclient.protocol = MagicMock(state='closed')
opc_repository.last_reconnection_time = datetime.now()
opc_repository.reconnection_interval = 3600
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
assert opc_repository._reconnect_thread is None
assert opc_repository._reconnect_task is None
assert is_success is False
assert error_data['opc_error_kind'] == 'connection_lost'
def test_write_data_after_failed_reconnect_schedules_again(opc_repository):
@pytest.mark.asyncio
async def test_write_data_after_failed_reconnect_schedules_again(opc_repository):
opc_repository._session_ready.clear()
opc_repository.reconnection_interval = 0
opc_repository.last_reconnection_time = None
opc_repository._reconnect_locked = MagicMock(
opc_repository._reconnect_locked = AsyncMock(
return_value=(False, {'message': 'connect failed'})
)
opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
if opc_repository._reconnect_thread is not None:
opc_repository._reconnect_thread.join(timeout=2)
await opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
await asyncio.sleep(0.1)
assert opc_repository._reconnect_locked.call_count == 1
assert not opc_repository._reconnect_thread_in_progress()
assert not opc_repository._reconnect_task_in_progress()
opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
if opc_repository._reconnect_thread is not None:
opc_repository._reconnect_thread.join(timeout=2)
await opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
await asyncio.sleep(0.1)
assert opc_repository._reconnect_locked.call_count == 2
def test_write_data_after_disconnect_does_not_schedule_reconnect(opc_repository, mock_client):
@pytest.mark.asyncio
async def test_write_data_after_disconnect_does_not_schedule_reconnect(opc_repository, mock_client):
opc_repository.client = mock_client
proto = MagicMock()
proto.state = 'closed'
mock_client.aio_obj.uaclient.protocol = proto
opc_repository._disconnection_fallback = MagicMock(return_value=[])
opc_repository.disconnect()
mock_client.uaclient = MagicMock(protocol=proto)
opc_repository._disconnection_fallback = AsyncMock(return_value=[])
await opc_repository.disconnect()
is_success, error_data = opc_repository.write_data(
is_success, error_data = await opc_repository.write_data(
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
)
assert opc_repository._reconnect_thread is None
assert opc_repository._reconnect_task is None
assert is_success is False
assert error_data['opc_error_kind'] == 'connection_lost'
def test_parallel_bad_writes_single_reconnect_task(opc_repository, mock_client):
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
@pytest.mark.asyncio
async def test_parallel_bad_writes_single_reconnect_task(opc_repository, mock_client):
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
opc_repository.client = mock_client
opc_repository.reconnection_interval = 0
opc_repository.last_reconnection_time = None
mock_node = MagicMock()
mock_node = AsyncMock()
mock_client.get_node = MagicMock(return_value=mock_node)
mock_node.write_value.side_effect = BadSessionIdInvalid()
opc_repository._start_reconnect = MagicMock()
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor:
futures = [
executor.submit(
opc_repository.write_data,
node,
value,
'float',
metadata['metadata'],
)
for node, value in (('ns=2;s=TestNode', 1.0), ('ns=2;s=TestNode2', 2.0))
]
results = [future.result() for future in futures]
connect_count = 0
assert 1 <= opc_repository._start_reconnect.call_count <= 2
assert mock_node.write_value.call_count == 2
async def slow_reconnect():
nonlocal connect_count
connect_count += 1
await asyncio.sleep(0.05)
opc_repository._session_ready.set()
return True, {}
opc_repository._reconnect_locked = slow_reconnect
results = await asyncio.gather(
opc_repository.write_data('ns=2;s=TestNode', 1.0, 'float', metadata['metadata']),
opc_repository.write_data('ns=2;s=TestNode2', 2.0, 'float', metadata['metadata']),
)
await asyncio.sleep(0.15)
assert connect_count <= 1
assert 1 <= mock_node.write_value.call_count <= 2
error_kinds = [r[1].get('opc_error_kind') for r in results]
assert error_kinds.count('session_bad') >= 1
assert all(k in ('session_bad', 'reconnect_in_progress') for k in error_kinds)
@pytest.mark.asyncio
@patch('laborious.utils.repository.opc_repository.datetime')
def test_reconnect_locked_sets_last_reconnection_time(mock_datetime, opc_repository):
async def test_reconnect_locked_sets_last_reconnection_time(mock_datetime, opc_repository):
mock_datetime.now = MagicMock(return_value=datetime(2025, 1, 1, 12, 0, 0))
opc_repository._disconnect_locked = MagicMock()
opc_repository._connect_locked = MagicMock(return_value=(True, {}))
opc_repository._disconnect_locked = AsyncMock()
opc_repository._connect_locked = AsyncMock(return_value=(True, {}))
result = opc_repository._reconnect_locked()
result = await opc_repository._reconnect_locked()
opc_repository._disconnect_locked.assert_called_once()
opc_repository._connect_locked.assert_called_once()

View File

@@ -4,13 +4,13 @@ from laborious.utils.connectors_config import (
build_minio_config,
build_mlflow_config,
build_opc_config,
build_plugin_store_config,
)
def test_build_mlflow_config_with_env_vars():
# Arrange
environ['MLFLOW_URL'] = 'http://test-host:8080'
environ['MLFLOW_HOST'] = 'http://test-host'
environ['MLFLOW_PORT'] = '8080'
environ['MLFLOW_USERNAME'] = 'test-user'
environ['MLFLOW_PASSWORD'] = 'test-pass'
@@ -18,25 +18,17 @@ def test_build_mlflow_config_with_env_vars():
config = build_mlflow_config()
# Assert
assert config['url'] == 'http://test-host:8080'
assert config['host'] == 'http://test-host'
assert config['port'] == 8080
assert config['username'] == 'test-user'
assert config['password'] == 'test-pass'
def test_build_mlflow_config_host_already_has_port():
environ['MLFLOW_URL'] = 'http://tracker.example.com:443'
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
environ.pop('MLFLOW_URL', None)
environ.pop('MLFLOW_HOST', None)
environ.pop('MLFLOW_PORT', None)
environ.pop('MLFLOW_USERNAME', None)
environ.pop('MLFLOW_PASSWORD', None)
@@ -44,30 +36,12 @@ def test_build_mlflow_config_with_defaults():
config = build_mlflow_config()
# Assert
assert config['url'] == 'http://localhost:5080'
assert config['host'] == 'http://localhost'
assert config['port'] == 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"}}'
@@ -122,8 +96,6 @@ def test_build_minio_config_with_env_vars():
environ['MINIO_SECRET_KEY'] = 'test-secret'
environ['MINIO_REGION_NAME'] = 'test-region'
environ['MINIO_DEFAULT_BUCKET'] = 'test-bucket'
# Isolate from IDE/CI env (e.g. VS Code may export MINIO_SECURE=true).
environ['MINIO_SECURE'] = 'false'
assert build_minio_config() == {
'endpoint_url': 'http://test-host',
'access_key': 'test-key',
@@ -140,8 +112,6 @@ def test_build_minio_config_with_defaults():
environ.pop('MINIO_SECRET_KEY', None)
environ.pop('MINIO_REGION_NAME', None)
environ.pop('MINIO_DEFAULT_BUCKET', None)
environ.pop('MINIO_SECURE', None)
environ.pop('MINIO_RETENTION_HOURS', None)
assert build_minio_config() == {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',

View File

@@ -1,12 +0,0 @@
from pandas import DataFrame
from laborious.utils.dataframe_debug import build_dataframe_debug_message
def test_build_dataframe_debug_message_skips_large_dataframe():
df = DataFrame({'a': [1, 2, 3]})
msg = build_dataframe_debug_message('payload', df, max_rows=1)
assert 'skipped because dataframe has 3 rows' in msg
assert '(max: 1)' in msg

View File

@@ -0,0 +1,14 @@
from sientia_do.temporal.worker.prepare_worker import build_queue_name
from laborious.workflows.drift import Drift
from laborious.workflows.minimal_retrain import MinimalRetrain
from laborious.workflows.predictions_batch import PredictionsBatch
from laborious.workflows.simple_metrics import SimpleMetrics
def test_runtime_scoped_queue_names():
runtime = 'prod-a'
assert build_queue_name(PredictionsBatch.__name__, runtime) == 'predictions_batch-prod-a-queue'
assert build_queue_name(MinimalRetrain.__name__, runtime) == 'minimal_retrain-prod-a-queue'
assert build_queue_name(Drift.__name__, runtime) == 'drift-prod-a-queue'
assert build_queue_name(SimpleMetrics.__name__, runtime) == 'simple_metrics-prod-a-queue'

View File

@@ -1,243 +0,0 @@
from unittest.mock import AsyncMock, MagicMock, patch
from pytest import mark, raises
from laborious.worker import worker
def _build_fake_activities():
inst = MagicMock()
inst.init_opc = MagicMock()
inst.shutdown = MagicMock()
inst.load_query_with_minio_offload = MagicMock()
inst.retrain_model = MagicMock()
inst.update_production_model = MagicMock()
inst.format_retrain_report = MagicMock()
inst.export_data_to_postgres = MagicMock()
inst.load_custom_query = MagicMock()
inst.calculate_simple_metrics = MagicMock()
inst.get_reference_data = MagicMock()
inst.calculate_drift = MagicMock()
inst.request_predict = MagicMock()
inst.request_transform = MagicMock()
inst.input_gate = MagicMock()
inst.mlflow_response_gate = MagicMock()
inst.mlflow_content_gate = MagicMock()
inst.format_transformed_data = MagicMock()
inst.format_prediction = MagicMock()
inst.format_default_prediction = MagicMock()
inst.write_opc_data = MagicMock()
inst.cleanup_minio_objects_expired = MagicMock()
inst.repeat_last_prediction = MagicMock()
inst.write_metrics = MagicMock()
inst.write_pi_web_api_data = MagicMock()
return inst
def _build_fake_worker(async_result=None, async_error: Exception | None = None):
w = MagicMock()
async def _run():
if async_error is not None:
raise async_error
return async_result
w.run = MagicMock(side_effect=_run)
return w
@patch('laborious.worker.worker.start_http_server')
def test_start_prometheus_server_success(mock_start_http):
with patch.object(worker.metrics.APP_UP, 'labels') as labels:
gauge = MagicMock()
labels.return_value = gauge
with patch('laborious.worker.worker.os.getenv', return_value='9090'):
worker.start_prometheus_server()
mock_start_http.assert_called_once_with(9090)
gauge.set.assert_called_once_with(1)
@patch('laborious.worker.worker.start_http_server', side_effect=RuntimeError('nope'))
def test_start_prometheus_server_error_exits(_mock_start_http):
with patch('laborious.worker.worker.os._exit', side_effect=SystemExit(1)) as m_exit:
with raises(SystemExit):
worker.start_prometheus_server()
m_exit.assert_called_once_with(1)
@mark.asyncio
async def test_main_missing_runtime_exits_fast(monkeypatch):
monkeypatch.setenv('RUNTIME', '')
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger') as m_logger,
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
assert m_logger.return_value.custom_critical.called
@mark.asyncio
async def test_main_plugin_install_failure(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(side_effect=RuntimeError('install failed'))
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger'),
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
@mark.asyncio
async def test_main_success_exit_zero(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(return_value=None)
fake_workers = [_build_fake_worker() for _ in range(4)]
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger'),
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler') as m_notif_cls,
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch('laborious.worker.worker.build_postgres_config', return_value={}),
patch('laborious.worker.worker.build_minio_config', return_value={}),
patch('laborious.worker.worker.build_opc_config', return_value={}),
patch('laborious.worker.worker.build_api_config', return_value={}),
patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()),
patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()),
patch('laborious.worker.worker.Runtime', return_value=MagicMock()),
patch(
'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock())
),
patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers) as m_prepare,
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(0)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
notif = m_notif_cls.return_value
notif.shutdown.assert_called_once()
fake_activities.shutdown.assert_called_once()
assert m_prepare.call_count == 4
prepare_calls = m_prepare.call_args_list
assert prepare_calls[0].kwargs['runtime'] == 'single'
assert prepare_calls[3].kwargs['runtime'] == 'single'
@mark.asyncio
async def test_main_worker_gather_error_exits_one(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(return_value=None)
fake_workers = [
_build_fake_worker(async_error=RuntimeError('boom')),
_build_fake_worker(),
_build_fake_worker(),
_build_fake_worker(),
]
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger') as m_logger,
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch('laborious.worker.worker.build_postgres_config', return_value={}),
patch('laborious.worker.worker.build_minio_config', return_value={}),
patch('laborious.worker.worker.build_opc_config', return_value={}),
patch('laborious.worker.worker.build_api_config', return_value={}),
patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()),
patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()),
patch('laborious.worker.worker.Runtime', return_value=MagicMock()),
patch(
'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock())
),
patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
assert m_logger.return_value.custom_error.called

View File

@@ -1,6 +1,6 @@
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pytest import fixture, mark, raises
from pytest import fixture, mark
from laborious.activities.activities import Activities
from laborious.workflows.sub_workflows.prediction_process import PredictionProcess
@@ -840,27 +840,3 @@ async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process):
retry_policy=ANY,
start_to_close_timeout=ANY,
)
@mark.asyncio
@patch('laborious.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
async def test_run_always_cleans_up_on_pipeline_exception(workflow_mock, prediction_process):
input_data = {
'metadata': metadata,
'data': {'last_timestamp': '2024-01-01'},
'model_id': 1,
'model_name': 'm',
'model_config': {},
'save_transform': False,
}
prediction_process._run_prediction_pipeline = AsyncMock(side_effect=RuntimeError('boom'))
with raises(RuntimeError):
await prediction_process.run(input_data)
workflow_mock.execute_activity_method.assert_called_once_with(
Activities.cleanup_minio_objects_expired,
{**metadata, 'data': input_data['data']},
retry_policy=ANY,
start_to_close_timeout=ANY,
)

View File

@@ -34,7 +34,8 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
'table_name': 'test_table',
'model_config': {
'target': 'test_target',
'retention_minutes': 0,
'transform_flavor': 'test_transform_flavor',
'predict_flavor': 'test_predict_flavor',
},
}
@@ -164,7 +165,8 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
'table_name': 'test_table',
'model_config': {
'target': 'test_target',
'retention_minutes': 0,
'transform_flavor': 'test_transform_flavor',
'predict_flavor': 'test_predict_flavor',
},
}
@@ -225,7 +227,8 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
'table_name': 'test_table',
'model_config': {
'target': 'test_target',
'retention_minutes': 0,
'transform_flavor': 'test_transform_flavor',
'predict_flavor': 'test_predict_flavor',
},
}

View File

@@ -1,319 +0,0 @@
#
# Default values for sientia-laborious-worker using the sientia-module chart (0.6.x).
# This is a YAML-formatted file.
# Declare variables to be passed into your templates.
#
projectName: &projectName "sientia-laborious-worker"
# -----------------------------------------------------------------------------
# Global configuration shared by all runtimes
# -----------------------------------------------------------------------------
global:
namespace: sientia
image:
repository: aignosi.azurecr.io/sientia-module
pullPolicy: Always
tag: "1.2.0"
commonLabels: {}
resources:
# Resource limits and requests are important for ResourceBasedTuner to work correctly.
# The tuner monitors system CPU and memory usage, so proper resource limits must be set.
limits:
cpu: 2000m
memory: 20Gi
requests:
cpu: 1000m
memory: 2Gi
livenessProbe:
exec:
command:
- sh
- -c
- |
curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1'
initialDelaySeconds: 1260
periodSeconds: 15
timeoutSeconds: 5
failureThreshold: 3
readinessProbe:
exec:
command:
- sh
- -c
- |
curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1'
initialDelaySeconds: 1200
periodSeconds: 10
timeoutSeconds: 3
failureThreshold: 2
autoscaling:
enabled: false
minReplicas: 1
maxReplicas: 100
targetCPUUtilizationPercentage: 80
# targetMemoryUtilizationPercentage: 80
# Environment variables shared by all runtimes.
env:
# Entrypoint variables
- name: GITHUB_REPO_URL
value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git"
- name: GITHUB_BRANCH
value: "release/SIENTIAPDE-1646"
- name: PYTHON_APP
value: "laborious.worker.worker"
- name: PYPI_SERVER
value: "http://library-distribution-server.library.svc.cluster.local:5000"
# Application variables
- name: POSTGRES_HOST
value: "paradedb-rw.paradedb.svc.cluster.local"
- name: POSTGRES_PORT
value: "5432"
- name: POSTGRES_USER
value: "postgres"
- name: POSTGRES_PASSWORD
value: "nFqc81y6kwmr2zuAIx43DhiOosFCVPpeEfTtTWZflkNjB2j1KtEeIANkhFR9mAX3"
- name: POSTGRES_DBNAME
value: "sientia"
- name: POSTGRES_MIN_CONNECTIONS
value: "20"
# max_connections = number_of_workers * max_concurrent_activities * safety_factor
# Example: 4 workers * 50 activities * 0.5 = 100 connections
- name: POSTGRES_MAX_CONNECTIONS
value: "100"
- name: MLFLOW_URL
value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local:80"
- name: MLFLOW_USERNAME
value: "aignosi"
- name: MLFLOW_PASSWORD
value: "1L0FP50j3ncp123"
# Plugin store (model-library-store Git + runtime packages).
- name: STORE_BASE_URL
value: "http://gitea-http.gitea.svc.cluster.local:3000"
- name: STORE_OWNER
value: "aignosi"
- name: STORE_REPO
value: "suse-model-store"
- name: STORE_USERNAME
valueFrom:
secretKeyRef:
name: sientia-plugin-store-credentials
key: username
- name: STORE_PASSWORD
valueFrom:
secretKeyRef:
name: sientia-plugin-store-credentials
key: password
- name: STORE_CACHE_TTL_SECONDS
value: "3600"
- name: OPC_ID
value: "1"
- name: OPC_SERVER_NAME
value: "default_server"
- name: OPC_URL
value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
- name: LOG_LEVEL
value: "DEBUG"
- name: HTTP_METRICS_PORT
value: "9090"
- name: HTTP_SDK_METRICS_PORT
value: "9091"
- name: PROJECT_NAME
value: "sientia-laborious"
- name: TEMPORAL_HOST
value: "temporal-frontend.temporal.svc.cluster.local:7233"
- name: TEMPORAL_NAMESPACE
value: "laborious"
- name: MONGODB_USERNAME
value: "root"
- name: MONGODB_PASSWORD
value: "wKZDbMNU1c"
- name: MONGODB_URL
value: "my-release-mongodb.mongodb.svc.cluster.local:27017"
- name: MONGODB_DATABASE
value: "sientia"
- name: MONGODB_TTL_INDEX_HOURS
value: "1"
- name: MINIO_ENDPOINT_URL
value: "minio.minio.svc.cluster.local:9000"
- name: MINIO_ACCESS_KEY
value: "admin"
- name: MINIO_SECRET_KEY
value: "LiArt4eNmJ"
- name: MINIO_DEFAULT_BUCKET
value: "sientia"
- name: MINIO_RETENTION_HOURS
value: "24"
- name: SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES
value: "0.5"
# Temporal worker tuning for PredictionsBatch.
# IMPORTANT: prefix must be PREDICTIONSBATCH_ (from class name PredictionsBatch).
- name: PREDICTIONSBATCH_MAX_CONCURRENT_WORKFLOW_TASKS
value: "20"
- name: PREDICTIONSBATCH_MAX_CONCURRENT_ACTIVITIES
value: "60"
- name: PREDICTIONSBATCH_ACTIVITY_EXECUTOR_MAX_WORKERS
value: "10"
- name: PREDICTIONSBATCH_MAX_CONCURRENT_LOCAL_ACTIVITIES
value: "20"
- name: PREDICTIONSBATCH_MAX_CACHED_WORKFLOWS
value: "200"
- name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM
value: "3"
- name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_INITIAL
value: "5"
- name: PREDICTIONSBATCH_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM
value: "15"
- name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM
value: "3"
- name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_INITIAL
value: "10"
- name: PREDICTIONSBATCH_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM
value: "30"
- name: MINIMALRETRAIN_MAX_CONCURRENT_ACTIVITIES
value: "5"
- name: MINIMALRETRAIN_ACTIVITY_EXECUTOR_MAX_WORKERS
value: "5"
- name: MINIMALRETRAIN_MAX_CONCURRENT_LOCAL_ACTIVITIES
value: "5"
- name: MINIMALRETRAIN_MAX_CACHED_WORKFLOWS
value: "5"
- name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MINIMUM
value: "5"
- name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_INITIAL
value: "5"
- name: MINIMALRETRAIN_WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM
value: "5"
- name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MINIMUM
value: "5"
- name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_INITIAL
value: "5"
- name: MINIMALRETRAIN_ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM
value: "5"
- name: PI_WEB_API_BASE_URL
value: "https://pivision.votorantimcimentos.com/piwebapi"
- name: PI_WEB_API_AUTH_TYPE
value: "basic"
- name: PI_WEB_API_AUTH_TOKEN
valueFrom:
secretKeyRef:
name: pi-web-api-auth-token
key: token
# Thread-pool size for non-runtime workers that also use prepare_worker.
- name: SIMPLEMETRICS_ACTIVITY_EXECUTOR_MAX_WORKERS
value: "20"
- name: DRIFT_ACTIVITY_EXECUTOR_MAX_WORKERS
value: "20"
# -----------------------------------------------------------------------------
# Runtimes configuration
# -----------------------------------------------------------------------------
# IMPORTANT:
# - The runtime name is used by the worker bootstrap to resolve plugins and task queues.
# - Keep runtime names in sync with the plugin-store runtime names.
runtimes:
basic:
replicas: 1
legacy:
replicas: 1
env:
- name: "GITHUB_BRANCH"
value: "main"
- name: "MLFLOW_HOST"
value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local"
- name: "MLFLOW_PORT"
value: "80"
# -----------------------------------------------------------------------------
# Chart-level configuration (applies to all runtimes)
# -----------------------------------------------------------------------------
imagePullSecrets:
- name: docker-hub-secret
nameOverride: *projectName
fullnameOverride: *projectName
serviceAccount:
create: true
automount: true
annotations: {}
name: *projectName
podAnnotations: {}
podLabels: {}
podSecurityContext: {}
securityContext: {}
volumes: []
volumeMounts: []
nodeSelector: {}
tolerations: []
affinity: {}
services:
sdk-metrics:
enabled: true
type: ClusterIP
port: 9091
targetPort: 9091
name: sdk-metrics
metrics:
enabled: true
type: ClusterIP
port: 9090
targetPort: 9090
name: metrics
# Configuração do ServiceMonitor para o Prometheus Operator
# ref: https://github.com/prometheus-operator/prometheus-operator
serviceMonitor:
enabled: true
endpoints:
- port: metrics
path: /metrics
interval: 30s
relabelings: []
- port: sdk-metrics
path: /metrics
interval: 30s
relabelings: []
additionalLabels:
release: kube-prometheus-stack
ssh:
enabled: true
secretName: git-ssh-key-sientia-laborious-worker
sshPath: /mnt/.ssh
knownHostsPath: /mnt/known_hosts
# kubectl create secret docker-registry docker-hub-secret --namespace sientia --docker-server=http://aignosi.azurecr.io --docker-username=aignosi --docker-password=<pwd>
#
# helm upgrade --install sientia-laborious-worker /home/grezewave/Documents/projects/sientia/sientia-core-applications/sientia-module -n sientia --create-namespace -f ./values.yaml
# kubectl create secret generic git-ssh-key-sientia-laborious-worker \
# --namespace sientia \
# --from-file=ssh-privatekey=git_key \
# --type=kubernetes.io/ssh-auth
# helm upgrade --install sientia-laborious-worker /home/grezewave/Documents/projects/sientia/sientia-core-applications/sientia-module -n sientia --create-namespace -f ./values.yaml