Compare commits
24 Commits
fix/SIENTI
...
695e4c07a6
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
695e4c07a6 | ||
|
|
6a8c41b328 | ||
|
|
773c980fc3 | ||
|
|
35efb67c89 | ||
|
|
acd925cf2a | ||
|
|
310ceea0d8 | ||
|
|
2d73ec9ec2 | ||
|
|
2142143ab9 | ||
|
|
c07457bfbe | ||
|
|
086b12492e | ||
|
|
4ea0754f0c | ||
|
|
ddb1618209 | ||
|
|
90f8bdda61 | ||
|
|
2ccda3e440 | ||
|
|
d856150e24 | ||
|
|
6569810756 | ||
|
|
7153f1da0d | ||
|
|
fcc8920a8b | ||
|
|
cd1be2430a | ||
|
|
7c1dae8ef6 | ||
|
|
2a6def4056 | ||
|
|
7d59fd7c8c | ||
|
|
e3636d4b88 | ||
|
|
638d5b70b4 |
16
.env.example
16
.env.example
@@ -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"
|
||||
|
||||
|
||||
18
.github/workflows/quality-gate.yml
vendored
18
.github/workflows/quality-gate.yml
vendored
@@ -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
4
.gitignore
vendored
@@ -54,5 +54,5 @@ catboost_info/
|
||||
mlruns/
|
||||
|
||||
relatorio*
|
||||
openspec/*
|
||||
.cursor/*
|
||||
openspec/
|
||||
.cursor/
|
||||
135
README.md
135
README.md
@@ -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
|
||||
|
||||
|
||||
@@ -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:** Jensen–Shannon NULLs for the drift happy-path scenario are fully traced (histogram out-of-range + `density=True`, vs NannyML’s 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 **Jensen–Shannon**) 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 Jensen–Shannon (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 project’s 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 workflow’s `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 workflow’s `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.
|
||||
@@ -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
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
# Specification: `sientia_model` — Jensen–Shannon 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 Jensen–Shannon 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 **Jensen–Shannon** 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 NannyML’s 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).
|
||||
842
e2e/conftest.py
842
e2e/conftest.py
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
);
|
||||
221
e2e/helpers.py
221
e2e/helpers.py
@@ -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)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"]
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
883
e2e/scenarios.md
883
e2e/scenarios.md
@@ -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
|
||||
|
||||
---
|
||||
@@ -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
|
||||
|
||||
@@ -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}'
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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'
|
||||
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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 wrapper’s `_model_impl` class does not extend `SientiaModel`.
|
||||
- All this MLflow-specific logic is local to `laborious_temporal` and partially duplicated in `sientia-dataops-model-manager`, which motivates the extraction of a shared MLflow repository in `sientia-dataops-library` (see `mlflow-shared-repository-migration-plan`).
|
||||
|
||||
**MLflow:** All MLflow ops → [[mlflow-shared-repository-migration-plan|shared repository]]. Laborious uses the interface; `SientiaModel` lifecycle is in sientia-model-library.
|
||||
|
||||
### Current vs Target — High-level Flow
|
||||
|
||||
```mermaid
|
||||
flowchart LR
|
||||
subgraph current [Current State — Laborious Temporal]
|
||||
direction TB
|
||||
TemporalWorker["Temporal Worker"]
|
||||
MlflowActivities["MLFlow Activities\nrequest_transform / request_predict / retrain_model"]
|
||||
MLFlowRepositoryNode["MLFlowRepository"]
|
||||
MLflowRegistry["MLflow Tracking + Registry"]
|
||||
RawModel["Loaded Model\n(sklearn / pyfunc / pytorch)"]
|
||||
|
||||
TemporalWorker -->|"start workflow\n(Temporal)"| MlflowActivities
|
||||
MlflowActivities -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode
|
||||
MLFlowRepositoryNode -->|"search_model_versions()\ncurrent_stage == 'Production'"| MLflowRegistry
|
||||
MLFlowRepositoryNode -->|"mlflow.*.load_model(model_uri)"| RawModel
|
||||
MLFlowRepositoryNode -->|"raw_model.predict(data)\nraw_model.fit(data)"| RawModel
|
||||
end
|
||||
|
||||
subgraph target [Target State — Laborious Temporal]
|
||||
direction TB
|
||||
TemporalWorker2["Temporal Worker"]
|
||||
MlflowActivities2["MLFlow Activities\n(same APIs)"]
|
||||
MLFlowRepositoryNode2["MLFlowRepository\n(wrapper-aware)"]
|
||||
MLflowRegistry2["MLflow Tracking + Registry\n(aliases enabled)"]
|
||||
SientiaWrapperNode["Wrapper Instance\n(extends SientiaModel)"]
|
||||
|
||||
TemporalWorker2 -->|"start workflow\n(Temporal)"| MlflowActivities2
|
||||
MlflowActivities2 -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode2
|
||||
MLFlowRepositoryNode2 -->|"get_model_version_by_alias('production')\n& models:/name@production"| MLflowRegistry2
|
||||
MLFlowRepositoryNode2 -->|"mlflow.pyfunc.load_model(...)"| SientiaWrapperNode
|
||||
MLFlowRepositoryNode2 -->|"wrapper.transform(...)\nwrapper.predict(...)\nwrapper.train()/retrain()"| SientiaWrapperNode
|
||||
SientiaWrapperNode -->|"store_model(...)\n(auto-register + update alias)"| MLflowRegistry2
|
||||
end
|
||||
```
|
||||
|
||||
### Current vs Target — Retrain Hot Path (Code Sketch)
|
||||
|
||||
Current retrain flow inside `MLFlowRepository.fit_models` / `retrain_model` (simplified):
|
||||
|
||||
```python
|
||||
data_model, _ = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type="transform",
|
||||
flavor=transform_flavor,
|
||||
load_wrapper=(transform_flavor == "pyfunc"),
|
||||
)
|
||||
|
||||
prediction_model, _ = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type="predict",
|
||||
flavor=predict_flavor,
|
||||
load_wrapper=(predict_flavor == "pyfunc"),
|
||||
)
|
||||
|
||||
treated_data_candidate = data_model.fit(data) # or data_model.predict(data)
|
||||
...
|
||||
prediction_model.fit(retrain_dataset) # direct fit on underlying model
|
||||
```
|
||||
|
||||
Target retrain flow when using wrappers that extend `SientiaModel`:
|
||||
|
||||
```python
|
||||
wrapper, _ = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
# model_type="predict", # wrapper owns both transformer + model stack
|
||||
# flavor="pyfunc", flavor will always be pyfunc, this parameter will be removed
|
||||
# load_wrapper=True, wrapper will always be loaded from _model_impl
|
||||
)
|
||||
|
||||
# First-time training or full retrain using public API
|
||||
wrapper.train(
|
||||
train_data=train_df, # features + target
|
||||
val_data=val_df,
|
||||
target=target_name,
|
||||
)
|
||||
|
||||
# Incremental retrain (when appropriate)
|
||||
wrapper.retrain(full_retrain_df)
|
||||
|
||||
transformed_df, trans_meta = wrapper.transform(raw_df)
|
||||
pred_df, pred_meta = wrapper.predict({}, transformed_df, params={})
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Requirements Mapping
|
||||
|
||||
### Functional Requirements
|
||||
|
||||
| ID | Requirement | Description | Impacted Areas |
|
||||
| ----- | -------------------------------------------------- | ------------------------------------------------------------------------------------------------ | ------------------------------------------------------ |
|
||||
| FR-01 | Runtime detection and installation | Align worker startup with `RUNTIME` and runtime installation via PluginStore | Worker |
|
||||
| FR-02 | Wrapper‑based training using public API | Retraining must call the public `retrain(...)` method of `SientiaModel` | `fit_models`, `retrain_model` paths |
|
||||
| FR-03 | Wrapper‑based inference using public API | Inference must call `predict(...)` and `transform(...)` on the wrapper; obtain wrappers via `SientiaMLflowRepository` | Activities, shared repository |
|
||||
| FR-04 | Use shared MLflow repository | Delegate all MLflow operations (load, runs, artifacts, promotion, production lookup) to `SientiaMLflowRepository`; do not implement in Laborious | [[mlflow-shared-repository-migration-plan]] |
|
||||
| FR-05 | Model configuration for training/retraining | `model_config` must define `target` and `retention_minutes`; all models are pyfunc + wrapper (no flavor selection) | `model_config` structures |
|
||||
|
||||
### Non-Functional Requirements
|
||||
|
||||
| ID | Requirement | Description | Impacted Areas |
|
||||
| ------ | -------------------------------- | ------------------------------------------------------------------------------------------------------------------------ | --------------------------------------- |
|
||||
| NFR-01 | Consistent public API usage | All wrapper interactions must go through public `SientiaModel` methods; metadata logging is handled by the shared repository | Activities |
|
||||
| NFR-02 | Fail fast when wrappers unavailable | Where wrappers are not yet available, raise an explicit error requesting model update | Shared repository, config |
|
||||
| NFR-03 | Testability | Enable unit tests to validate wrapper-based flows and integration with shared repository | `tests/laborious` |
|
||||
| NFR-04 | Operational safety | Retraining and promotion behavior must remain auditable and robust | Retrain & promotion flows |
|
||||
|
||||
---
|
||||
|
||||
## Target Architecture
|
||||
|
||||
### Wrapper-centric model interactions
|
||||
|
||||
The target state for `laborious_temporal` is:
|
||||
|
||||
- All training and retraining logic goes through the wrapper’s public methods: `train(...)` and `retrain(...)`.
|
||||
- Inference and transformation use `transform(df)` and `predict(context, df, params)`.
|
||||
- All MLflow operations (load, runs, artifacts, promotion, production lookup) go through `SientiaMLflowRepository` (see [[mlflow-shared-repository-migration-plan]]).
|
||||
|
||||
### Use of Shared MLflow Repository
|
||||
|
||||
Laborious obtains wrappers and performs all MLflow operations via `SientiaMLflowRepository`. The shared repository (see [[mlflow-shared-repository-migration-plan]]) owns: alias-based URIs, pyfunc loading, wrapper extraction, promotion, runs, artifacts, and metadata logging. Laborious activities call `repo.load_wrapper(...)` and then the wrapper’s public methods (`transform`, `predict`, `train`, `retrain`); they do not implement MLflow logic.
|
||||
|
||||
---
|
||||
|
||||
## Implementation Plan
|
||||
|
||||
### Phase 0 — Design and Configuration Alignment
|
||||
|
||||
- **P0-01**: Catalog model types and flavors used by `laborious_temporal`:
|
||||
- For each active model:
|
||||
- Flavor will be always pyfunc
|
||||
- Whether a `_model_impl` wrapper is already present and extends `SientiaModel`.
|
||||
- **P0-02**: Define configuration fields in `model_config` for wrapper usage:
|
||||
- Example:
|
||||
- `target` field for training/retraining (already partially present).
|
||||
- `retention_minutes` field for model retention in minutes (unchanged).
|
||||
|
||||
### Phase 1 — Inference via SientiaModel Public API (using shared repository)
|
||||
|
||||
- **P1-01**: Replace `download_model` with `SientiaMLflowRepository.load_wrapper(...)`; remove local MLflow loading logic.
|
||||
- **P1-02**: Update `get_cached_operation` to call `wrapper.transform(data)` and `wrapper.predict({}, data)`; unpack returned metadata; metadata logging is handled by the shared repository.
|
||||
- **P1-03**: Ensure activities (`request_transform`, `request_predict`) remain unchanged externally (inputs/outputs unchanged).
|
||||
|
||||
### Phase 2 — Retraining via SientiaModel.retrain
|
||||
|
||||
- **P2-01**: Refactor `fit_models` to call `wrapper.retrain(data)` and `wrapper.train(...)`; use `SientiaMLflowRepository` for runs, metrics, artifacts, and promotion (no local MLflow logic).
|
||||
|
||||
### Phase 3 — MLflow Logging and Promotion (via shared repository)
|
||||
|
||||
- **P3-01**: In `create_new_experiment`, use the wrapper’s `store_model(...)` for model artifacts; use `SientiaMLflowRepository` for runs, metrics, and any additional MLflow operations.
|
||||
- **P3-02**: Use `SientiaMLflowRepository.promote_to_alias(...)` for production promotion; model registration is auto-handled by wrappers.
|
||||
|
||||
### Phase 4 — Runtime Alignment
|
||||
|
||||
- **P4-01**:`laborious_temporal` is also deployed via the runtime-aware Helm chart:
|
||||
- Read `RUNTIME` env var.
|
||||
- Install runtime via PluginStore before starting Temporal workers.
|
||||
- **P4-02**: Standardize worker queue name as `{project_name}-{runtime}-queue`.
|
||||
- **P4-03**: Fix quality pipelines to a single runtime (to be decided).
|
||||
|
||||
---
|
||||
|
||||
## Testing Strategy
|
||||
|
||||
- **T1 — Unit tests**
|
||||
- Add tests in `tests/laborious/utils/repository/test_model_repository.py` to cover:
|
||||
- Wrapper-based `get_cached_operation` for both `transform` and `predict` (using shared repository).
|
||||
- Wrapper-based `fit_models` and `retrain_model` paths calling `train` and `retrain` respectively.
|
||||
- Code coverage must be 100%.
|
||||
- **T2 — Integration tests**
|
||||
- Use existing end-to-end tests under `e2e/`:
|
||||
- Configure a model with `SientiaModel` wrapper (loaded via shared repository).
|
||||
- Run full prediction and retrain workflows; compare predictions, retrain outcomes, and MLflow artifacts.
|
||||
|
||||
---
|
||||
|
||||
## Rollout and Migration Strategy
|
||||
|
||||
- **R1 — Create models with the new architecture**
|
||||
- Update or create models in the modeling pipeline so they:
|
||||
- Use wrappers that extend `SientiaModel`.
|
||||
- Correctly implement `train`, `retrain`, `predict`, `transform`, and `store_model`.
|
||||
- Produce structured metadata in `transform_meta` and `pred_meta`.
|
||||
- Publish these models to a test store (or dedicated branch/experiment) for initial validation.
|
||||
|
||||
- **R2 — Provision runtimes for the new models**
|
||||
- Configure and install dedicated runtimes for the new models:
|
||||
- Ensure runtime dependencies (Python and system libraries) are available via PluginStore/runtime installer.
|
||||
- Validate that each runtime can:
|
||||
- Load the wrapper through MLflow.
|
||||
- Execute `transform` and `predict` end-to-end on sample data.
|
||||
|
||||
- **R3 — Run new models in real workflows**
|
||||
- Integrate the new wrapper-based models into real Laborious workflows, initially in non-critical environments:
|
||||
- Route only a subset of flows or entities to the new models.
|
||||
- Monitor logs (including metadata), business metrics, and retraining/promotion behavior.
|
||||
- Promote these models to production using aliases (`@production`) in MLflow 3+.
|
||||
|
||||
- **R4 — Migrate remaining models progressively**
|
||||
- Define migration waves by model family:
|
||||
- For each existing model:
|
||||
- Create or adapt a `SientiaModel` wrapper.
|
||||
- Provision the corresponding runtime.
|
||||
- Execute the test cycle (T1–T2) from the previous section.
|
||||
- Update aliases so production traffic uses the new wrapper.
|
||||
- After all models are migrated:
|
||||
- Remove legacy stage-based paths (`Production`) and non-wrapper models.
|
||||
- Simplify the codebase to assume wrappers + aliases only.
|
||||
|
||||
---
|
||||
|
||||
## Related Documents
|
||||
|
||||
- [[mlflow-shared-repository-migration-plan|MLflow Shared Repository Migration Plan]] — Concepts implemented in the common interface (production lookup, wrapper loading, promotion, metadata logging, etc.)
|
||||
- [[model-manager-plugin-store-migration-plan|Model Manager PluginStore Migration Plan]]
|
||||
- [[analytics-implementation-plan|Runtime Analytics Helm Implementation Plan]]
|
||||
- [[analytics|Runtime Analytics Architecture and Analysis]]
|
||||
- [[../model-plugin-system/06-end-to-end-flow|Model Plugin System — End-to-End Flow]]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}',
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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'),
|
||||
|
||||
@@ -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))
|
||||
|
||||
32
laborious/utils/repository/minio_manager.py
Normal file
32
laborious/utils/repository/minio_manager.py
Normal 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)
|
||||
1493
laborious/utils/repository/model_repository.py
Normal file
1493
laborious/utils/repository/model_repository.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
18
requirements-light.txt
Normal 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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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})
|
||||
|
||||
1715
tests/laborious/utils/repository/test_model_repository.py
Normal file
1715
tests/laborious/utils/repository/test_model_repository.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
|
||||
14
tests/laborious/worker/test_runtime_task_queues.py
Normal file
14
tests/laborious/worker/test_runtime_task_queues.py
Normal 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'
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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',
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
319
values.yaml
319
values.yaml
@@ -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
|
||||
Reference in New Issue
Block a user