SIENTIAPDE-1646 Sync full repo from local release/SIENTIAPDE-1646 (272e02d)
This commit is contained in:
22
.env.example
22
.env.example
@@ -11,6 +11,22 @@ MLFLOW_PORT="80"
|
|||||||
MLFLOW_USERNAME="aignosi"
|
MLFLOW_USERNAME="aignosi"
|
||||||
MLFLOW_PASSWORD="mlflow_password"
|
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_ID="1"
|
||||||
OPC_URL="opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
|
OPC_URL="opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
|
||||||
|
|
||||||
@@ -27,3 +43,9 @@ MONGODB_PASSWORD="mongo_db_password"
|
|||||||
MONGODB_URL="my-release-mongodb.mongodb.svc.cluster.local:27017"
|
MONGODB_URL="my-release-mongodb.mongodb.svc.cluster.local:27017"
|
||||||
MONGODB_DATABASE="sientia"
|
MONGODB_DATABASE="sientia"
|
||||||
MONGODB_TTL_INDEX_HOURS="1"
|
MONGODB_TTL_INDEX_HOURS="1"
|
||||||
|
|
||||||
|
MINIO_ENDPOINT_URL="http://localhost:9000"
|
||||||
|
MINIO_ACCESS_KEY="sientia"
|
||||||
|
MINIO_SECRET_KEY="sientia"
|
||||||
|
MINIO_REGION_NAME="sa-east-1"
|
||||||
|
MINIO_DEFAULT_BUCKET="sientia"
|
||||||
69
.github/workflows/quality-gate.yml
vendored
69
.github/workflows/quality-gate.yml
vendored
@@ -1,74 +1,17 @@
|
|||||||
name: Quality gate
|
name: Quality gate
|
||||||
|
|
||||||
on:
|
on:
|
||||||
push:
|
|
||||||
branches:
|
|
||||||
- main
|
|
||||||
pull_request:
|
pull_request:
|
||||||
branches:
|
branches:
|
||||||
- main
|
- main
|
||||||
types: [ opened, synchronize, reopened ]
|
types: [ opened, synchronize, reopened ]
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
sonar:
|
quality-gate:
|
||||||
name: SonarQube Analysis
|
uses: Aignosi/github_workflow_templates/.github/workflows/dataops-module-quality-gate.yml@main
|
||||||
runs-on: ubuntu-latest
|
|
||||||
permissions: write-all
|
permissions: write-all
|
||||||
steps:
|
|
||||||
- name: ⬇️ Checkout Code
|
|
||||||
uses: actions/checkout@v4
|
|
||||||
with:
|
with:
|
||||||
fetch-depth: 0
|
project_name: 'laborious'
|
||||||
persist-credentials: false
|
repositories: 'sientia-dataops-library, sientia-mlops-library'
|
||||||
|
requirements_file: 'requirements-light.txt'
|
||||||
- name: Generate App Token
|
secrets: inherit
|
||||||
id: generate-app-token
|
|
||||||
uses: actions/create-github-app-token@v1
|
|
||||||
with:
|
|
||||||
app-id: ${{ secrets.APP_ID }}
|
|
||||||
private-key: ${{ secrets.APP_PRIVATE_KEY }}
|
|
||||||
owner: 'Aignosi'
|
|
||||||
repositories: 'sientia-dataops-library,sientia-mlops-library'
|
|
||||||
|
|
||||||
- name: Prepare requirements.txt
|
|
||||||
id: prepare-requirements
|
|
||||||
run: |
|
|
||||||
sed -e "s|git+ssh://git@github.com/|git+https://github.com/|g" \
|
|
||||||
-e "s|git@github.com:|git+https://github.com/|g" \
|
|
||||||
requirements.txt > requirements_prepared.txt
|
|
||||||
echo "PROCESSED_REQUIREMENTS_FILE=requirements_prepared.txt" >> $GITHUB_OUTPUT
|
|
||||||
|
|
||||||
- name: Configure Git to use App Token
|
|
||||||
env:
|
|
||||||
GH_APP_TOKEN: ${{ steps.generate-app-token.outputs.token }}
|
|
||||||
run: |
|
|
||||||
git config --global url."https://oauth2:${GH_APP_TOKEN}@github.com/".insteadOf "https://github.com/"
|
|
||||||
|
|
||||||
- name: 🔧 Setup Python
|
|
||||||
uses: actions/setup-python@v4
|
|
||||||
with:
|
|
||||||
python-version: "3.11"
|
|
||||||
|
|
||||||
- name: 🗄️ Cache Python dependencies
|
|
||||||
uses: actions/cache@v3
|
|
||||||
with:
|
|
||||||
path: ~/.cache/pip
|
|
||||||
key: ${{ runner.os }}-pip-${{ hashFiles(steps.prepare-requirements.outputs.PROCESSED_REQUIREMENTS_FILE) }}
|
|
||||||
restore-keys: |
|
|
||||||
${{ runner.os }}-pip-
|
|
||||||
|
|
||||||
- name: 📦 Install Dependencies
|
|
||||||
run: |
|
|
||||||
python -m pip install --upgrade pip
|
|
||||||
pip install -r ${{ steps.prepare-requirements.outputs.PROCESSED_REQUIREMENTS_FILE }}
|
|
||||||
pip install pytest pytest-cov pytest-asyncio
|
|
||||||
|
|
||||||
- name: 🧪 Run Tests with Pytest
|
|
||||||
run: |
|
|
||||||
pytest tests --junitxml=pytest.xml --cov=laborious --cov-report=xml --cov-report=term
|
|
||||||
|
|
||||||
- name: Run SonarQube Analysis
|
|
||||||
uses: SonarSource/sonarqube-scan-action@v5
|
|
||||||
env:
|
|
||||||
SONAR_TOKEN: ${{ secrets.SONAR_TOKEN }}
|
|
||||||
SONAR_HOST_URL: ${{ secrets.SONAR_HOST_URL }}
|
|
||||||
25
.github/workflows/release.yml
vendored
Normal file
25
.github/workflows/release.yml
vendored
Normal file
@@ -0,0 +1,25 @@
|
|||||||
|
name: Create Release on Merge to Main
|
||||||
|
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
types: [closed]
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
workflow_dispatch:
|
||||||
|
inputs:
|
||||||
|
version:
|
||||||
|
description: 'Version to release'
|
||||||
|
required: false
|
||||||
|
type: string
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
release:
|
||||||
|
if: |
|
||||||
|
(github.event_name == 'pull_request' && github.event.pull_request.merged == true) ||
|
||||||
|
github.event_name == 'workflow_dispatch'
|
||||||
|
uses: Aignosi/github_workflow_templates/.github/workflows/dataops-module-release.yml@main
|
||||||
|
permissions: write-all
|
||||||
|
with:
|
||||||
|
project_name: 'laborious'
|
||||||
|
release_version: ${{ github.event.inputs.version || '' }}
|
||||||
|
secrets: inherit
|
||||||
10
.gitignore
vendored
10
.gitignore
vendored
@@ -37,6 +37,7 @@ __pycache__/
|
|||||||
# Ignorar coverage
|
# Ignorar coverage
|
||||||
htmlcov/
|
htmlcov/
|
||||||
.coverage
|
.coverage
|
||||||
|
coverage.xml
|
||||||
|
|
||||||
# git keys
|
# git keys
|
||||||
git_key*
|
git_key*
|
||||||
@@ -46,3 +47,12 @@ git_log
|
|||||||
.env
|
.env
|
||||||
|
|
||||||
tmp/
|
tmp/
|
||||||
|
catboost_info/
|
||||||
|
|
||||||
|
.ruff_cache/
|
||||||
|
.mypy_cache/
|
||||||
|
mlruns/
|
||||||
|
|
||||||
|
relatorio*
|
||||||
|
openspec/*
|
||||||
|
.cursor/*
|
||||||
149
docs/E2E_TEST_REPORT.md
Normal file
149
docs/E2E_TEST_REPORT.md
Normal file
@@ -0,0 +1,149 @@
|
|||||||
|
# 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.
|
||||||
164
docs/opc-communication.md
Normal file
164
docs/opc-communication.md
Normal file
@@ -0,0 +1,164 @@
|
|||||||
|
# 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.
|
||||||
|
|
||||||
|
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).
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
```text
|
||||||
|
Worker (long-lived)
|
||||||
|
└── OpcRepository per OPC server id (from OPC_CONFIG / env)
|
||||||
|
├── 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
|
||||||
|
|
||||||
|
Temporal activity write_opc_data
|
||||||
|
└── OPC.manage_output_tags → write_data per tag (sequential per activity)
|
||||||
|
```
|
||||||
|
|
||||||
|
One worker process holds one `OpcRepository` instance per configured server. Multiple Temporal activities can call `write_data` concurrently on the same repository.
|
||||||
|
|
||||||
|
## Connection lifecycle
|
||||||
|
|
||||||
|
| Phase | Behavior |
|
||||||
|
|-------|----------|
|
||||||
|
| 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 |
|
||||||
|
|
||||||
|
### Session and channel timeouts
|
||||||
|
|
||||||
|
Requested session and secure-channel lifetime: **10 minutes** (`OPC_UA_SESSION_AND_CHANNEL_TIMEOUT_MS` in `opc_repository.py`). The server may revise these values; negotiated values are logged after connect and exposed as `opc_session_revised_timeout_milliseconds`.
|
||||||
|
|
||||||
|
### 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.
|
||||||
|
|
||||||
|
## Concurrency: connection lock and session readiness
|
||||||
|
|
||||||
|
To allow **multiple concurrent writes** when the session is healthy, but **block all writes** while the connection is being torn down or re-established:
|
||||||
|
|
||||||
|
| 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 methods (caller holds `_connection_lock` for `_*_locked` helpers):**
|
||||||
|
|
||||||
|
| Method | Role |
|
||||||
|
|--------|------|
|
||||||
|
| `_create_client()` | Create asyncua `Client` + optional `set_security`; raises if `client` already exists |
|
||||||
|
| `_open_session()` | `client.connect()` + metrics; raises if session already open or client missing |
|
||||||
|
| `_connect_locked()` | `_create_client()` (when needed) + `_open_session()`; raises if already connected |
|
||||||
|
| `_disconnect_locked()` | Teardown session and clear `client` |
|
||||||
|
| `_reconnect_locked()` | `_disconnect_locked()` + `_connect_locked()`; sets `last_reconnection_time` |
|
||||||
|
|
||||||
|
Public `connect()` / `disconnect()` acquire the lock and call `_connect_locked()` / `_disconnect_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).
|
||||||
|
|
||||||
|
**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()`.
|
||||||
|
3. `_session_ready` is set on successful `_open_session()`.
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
## Tier-1 `Bad*` errors and reconnect
|
||||||
|
|
||||||
|
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
|
||||||
|
- No reconnect task is already running.
|
||||||
|
|
||||||
|
There is **no write retry**: the failed export is not sent again in the same activity.
|
||||||
|
|
||||||
|
## Prediction confidence and PostgreSQL comments
|
||||||
|
|
||||||
|
| `prediction_confidence` | Meaning |
|
||||||
|
|-------------------------|---------|
|
||||||
|
| (unchanged) | Successful OPC export |
|
||||||
|
| **12** | Generic OPC write failure (`OPC_WRITTING_ERROR_CONFIDENCE`) |
|
||||||
|
| **14** | Tier-1 session/channel `Bad*` on export (`OPC_SESSION_BAD_CONFIDENCE`) |
|
||||||
|
| **14** | Write while reconnect in progress (`OPC_SESSION_BAD_CONFIDENCE`, comment `OPC UA reconnect in progress`) |
|
||||||
|
| **13** | PI Web API write failure (separate path) |
|
||||||
|
|
||||||
|
Session/channel errors use a stable comment for counting:
|
||||||
|
|
||||||
|
```text
|
||||||
|
OPC UA session/channel error: BadSessionIdInvalid
|
||||||
|
```
|
||||||
|
|
||||||
|
Reconnect-in-progress exports use:
|
||||||
|
|
||||||
|
```text
|
||||||
|
OPC UA reconnect in progress
|
||||||
|
```
|
||||||
|
|
||||||
|
Example SQL:
|
||||||
|
|
||||||
|
```sql
|
||||||
|
SELECT count(*) FROM predictions WHERE prediction_confidence = 14;
|
||||||
|
SELECT count(*) FROM predictions WHERE comments LIKE 'OPC UA session/channel error:%';
|
||||||
|
```
|
||||||
|
|
||||||
|
## Prometheus metrics (`opc_*`)
|
||||||
|
|
||||||
|
Defined in [`laborious/metrics.py`](../laborious/metrics.py). Do not rename in production without a dashboard migration.
|
||||||
|
|
||||||
|
| Metric | Purpose |
|
||||||
|
|--------|---------|
|
||||||
|
| `opc_connections_initiated_total` | Connection attempts |
|
||||||
|
| `opc_connections_failed_total` | Failed connects |
|
||||||
|
| `opc_connection_status` | Gauge 1=connected, 0=disconnected |
|
||||||
|
| `opc_session_created_total` | Session established after connect |
|
||||||
|
| `opc_session_closed_total` | Disconnect initiated |
|
||||||
|
| `opc_session_revised_timeout_milliseconds` | Negotiated session timeout (ms) |
|
||||||
|
| `opc_write_attempts_total` | Per write; label `result` = `OK` or exception name |
|
||||||
|
| `opc_write_inter_arrival_over_session_timeout_total` | Successful writes spaced longer than revised session timeout |
|
||||||
|
|
||||||
|
Legacy activity metrics: `laborious_prediction_opc_writing_count`, `laborious_prediction_opc_writing_response_time_monitor`.
|
||||||
|
|
||||||
|
## Environment variables
|
||||||
|
|
||||||
|
| Variable | Default | Description |
|
||||||
|
|----------|---------|-------------|
|
||||||
|
| `OPC_CONFIG` | — | JSON map of server configs (overrides single-server env) |
|
||||||
|
| `OPC_ID` | `1` | Server id |
|
||||||
|
| `OPC_URL` | `opc.tcp://localhost:4840` | Endpoint |
|
||||||
|
| `OPC_SERVER_NAME` | `default_server` | Label for metrics/logs |
|
||||||
|
| `OPC_SERVER_URI` | same as URL | Application URI / cert SAN |
|
||||||
|
| `OPC_CERT_PATH` | — | Client certificate (secure mode) |
|
||||||
|
| `OPC_PRIVATE_KEY_PATH` | — | Client private key |
|
||||||
|
| `OPC_SERVER_CERT_PATH` | — | Server certificate |
|
||||||
|
| `OPC_RECONNECTION_INTERVAL` | `120` | Minimum seconds between reconnects |
|
||||||
|
|
||||||
|
## Operations checklist
|
||||||
|
|
||||||
|
- Correlate `BadSessionIdInvalid` in `opc_write_attempts_total` with `opc_session_closed_total` / `opc_session_created_total` (reconnect may finish after the row is stored with confidence 14).
|
||||||
|
- Use confidence **14** and comment prefix for session invalidation rates; use **12** for other OPC failures.
|
||||||
|
- Respect `OPC_RECONNECTION_INTERVAL` under parallel load; bursts of confidence 14 are expected until the next successful cycle.
|
||||||
|
|
||||||
|
## Related tests
|
||||||
|
|
||||||
|
- Unit: [`tests/laborious/utils/repository/test_opc_repository.py`](../tests/laborious/utils/repository/test_opc_repository.py)
|
||||||
|
- Unit: [`tests/laborious/activities/test_opc.py`](../tests/laborious/activities/test_opc.py)
|
||||||
|
- E2E (mock OPC): [`e2e/test_predictions_batch_format_export.py`](../e2e/test_predictions_batch_format_export.py)
|
||||||
|
- E2E (in-process asyncua server + real `OpcRepository`): [`e2e/test_opc_real_server.py`](../e2e/test_opc_real_server.py) — scenarios 3.1.2, 3.2.2, 3.2.4, 3.2.5
|
||||||
|
- Scenarios: [`e2e/scenarios.md`](../e2e/scenarios.md)
|
||||||
55
docs/sientia_model_drift_jensen_shannon_change.md
Normal file
55
docs/sientia_model_drift_jensen_shannon_change.md
Normal file
@@ -0,0 +1,55 @@
|
|||||||
|
# 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).
|
||||||
3
e2e/__init__.py
Normal file
3
e2e/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
"""
|
||||||
|
End-to-end tests for laborious temporal workflows.
|
||||||
|
"""
|
||||||
603
e2e/conftest.py
Normal file
603
e2e/conftest.py
Normal file
@@ -0,0 +1,603 @@
|
|||||||
|
"""Pytest configuration and fixtures for E2E tests."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
from sqlalchemy import create_engine
|
||||||
|
from testcontainers.core.container import DockerContainer
|
||||||
|
from testcontainers.minio import MinioContainer
|
||||||
|
from testcontainers.postgres import PostgresContainer
|
||||||
|
from temporalio.testing import WorkflowEnvironment
|
||||||
|
from temporalio.worker import Worker
|
||||||
|
|
||||||
|
from e2e.opc_test_server import OpcE2ETestServer
|
||||||
|
from laborious.activities.activities import Activities
|
||||||
|
from laborious.workflows.drift import Drift
|
||||||
|
from laborious.workflows.minimal_retrain import MinimalRetrain
|
||||||
|
from laborious.workflows.predictions_batch import PredictionsBatch
|
||||||
|
from laborious.workflows.simple_metrics import SimpleMetrics
|
||||||
|
from laborious.workflows.sub_workflows.format_and_export_prediction import (
|
||||||
|
FormatAndExportPrediction,
|
||||||
|
)
|
||||||
|
from laborious.workflows.sub_workflows.prediction_process import PredictionProcess
|
||||||
|
from sientia_do.notifications.handlers import CoreNotificationHandler
|
||||||
|
from sientia_do.observability.logger import Logger
|
||||||
|
from sientia_do.observability.metrics_controller import MetricsController
|
||||||
|
|
||||||
|
# Single source of truth for the test database schema. Mirrors the production
|
||||||
|
# DDL for ``sientia_data`` so any production change can be pasted directly into
|
||||||
|
# this file (see ``e2e/db_schema.sql``) without touching Python.
|
||||||
|
DB_SCHEMA_SQL_PATH = Path(__file__).parent / 'db_schema.sql'
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='session')
|
||||||
|
def postgres_container():
|
||||||
|
"""PostgreSQL testcontainer used by all E2E tests."""
|
||||||
|
postgres = PostgresContainer('postgres:15')
|
||||||
|
postgres.start()
|
||||||
|
yield postgres
|
||||||
|
postgres.stop()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='session')
|
||||||
|
def minio_container():
|
||||||
|
"""MinIO testcontainer used by E2E offload and payload retrieval paths."""
|
||||||
|
minio = MinioContainer()
|
||||||
|
minio.start()
|
||||||
|
yield minio
|
||||||
|
minio.stop()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='session')
|
||||||
|
def mongo_container():
|
||||||
|
"""MongoDB testcontainer used by real CoreNotificationHandler."""
|
||||||
|
mongo = DockerContainer('mongo:7').with_exposed_ports(27017)
|
||||||
|
mongo.start()
|
||||||
|
yield mongo
|
||||||
|
mongo.stop()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def postgres_engine(postgres_container):
|
||||||
|
"""SQLAlchemy engine bound to the PostgreSQL testcontainer."""
|
||||||
|
engine = create_engine(postgres_container.get_connection_url())
|
||||||
|
yield engine
|
||||||
|
engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
|
def _create_schema_and_tables(engine):
|
||||||
|
"""
|
||||||
|
Create all schemas/tables required by workflow and activity paths.
|
||||||
|
|
||||||
|
Loads the DDL from ``e2e/db_schema.sql`` (single source of truth that
|
||||||
|
mirrors the production schema). The SQL file is executed via the raw
|
||||||
|
DBAPI cursor so multi-statement DDL is supported.
|
||||||
|
"""
|
||||||
|
sql_text = DB_SCHEMA_SQL_PATH.read_text(encoding='utf-8')
|
||||||
|
with engine.begin() as conn:
|
||||||
|
conn.exec_driver_sql(sql_text)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(autouse=True)
|
||||||
|
def setup_postgres_schema_and_tables(postgres_engine):
|
||||||
|
"""Ensure required schema and tables exist before each E2E test."""
|
||||||
|
_create_schema_and_tables(postgres_engine)
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def mock_logger():
|
||||||
|
"""Logger double with readable console output for E2E runs."""
|
||||||
|
logger = MagicMock(spec=Logger)
|
||||||
|
logger.info = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
|
||||||
|
logger.debug = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
|
||||||
|
logger.error = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
|
||||||
|
logger.warning = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
|
||||||
|
logger.custom_info = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
|
||||||
|
logger.custom_debug = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
|
||||||
|
logger.custom_error = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
|
||||||
|
logger.custom_warning = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
|
||||||
|
return logger
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def metrics_controller(mock_logger):
|
||||||
|
"""Real metrics controller for E2E observability paths."""
|
||||||
|
return MetricsController(logger=mock_logger)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def notification_handler(mock_logger, mongo_container):
|
||||||
|
"""Real notification handler using MongoDB testcontainer."""
|
||||||
|
mongo_port = mongo_container.get_exposed_port(27017)
|
||||||
|
handler = CoreNotificationHandler(
|
||||||
|
connection_string=f'mongodb://localhost:{mongo_port}',
|
||||||
|
database='test_db',
|
||||||
|
logger=mock_logger,
|
||||||
|
project_name='laborious',
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
yield handler
|
||||||
|
finally:
|
||||||
|
handler.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def notification_inserts(notification_handler):
|
||||||
|
"""Spy on real Mongo insert calls issued by notification handler."""
|
||||||
|
collection = notification_handler.mongo_collection
|
||||||
|
original_insert_one = collection.insert_one
|
||||||
|
spy = MagicMock(wraps=original_insert_one)
|
||||||
|
collection.insert_one = spy
|
||||||
|
try:
|
||||||
|
yield spy
|
||||||
|
finally:
|
||||||
|
collection.insert_one = original_insert_one
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeModelWrapper:
|
||||||
|
"""External MLflow wrapper double used by repository stub."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.transform = MagicMock(side_effect=self._default_transform)
|
||||||
|
self.predict = MagicMock(side_effect=self._default_predict)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _default_transform(data: pd.DataFrame):
|
||||||
|
result = pd.DataFrame(
|
||||||
|
{
|
||||||
|
'feature_1': [0.234] * len(data),
|
||||||
|
'feature_2': [0.783] * len(data),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
result.index = data.index
|
||||||
|
return result, {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _default_predict(_params: dict, data: pd.DataFrame):
|
||||||
|
pred = pd.DataFrame([0.5] * len(data), columns=['placeholder'])
|
||||||
|
pred.index = data.index
|
||||||
|
return pred, {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def mlflow_repository_stub():
|
||||||
|
"""External MLflow repository stub."""
|
||||||
|
repo = MagicMock()
|
||||||
|
wrapper = _FakeModelWrapper()
|
||||||
|
repo.stub_wrapper = wrapper
|
||||||
|
repo.get_cached_model = MagicMock(return_value=wrapper)
|
||||||
|
repo._client = MagicMock()
|
||||||
|
return repo
|
||||||
|
|
||||||
|
|
||||||
|
class _FakePIWebAPIClient:
|
||||||
|
"""External PI Web API client stub with deterministic responses."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._responses = None
|
||||||
|
self.write_value = MagicMock(side_effect=self._write_value)
|
||||||
|
self.close = MagicMock()
|
||||||
|
|
||||||
|
def set_side_effect(self, side_effect):
|
||||||
|
self._responses = side_effect
|
||||||
|
|
||||||
|
def _write_value(self, web_ids, value, metadata=None, **kwargs):
|
||||||
|
if isinstance(self._responses, Exception):
|
||||||
|
raise self._responses
|
||||||
|
if isinstance(self._responses, list):
|
||||||
|
item = self._responses.pop(0)
|
||||||
|
if isinstance(item, Exception):
|
||||||
|
raise item
|
||||||
|
return item
|
||||||
|
if callable(self._responses):
|
||||||
|
return self._responses(web_ids=web_ids, value=value, metadata=metadata, **kwargs)
|
||||||
|
return [{'WebId': wid, 'Errors': []} for wid in web_ids]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def pi_web_api_client_stub():
|
||||||
|
"""PI Web API stub fixture."""
|
||||||
|
return _FakePIWebAPIClient()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def opc_repository_stub():
|
||||||
|
"""OPC external dependency stub."""
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.write_data = MagicMock(return_value=(True, {'response_time': 0.1}))
|
||||||
|
repo.disconnect = MagicMock()
|
||||||
|
return repo
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
def plugin_store_stub():
|
||||||
|
"""Plugin store external dependency stub."""
|
||||||
|
return MagicMock()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def test_activities(
|
||||||
|
postgres_container,
|
||||||
|
minio_container,
|
||||||
|
mock_logger,
|
||||||
|
notification_handler,
|
||||||
|
metrics_controller,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
plugin_store_stub,
|
||||||
|
pi_web_api_client_stub,
|
||||||
|
opc_repository_stub,
|
||||||
|
):
|
||||||
|
"""Activities with real infra and external-system stubs only."""
|
||||||
|
minio_client = minio_container.get_client()
|
||||||
|
if not minio_client.bucket_exists('test-bucket'):
|
||||||
|
minio_client.make_bucket('test-bucket')
|
||||||
|
minio_port = minio_container.get_exposed_port(9000)
|
||||||
|
|
||||||
|
activities = Activities(
|
||||||
|
postgres_config={
|
||||||
|
'host': 'localhost',
|
||||||
|
'port': int(postgres_container.get_exposed_port(5432)),
|
||||||
|
'user': postgres_container.username,
|
||||||
|
'password': postgres_container.password,
|
||||||
|
'dbname': postgres_container.dbname,
|
||||||
|
'min_connections': 1,
|
||||||
|
'max_connections': 5,
|
||||||
|
},
|
||||||
|
plugin_store=plugin_store_stub,
|
||||||
|
minio_config={
|
||||||
|
'endpoint_url': f'localhost:{minio_port}',
|
||||||
|
'access_key': 'minioadmin',
|
||||||
|
'secret_key': 'minioadmin',
|
||||||
|
'default_bucket': 'test-bucket',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
|
},
|
||||||
|
opc_config={},
|
||||||
|
pi_web_api_config={
|
||||||
|
'base_url': 'http://localhost:8080',
|
||||||
|
'auth_type': 'bearer',
|
||||||
|
'auth_token': 'test_token',
|
||||||
|
},
|
||||||
|
logger=mock_logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
|
mlflow_repository=mlflow_repository_stub,
|
||||||
|
)
|
||||||
|
activities.pi_web_api_client = pi_web_api_client_stub
|
||||||
|
activities.opc_repository = {'1': opc_repository_stub}
|
||||||
|
try:
|
||||||
|
yield activities
|
||||||
|
finally:
|
||||||
|
activities.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def test_activities_real_minio(
|
||||||
|
postgres_container,
|
||||||
|
minio_container,
|
||||||
|
mock_logger,
|
||||||
|
notification_handler,
|
||||||
|
metrics_controller,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
plugin_store_stub,
|
||||||
|
pi_web_api_client_stub,
|
||||||
|
opc_repository_stub,
|
||||||
|
):
|
||||||
|
"""Compatibility alias for offload tests."""
|
||||||
|
minio_client = minio_container.get_client()
|
||||||
|
if not minio_client.bucket_exists('test-bucket'):
|
||||||
|
minio_client.make_bucket('test-bucket')
|
||||||
|
minio_port = minio_container.get_exposed_port(9000)
|
||||||
|
|
||||||
|
activities = Activities(
|
||||||
|
postgres_config={
|
||||||
|
'host': 'localhost',
|
||||||
|
'port': int(postgres_container.get_exposed_port(5432)),
|
||||||
|
'user': postgres_container.username,
|
||||||
|
'password': postgres_container.password,
|
||||||
|
'dbname': postgres_container.dbname,
|
||||||
|
'min_connections': 1,
|
||||||
|
'max_connections': 5,
|
||||||
|
},
|
||||||
|
plugin_store=plugin_store_stub,
|
||||||
|
minio_config={
|
||||||
|
'endpoint_url': f'localhost:{minio_port}',
|
||||||
|
'access_key': 'minioadmin',
|
||||||
|
'secret_key': 'minioadmin',
|
||||||
|
'default_bucket': 'test-bucket',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
|
},
|
||||||
|
opc_config={},
|
||||||
|
pi_web_api_config={
|
||||||
|
'base_url': 'http://localhost:8080',
|
||||||
|
'auth_type': 'bearer',
|
||||||
|
'auth_token': 'test_token',
|
||||||
|
},
|
||||||
|
logger=mock_logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
|
mlflow_repository=mlflow_repository_stub,
|
||||||
|
)
|
||||||
|
activities.pi_web_api_client = pi_web_api_client_stub
|
||||||
|
activities.opc_repository = {'1': opc_repository_stub}
|
||||||
|
try:
|
||||||
|
yield activities
|
||||||
|
finally:
|
||||||
|
activities.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
def _worker_activity_list(test_activities: Activities):
|
||||||
|
"""List of registered activity callables used by Temporal worker in E2E."""
|
||||||
|
return [
|
||||||
|
test_activities.load_query_with_minio_offload,
|
||||||
|
test_activities.cleanup_minio_objects_expired,
|
||||||
|
test_activities.input_gate,
|
||||||
|
test_activities.request_transform,
|
||||||
|
test_activities.mlflow_response_gate,
|
||||||
|
test_activities.mlflow_content_gate,
|
||||||
|
test_activities.request_predict,
|
||||||
|
test_activities.repeat_last_prediction,
|
||||||
|
test_activities.format_prediction,
|
||||||
|
test_activities.format_transformed_data,
|
||||||
|
test_activities.format_default_prediction,
|
||||||
|
test_activities.write_pi_web_api_data,
|
||||||
|
test_activities.write_opc_data,
|
||||||
|
test_activities.export_data_to_postgres,
|
||||||
|
test_activities.export_payload_to_postgres,
|
||||||
|
test_activities.write_metrics,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_test_env():
|
||||||
|
"""Temporal test environment with time-skipping."""
|
||||||
|
env = await WorkflowEnvironment.start_time_skipping()
|
||||||
|
async with env:
|
||||||
|
yield env
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker(temporal_test_env, test_activities):
|
||||||
|
"""Temporal worker for full predictions-batch and child workflows."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction],
|
||||||
|
activities=_worker_activity_list(test_activities),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker_real_minio(temporal_test_env, test_activities_real_minio):
|
||||||
|
"""Temporal worker alias for tests that emphasize MinIO behavior."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction],
|
||||||
|
activities=_worker_activity_list(test_activities_real_minio),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
|
def _drift_worker_activity_list(test_activities: Activities):
|
||||||
|
"""Activity callables registered on the drift Temporal worker."""
|
||||||
|
return [
|
||||||
|
test_activities.load_custom_query,
|
||||||
|
test_activities.get_reference_data,
|
||||||
|
test_activities.calculate_drift,
|
||||||
|
test_activities.export_data_to_postgres,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker_drift(temporal_test_env, test_activities):
|
||||||
|
"""Temporal worker registered with the Drift workflow and its activities."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[Drift],
|
||||||
|
activities=_drift_worker_activity_list(test_activities),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
|
def _simple_metrics_worker_activity_list(test_activities: Activities):
|
||||||
|
"""Activity callables registered on the simple-metrics Temporal worker."""
|
||||||
|
return [
|
||||||
|
test_activities.load_custom_query,
|
||||||
|
test_activities.calculate_simple_metrics,
|
||||||
|
test_activities.export_data_to_postgres,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker_simple_metrics(temporal_test_env, test_activities):
|
||||||
|
"""Temporal worker registered with the SimpleMetrics workflow and its activities."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[SimpleMetrics],
|
||||||
|
activities=_simple_metrics_worker_activity_list(test_activities),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
|
def _minimal_retrain_worker_activity_list(test_activities: Activities):
|
||||||
|
"""Activity callables registered on the minimal-retrain Temporal worker."""
|
||||||
|
return [
|
||||||
|
test_activities.load_query_with_minio_offload,
|
||||||
|
test_activities.retrain_model,
|
||||||
|
test_activities.update_production_model,
|
||||||
|
test_activities.format_retrain_report,
|
||||||
|
test_activities.export_data_to_postgres,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker_minimal_retrain(temporal_test_env, test_activities):
|
||||||
|
"""Temporal worker registered with the MinimalRetrain workflow and its activities."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[MinimalRetrain],
|
||||||
|
activities=_minimal_retrain_worker_activity_list(test_activities),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
|
def _connect_activities_to_opc_server(activities: Activities, server_url: str) -> None:
|
||||||
|
"""
|
||||||
|
Initialize OPC repositories and block until the E2E server session is ready.
|
||||||
|
|
||||||
|
Runs synchronously (typically via ``asyncio.to_thread``) so the asyncua test
|
||||||
|
server event loop is not blocked during ``Client.connect()``.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
activities (Activities): Worker activities under test.
|
||||||
|
server_url (str): ``opc.tcp://`` URL from ``OpcE2ETestServer``.
|
||||||
|
"""
|
||||||
|
activities.init_opc()
|
||||||
|
repo = activities.opc_repository['1']
|
||||||
|
deadline = time.monotonic() + 30.0
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
if repo._session_ready.is_set():
|
||||||
|
return
|
||||||
|
connected, _ = repo.connect()
|
||||||
|
if connected:
|
||||||
|
return
|
||||||
|
time.sleep(0.5)
|
||||||
|
raise RuntimeError(f'Could not connect OpcRepository to OPC E2E server at {server_url}')
|
||||||
|
|
||||||
|
|
||||||
|
def _build_e2e_opc_config(server_url: str) -> dict[str, dict]:
|
||||||
|
"""
|
||||||
|
OPC server config for E2E Activities pointing at an in-process asyncua server.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
server_url (str): ``opc.tcp://`` endpoint from ``OpcE2ETestServer``.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
dict: ``opc_config`` payload for ``Activities`` (server id ``1``).
|
||||||
|
"""
|
||||||
|
return {
|
||||||
|
'1': {
|
||||||
|
'id': '1',
|
||||||
|
'server_name': 'e2e_opcua',
|
||||||
|
'url': server_url,
|
||||||
|
'server_uri': server_url,
|
||||||
|
'cert_path': None,
|
||||||
|
'private_key_path': None,
|
||||||
|
'server_cert_path': None,
|
||||||
|
'reconnection_interval': 0,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
async def opc_e2e_server():
|
||||||
|
"""In-process asyncua server with writable prediction/confidence nodes."""
|
||||||
|
server = OpcE2ETestServer()
|
||||||
|
await server.start()
|
||||||
|
await asyncio.sleep(0.5)
|
||||||
|
try:
|
||||||
|
yield server
|
||||||
|
finally:
|
||||||
|
await server.stop()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def test_activities_real_opc(
|
||||||
|
postgres_container,
|
||||||
|
minio_container,
|
||||||
|
mock_logger,
|
||||||
|
notification_handler,
|
||||||
|
metrics_controller,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
plugin_store_stub,
|
||||||
|
pi_web_api_client_stub,
|
||||||
|
opc_e2e_server: OpcE2ETestServer,
|
||||||
|
):
|
||||||
|
"""Activities with real OpcRepository connected to the in-process OPC UA server."""
|
||||||
|
minio_client = minio_container.get_client()
|
||||||
|
if not minio_client.bucket_exists('test-bucket'):
|
||||||
|
minio_client.make_bucket('test-bucket')
|
||||||
|
minio_port = minio_container.get_exposed_port(9000)
|
||||||
|
|
||||||
|
activities = Activities(
|
||||||
|
postgres_config={
|
||||||
|
'host': 'localhost',
|
||||||
|
'port': int(postgres_container.get_exposed_port(5432)),
|
||||||
|
'user': postgres_container.username,
|
||||||
|
'password': postgres_container.password,
|
||||||
|
'dbname': postgres_container.dbname,
|
||||||
|
'min_connections': 1,
|
||||||
|
'max_connections': 5,
|
||||||
|
},
|
||||||
|
plugin_store=plugin_store_stub,
|
||||||
|
minio_config={
|
||||||
|
'endpoint_url': f'localhost:{minio_port}',
|
||||||
|
'access_key': 'minioadmin',
|
||||||
|
'secret_key': 'minioadmin',
|
||||||
|
'default_bucket': 'test-bucket',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
|
},
|
||||||
|
opc_config=_build_e2e_opc_config(opc_e2e_server.url),
|
||||||
|
pi_web_api_config={
|
||||||
|
'base_url': 'http://localhost:8080',
|
||||||
|
'auth_type': 'bearer',
|
||||||
|
'auth_token': 'test_token',
|
||||||
|
},
|
||||||
|
logger=mock_logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
|
mlflow_repository=mlflow_repository_stub,
|
||||||
|
)
|
||||||
|
activities.pi_web_api_client = pi_web_api_client_stub
|
||||||
|
await asyncio.to_thread(_connect_activities_to_opc_server, activities, opc_e2e_server.url)
|
||||||
|
try:
|
||||||
|
yield activities
|
||||||
|
finally:
|
||||||
|
await asyncio.to_thread(_teardown_real_opc_activities, activities)
|
||||||
|
|
||||||
|
|
||||||
|
def _teardown_real_opc_activities(activities: Activities) -> None:
|
||||||
|
"""Disconnect OPC sessions and shut down activities (sync, for asyncio.to_thread)."""
|
||||||
|
for opc_repo in activities.opc_repository.values():
|
||||||
|
opc_repo.disconnect()
|
||||||
|
activities.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope='function')
|
||||||
|
async def temporal_worker_real_opc(temporal_test_env, test_activities_real_opc):
|
||||||
|
"""Temporal worker using real OpcRepository against the in-process OPC UA server."""
|
||||||
|
with ThreadPoolExecutor(max_workers=32) as activity_executor:
|
||||||
|
async with Worker(
|
||||||
|
temporal_test_env.client,
|
||||||
|
task_queue='test-queue',
|
||||||
|
workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction],
|
||||||
|
activities=_worker_activity_list(test_activities_real_opc),
|
||||||
|
activity_executor=activity_executor,
|
||||||
|
) as worker:
|
||||||
|
yield worker
|
||||||
|
|
||||||
|
|
||||||
109
e2e/db_schema.sql
Normal file
109
e2e/db_schema.sql
Normal file
@@ -0,0 +1,109 @@
|
|||||||
|
-- =============================================================================
|
||||||
|
-- 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
|
||||||
|
);
|
||||||
368
e2e/helpers.py
Normal file
368
e2e/helpers.py
Normal file
@@ -0,0 +1,368 @@
|
|||||||
|
"""
|
||||||
|
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):
|
||||||
|
"""
|
||||||
|
Start a workflow and wait for its result.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
client: Temporal client from WorkflowEnvironment.
|
||||||
|
workflow_run: Workflow run method (e.g. PredictionsBatch.run).
|
||||||
|
input_data: Workflow input payload.
|
||||||
|
workflow_id: Unique workflow id.
|
||||||
|
timeout: Max seconds to wait for completion.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
Workflow result value.
|
||||||
|
"""
|
||||||
|
handle = await client.start_workflow(
|
||||||
|
workflow_run,
|
||||||
|
input_data,
|
||||||
|
id=workflow_id,
|
||||||
|
task_queue='test-queue',
|
||||||
|
)
|
||||||
|
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:
|
||||||
|
"""
|
||||||
|
Replace laborious_data rows for a model_id with one row per value (sensor_1..n).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
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}'))
|
||||||
|
values_sql = []
|
||||||
|
for i, value in enumerate(values):
|
||||||
|
values_sql.append(f"""
|
||||||
|
({model_id}, 'sensor_{i + 1}', {value}, '{data_timestamp}', '{data_timestamp}')
|
||||||
|
""")
|
||||||
|
insert_sql = f"""
|
||||||
|
INSERT INTO sientia_data.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_contains: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Assert exactly one prediction row exists for model_id with expected columns.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
postgres_engine: SQLAlchemy engine.
|
||||||
|
model_id: Expected model_id.
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
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'ORDER BY created_at ASC'
|
||||||
|
)
|
||||||
|
)
|
||||||
|
prediction_rows = result_query.fetchall()
|
||||||
|
assert len(prediction_rows) == 1, f'Expected one prediction record, got {len(prediction_rows)}'
|
||||||
|
row = prediction_rows[0]
|
||||||
|
assert row[0] == model_id, f'Expected model_id={model_id}, got {row[0]}'
|
||||||
|
assert row[1] == prediction or Decimal(str(row[1])) == Decimal(str(prediction)), (
|
||||||
|
f'Expected prediction={prediction}, got {row[1]}'
|
||||||
|
)
|
||||||
|
assert row[2] == prediction_confidence or Decimal(str(row[2])) == Decimal(
|
||||||
|
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_contains is not None:
|
||||||
|
assert comments_contains in actual_comments, (
|
||||||
|
f"Expected comments to contain '{comments_contains}', got '{actual_comments}'"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
assert actual_comments == comments, f"Expected comments='{comments}', got '{actual_comments}'"
|
||||||
|
|
||||||
|
|
||||||
|
def assert_continue(
|
||||||
|
postgres_engine: Engine,
|
||||||
|
model_id: int,
|
||||||
|
prediction_confidence: Decimal = Decimal(2),
|
||||||
|
comments: str = 'Input data with bad quality',
|
||||||
|
) -> None:
|
||||||
|
"""Assert one default-style prediction row after CONTINUE gate path."""
|
||||||
|
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}'
|
||||||
|
)
|
||||||
|
)
|
||||||
|
prediction_rows = result_query.fetchall()
|
||||||
|
assert len(prediction_rows) == 1, 'Expected one prediction record despite warnings'
|
||||||
|
row = prediction_rows[0]
|
||||||
|
assert row[1] == 0, f'Expected prediction=0, got {row[1]}'
|
||||||
|
assert row[2] == prediction_confidence, (
|
||||||
|
f'Expected prediction_confidence={prediction_confidence}, got {row[2]}'
|
||||||
|
)
|
||||||
|
assert row[3] == 'Bad', f"Expected prediction_status='Bad', got {row[3]}"
|
||||||
|
assert row[4] == comments, f"Expected comments='{comments}', got {row[4]}"
|
||||||
|
|
||||||
|
|
||||||
|
def assert_stop(postgres_engine: Engine, model_id: int) -> None:
|
||||||
|
"""Assert no prediction rows for model_id."""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
result_query = conn.execute(
|
||||||
|
text(f'SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = {model_id}')
|
||||||
|
)
|
||||||
|
count = result_query.scalar()
|
||||||
|
assert count == 0, f'Expected no predictions, but found {count} records'
|
||||||
|
|
||||||
|
|
||||||
|
def assert_repeat(postgres_engine: Engine, model_id: int, last_prediction: tuple) -> None:
|
||||||
|
"""
|
||||||
|
Assert two prediction rows for model_id both match last_prediction.
|
||||||
|
|
||||||
|
Rows are compared in created_at order for stability.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
postgres_engine: SQLAlchemy engine.
|
||||||
|
model_id: Model id.
|
||||||
|
last_prediction: Tuple (model_id, prediction, confidence, status) to match both rows.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
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'ORDER BY created_at ASC'
|
||||||
|
)
|
||||||
|
)
|
||||||
|
prediction_rows = result_query.fetchall()
|
||||||
|
assert len(prediction_rows) == 2, 'Expected two prediction records'
|
||||||
|
assert prediction_rows[0] == last_prediction, (
|
||||||
|
f'Expected first row {last_prediction}, got {prediction_rows[0]}'
|
||||||
|
)
|
||||||
|
assert prediction_rows[1] == last_prediction, (
|
||||||
|
f'Expected second row {last_prediction}, got {prediction_rows[1]}'
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
189
e2e/opc_test_server.py
Normal file
189
e2e/opc_test_server.py
Normal file
@@ -0,0 +1,189 @@
|
|||||||
|
"""
|
||||||
|
In-process OPC UA server for E2E tests (asyncua).
|
||||||
|
|
||||||
|
Provides writable prediction/confidence nodes and optional write faults
|
||||||
|
(Tier-1 BadSessionIdInvalid via PreWrite callback).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import socket
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from asyncua import Server, ua
|
||||||
|
from asyncua.common.callback import CallbackType
|
||||||
|
from asyncua.common.utils import ServiceError
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from asyncua.common.node import Node
|
||||||
|
|
||||||
|
|
||||||
|
UNKNOWN_NODE_ID = 'ns=99;i=9999'
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class OpcE2ENodeIds:
|
||||||
|
"""NodeId strings used in opc_output_config for E2E workflows."""
|
||||||
|
|
||||||
|
prediction: str
|
||||||
|
confidence: str
|
||||||
|
unknown: str = UNKNOWN_NODE_ID
|
||||||
|
|
||||||
|
|
||||||
|
class OpcE2ETestServer:
|
||||||
|
"""
|
||||||
|
Ephemeral asyncua server with Laborious E2E variables and controllable faults.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
host: Bind address (default 127.0.0.1).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, host: str = '127.0.0.1') -> None:
|
||||||
|
self._host = host
|
||||||
|
self._server: Server | None = None
|
||||||
|
self._prediction_node: Node | None = None
|
||||||
|
self._confidence_node: Node | None = None
|
||||||
|
self._session_bad_on_write = False
|
||||||
|
self._url: str | None = None
|
||||||
|
self._node_ids: OpcE2ENodeIds | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def url(self) -> str:
|
||||||
|
if self._url is None:
|
||||||
|
raise RuntimeError('OPC E2E server is not started')
|
||||||
|
return self._url
|
||||||
|
|
||||||
|
@property
|
||||||
|
def node_ids(self) -> OpcE2ENodeIds:
|
||||||
|
if self._node_ids is None:
|
||||||
|
raise RuntimeError('OPC E2E server is not started')
|
||||||
|
return self._node_ids
|
||||||
|
|
||||||
|
def set_session_bad_on_write(self, enabled: bool) -> None:
|
||||||
|
"""
|
||||||
|
When enabled, every client Write is rejected with BadSessionIdInvalid.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
enabled (bool): Turn Tier-1 session fault injection on or off.
|
||||||
|
"""
|
||||||
|
self._session_bad_on_write = enabled
|
||||||
|
|
||||||
|
async def start(self) -> OpcE2ENodeIds:
|
||||||
|
"""
|
||||||
|
Start the OPC UA server on a free TCP port.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
OpcE2ENodeIds: NodeId strings for prediction and confidence tags.
|
||||||
|
"""
|
||||||
|
port = _free_port(self._host)
|
||||||
|
self._url = f'opc.tcp://{self._host}:{port}/freeopcua/server/'
|
||||||
|
|
||||||
|
server = Server()
|
||||||
|
server.set_endpoint(self._url)
|
||||||
|
await server.init()
|
||||||
|
server.iserver.callback_service.addListener(
|
||||||
|
CallbackType.PreWrite,
|
||||||
|
self._pre_write_callback,
|
||||||
|
)
|
||||||
|
|
||||||
|
idx = await server.register_namespace('http://sientia.test/laborious-e2e')
|
||||||
|
e2e_object = await server.nodes.objects.add_object(idx, 'LaboriousE2E')
|
||||||
|
prediction = await e2e_object.add_variable(
|
||||||
|
idx,
|
||||||
|
'Prediction',
|
||||||
|
ua.Variant(0.0, ua.VariantType.Float),
|
||||||
|
)
|
||||||
|
confidence = await e2e_object.add_variable(
|
||||||
|
idx,
|
||||||
|
'Confidence',
|
||||||
|
ua.Variant(0.0, ua.VariantType.Float),
|
||||||
|
)
|
||||||
|
await prediction.set_writable()
|
||||||
|
await confidence.set_writable()
|
||||||
|
|
||||||
|
await server.start()
|
||||||
|
self._server = server
|
||||||
|
self._prediction_node = prediction
|
||||||
|
self._confidence_node = confidence
|
||||||
|
self._node_ids = OpcE2ENodeIds(
|
||||||
|
prediction=prediction.nodeid.to_string(),
|
||||||
|
confidence=confidence.nodeid.to_string(),
|
||||||
|
)
|
||||||
|
return self._node_ids
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Stop the OPC UA server and release the listening port."""
|
||||||
|
if self._server is not None:
|
||||||
|
await self._server.stop()
|
||||||
|
self._server = None
|
||||||
|
self._prediction_node = None
|
||||||
|
self._confidence_node = None
|
||||||
|
self._url = None
|
||||||
|
self._node_ids = None
|
||||||
|
self._session_bad_on_write = False
|
||||||
|
|
||||||
|
async def read_prediction(self) -> float:
|
||||||
|
"""
|
||||||
|
Read the current prediction variable value from the address space.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
float: Stored prediction value.
|
||||||
|
"""
|
||||||
|
if self._prediction_node is None:
|
||||||
|
raise RuntimeError('OPC E2E server is not started')
|
||||||
|
value = await self._prediction_node.read_value()
|
||||||
|
return float(value)
|
||||||
|
|
||||||
|
async def read_confidence(self) -> float:
|
||||||
|
"""
|
||||||
|
Read the current confidence variable value from the address space.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
float: Stored confidence value.
|
||||||
|
"""
|
||||||
|
if self._confidence_node is None:
|
||||||
|
raise RuntimeError('OPC E2E server is not started')
|
||||||
|
value = await self._confidence_node.read_value()
|
||||||
|
return float(value)
|
||||||
|
|
||||||
|
async def _pre_write_callback(self, _event, _service) -> None:
|
||||||
|
if self._session_bad_on_write:
|
||||||
|
raise ServiceError(ua.StatusCodes.BadSessionIdInvalid)
|
||||||
|
|
||||||
|
|
||||||
|
def _free_port(host: str) -> int:
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||||
|
sock.bind((host, 0))
|
||||||
|
return int(sock.getsockname()[1])
|
||||||
|
|
||||||
|
|
||||||
|
def build_opc_output_config(
|
||||||
|
node_ids: OpcE2ENodeIds,
|
||||||
|
*,
|
||||||
|
prediction_tag: str | None = None,
|
||||||
|
confidence_tag: str | None = None,
|
||||||
|
prediction_only: bool = False,
|
||||||
|
server_key: str = '1',
|
||||||
|
) -> dict[str, dict]:
|
||||||
|
"""
|
||||||
|
Build opc_output_config for PredictionsBatch using real server NodeIds.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node_ids (OpcE2ENodeIds): Node ids from OpcE2ETestServer.
|
||||||
|
prediction_tag (str | None): Override prediction NodeId (default: node_ids.prediction).
|
||||||
|
confidence_tag (str | None): Override confidence NodeId (default: node_ids.confidence).
|
||||||
|
prediction_only (bool): When True, omit confidence_tags (single write per activity).
|
||||||
|
server_key (str): OPC server id key in opc_output_config.
|
||||||
|
|
||||||
|
Return:
|
||||||
|
dict: opc_output_config payload for workflow input.
|
||||||
|
"""
|
||||||
|
pred = prediction_tag if prediction_tag is not None else node_ids.prediction
|
||||||
|
conf = confidence_tag if confidence_tag is not None else node_ids.confidence
|
||||||
|
server_config: dict = {
|
||||||
|
'prediction_tags': {pred: {'data_type': 'float'}},
|
||||||
|
}
|
||||||
|
if not prediction_only:
|
||||||
|
server_config['confidence_tags'] = {conf: {'data_type': 'float'}}
|
||||||
|
return {server_key: server_config}
|
||||||
14
e2e/scenario_inputs/drift_base.json
Normal file
14
e2e/scenario_inputs/drift_base.json
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
{
|
||||||
|
"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"
|
||||||
|
}
|
||||||
|
}
|
||||||
37
e2e/scenario_inputs/format_export_base.json
Normal file
37
e2e/scenario_inputs/format_export_base.json
Normal file
@@ -0,0 +1,37 @@
|
|||||||
|
{
|
||||||
|
"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"]
|
||||||
|
}
|
||||||
45
e2e/scenario_inputs/main_happy_path.json
Normal file
45
e2e/scenario_inputs/main_happy_path.json
Normal file
@@ -0,0 +1,45 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"metadata": {
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"workflow_name": "predictions_batch"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"query": "SELECT timestamp, variable, value, created_at FROM 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"]
|
||||||
|
}
|
||||||
45
e2e/scenario_inputs/main_invalid_datetime.json
Normal file
45
e2e/scenario_inputs/main_invalid_datetime.json
Normal file
@@ -0,0 +1,45 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"metadata": {
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"workflow_name": "predictions_batch"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"query": "SELECT timestamp, variable, value FROM 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"]
|
||||||
|
}
|
||||||
16
e2e/scenario_inputs/main_missing_required.json
Normal file
16
e2e/scenario_inputs/main_missing_required.json
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"metadata": {
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"workflow_name": "predictions_batch"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"schema": "sientia_data",
|
||||||
|
"table_name": "predictions",
|
||||||
|
"transform_table_name": "transformed_data"
|
||||||
|
}
|
||||||
44
e2e/scenario_inputs/main_sql_error.json
Normal file
44
e2e/scenario_inputs/main_sql_error.json
Normal file
@@ -0,0 +1,44 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"metadata": {
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"workflow_name": "predictions_batch"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"query": "SELECT * FROM nonexistent_table WHERE invalid_syntax =",
|
||||||
|
"schema": "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
|
||||||
|
}
|
||||||
|
}
|
||||||
12
e2e/scenario_inputs/minimal_retrain_base.json
Normal file
12
e2e/scenario_inputs/minimal_retrain_base.json
Normal file
@@ -0,0 +1,12 @@
|
|||||||
|
{
|
||||||
|
"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"
|
||||||
|
}
|
||||||
|
}
|
||||||
11
e2e/scenario_inputs/minio_offload_load_query.json
Normal file
11
e2e/scenario_inputs/minio_offload_load_query.json
Normal file
@@ -0,0 +1,11 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"schedule_name": "test-schedule",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"model_id": "{{MODEL_ID}}",
|
||||||
|
"workflow_name": "predictions_batch"
|
||||||
|
},
|
||||||
|
"query": "SELECT timestamp, variable, value, created_at FROM sientia_data.laborious_data WHERE model_id = {{MODEL_ID}}",
|
||||||
|
"model_name": "test_model",
|
||||||
|
"datetime_columns": ["timestamp", "created_at"]
|
||||||
|
}
|
||||||
37
e2e/scenario_inputs/minio_offload_workflow.json
Normal file
37
e2e/scenario_inputs/minio_offload_workflow.json
Normal file
@@ -0,0 +1,37 @@
|
|||||||
|
{
|
||||||
|
"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"]
|
||||||
|
}
|
||||||
39
e2e/scenario_inputs/prediction_process_base.json
Normal file
39
e2e/scenario_inputs/prediction_process_base.json
Normal file
@@ -0,0 +1,39 @@
|
|||||||
|
{
|
||||||
|
"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"]
|
||||||
|
}
|
||||||
14
e2e/scenario_inputs/simple_metrics_base.json
Normal file
14
e2e/scenario_inputs/simple_metrics_base.json
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
{
|
||||||
|
"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"
|
||||||
|
}
|
||||||
|
}
|
||||||
455
e2e/scenarios.md
Normal file
455
e2e/scenarios.md
Normal file
@@ -0,0 +1,455 @@
|
|||||||
|
# E2E Scenario Documentation - Predictions Batch
|
||||||
|
|
||||||
|
This document describes the end-to-end scenarios for `predictions_batch` and its child workflows:
|
||||||
|
`prediction_process` and `format_and_export_prediction`.
|
||||||
|
|
||||||
|
It is a functional reference of scenario behavior, inputs, and expected outcomes.
|
||||||
|
|
||||||
|
## Execution Context
|
||||||
|
|
||||||
|
- 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`).
|
||||||
|
|
||||||
|
### 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
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Main Workflow Scenarios
|
||||||
|
Source: `e2e/test_predictions_batch_main_workflow.py`
|
||||||
|
|
||||||
|
### 1.1.1 Happy Path - Complete Success
|
||||||
|
**Summary**: Full workflow succeeds with valid query and default gate behavior.
|
||||||
|
|
||||||
|
**Description**:
|
||||||
|
- Query returns rows for a model.
|
||||||
|
- `prediction_process` runs transform and predict paths.
|
||||||
|
- Final prediction and transformed data are persisted.
|
||||||
|
|
||||||
|
**Expected Outcome**:
|
||||||
|
- Exactly one prediction row is created.
|
||||||
|
- Transform rows are created.
|
||||||
|
- Confidence/status/comments are success values.
|
||||||
|
|
||||||
|
### 1.2.1 SQL Query Execution Error
|
||||||
|
**Summary**: Invalid SQL leads to no persisted prediction.
|
||||||
|
|
||||||
|
**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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Prediction Process Scenarios
|
||||||
|
Source: `e2e/test_predictions_batch_prediction_process.py`
|
||||||
|
|
||||||
|
### 2.1 Input Gate Path Decisions
|
||||||
|
|
||||||
|
#### 2.1.1 CONTINUE
|
||||||
|
**Summary**: Input filter flags quality issue but allows continuation via default path.
|
||||||
|
|
||||||
|
**Description**:
|
||||||
|
- Input gate returns `CONTINUE`.
|
||||||
|
- MLFlow transform/predict are skipped.
|
||||||
|
- Export path persists default-style prediction with warning context.
|
||||||
|
|
||||||
|
#### 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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Format and Export Scenarios
|
||||||
|
Source: `e2e/test_predictions_batch_format_export.py`
|
||||||
|
|
||||||
|
### 3.1 Output Combination Scenarios
|
||||||
|
|
||||||
|
#### 3.1.1 Default prediction export
|
||||||
|
**Summary**: Non-`None` path flag uses `format_default_prediction`.
|
||||||
|
|
||||||
|
**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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. MinIO Offload Scenarios
|
||||||
|
Source: `e2e/test_minio_offload.py`
|
||||||
|
|
||||||
|
### 4.1.1 Forced offload to MinIO
|
||||||
|
**Summary**: Very low threshold forces parquet upload.
|
||||||
|
|
||||||
|
**Description**:
|
||||||
|
- Payload is offloaded (`object_key` present, inline data absent/empty).
|
||||||
|
- Object is present in MinIO under `prediction_datasets/...`.
|
||||||
|
- Retrieval reconstructs the dataframe.
|
||||||
|
|
||||||
|
### 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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Drift Workflow Scenarios
|
||||||
|
Source: `e2e/test_drift.py`
|
||||||
|
|
||||||
|
The drift suite drives the **real** `sientia_model.analytics.drift_analysis.DriftAnalysis`
|
||||||
|
analyzer (no stubs / mocks). 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)
|
||||||
|
```
|
||||||
|
|
||||||
|
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:
|
||||||
|
|
||||||
|
`id, model_id, feature, method, value, alert, chunk_index, chunk_start_date, chunk_end_date, accurate, timestamp, created_at`.
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
### 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.
|
||||||
|
|
||||||
|
**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.
|
||||||
|
|
||||||
|
#### 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 Outcome**:
|
||||||
|
- All persisted rows carry `accurate=False`.
|
||||||
|
- A `MODEL_METRICS_REFERENCE_DATA_WARNING` notification is emitted to MongoDB.
|
||||||
|
|
||||||
|
### 5.2 Failure paths
|
||||||
|
|
||||||
|
#### D.3.1 Empty target data short-circuits the workflow
|
||||||
|
**Summary**: `load_custom_query` returns no rows.
|
||||||
|
|
||||||
|
**Expected Outcome**:
|
||||||
|
- The workflow returns early and writes nothing to `sientia_data.drift_metrics`.
|
||||||
|
|
||||||
|
### 5.3 Configuration paths
|
||||||
|
|
||||||
|
#### 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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Simple Metrics Workflow Scenarios
|
||||||
|
Source: `e2e/test_simple_metrics.py`
|
||||||
|
|
||||||
|
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``.
|
||||||
|
|
||||||
|
### 6.1 Happy paths
|
||||||
|
|
||||||
|
#### 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 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.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. Minimal Retrain Workflow Scenarios
|
||||||
|
Source: `e2e/test_minimal_retrain.py`
|
||||||
|
|
||||||
|
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``.
|
||||||
|
|
||||||
|
### 7.1 Happy path
|
||||||
|
|
||||||
|
#### 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).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Input Contract Reference
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
Optional outputs:
|
||||||
|
- `opc_output_config`
|
||||||
|
- `pi_web_api_output_config`
|
||||||
74
e2e/test_child_workflows_e2e.py
Normal file
74
e2e/test_child_workflows_e2e.py
Normal file
@@ -0,0 +1,74 @@
|
|||||||
|
"""
|
||||||
|
Direct E2E execution of child workflows (smaller surface than PredictionsBatch).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import text
|
||||||
|
from temporalio.testing import WorkflowEnvironment
|
||||||
|
from temporalio.worker import Worker
|
||||||
|
|
||||||
|
from e2e.helpers import make_workflow_id, start_and_await_workflow
|
||||||
|
from laborious.workflows.sub_workflows.format_and_export_prediction import FormatAndExportPrediction
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_format_and_export_prediction_default_path_e2e(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Run FormatAndExportPrediction with path_flag set (format_default_prediction path).
|
||||||
|
"""
|
||||||
|
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}'))
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
'metadata': {
|
||||||
|
'model_id': model_id,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'subworkflow.format_and_export_prediction',
|
||||||
|
}
|
||||||
|
}
|
||||||
|
input_data = {
|
||||||
|
'metadata': metadata,
|
||||||
|
'path_flag': 'CONTINUE',
|
||||||
|
'data': {'last_timestamp': '2024-01-01 12:00:00+00:00'},
|
||||||
|
'prediction_confidence': 2,
|
||||||
|
'timestamp': '2024-01-01 12:00:00+00:00',
|
||||||
|
'model_id': model_id,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schema': 'sientia_data',
|
||||||
|
'table_name': 'predictions',
|
||||||
|
'transform_table_name': 'transformed_data',
|
||||||
|
'comment': 'e2e child workflow default path',
|
||||||
|
'opc_output_config': {},
|
||||||
|
'pi_web_api_output_config': {},
|
||||||
|
'prediction_store_policy': 'lts:1',
|
||||||
|
}
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
FormatAndExportPrediction.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('e2e-format-export-child'),
|
||||||
|
)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
text(
|
||||||
|
f'SELECT prediction, prediction_confidence, prediction_status, comments '
|
||||||
|
f'FROM sientia_data.predictions WHERE model_id = {model_id}'
|
||||||
|
)
|
||||||
|
).fetchone()
|
||||||
|
assert row is not None
|
||||||
|
assert row[0] == 0
|
||||||
|
assert row[1] == Decimal(2)
|
||||||
|
assert row[2] == 'Bad'
|
||||||
|
assert row[3] == 'e2e child workflow default path'
|
||||||
600
e2e/test_drift.py
Normal file
600
e2e/test_drift.py
Normal file
@@ -0,0 +1,600 @@
|
|||||||
|
"""
|
||||||
|
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}'
|
||||||
|
)
|
||||||
430
e2e/test_minimal_retrain.py
Normal file
430
e2e/test_minimal_retrain.py
Normal file
@@ -0,0 +1,430 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
108
e2e/test_minio_offload.py
Normal file
108
e2e/test_minio_offload.py
Normal file
@@ -0,0 +1,108 @@
|
|||||||
|
"""
|
||||||
|
E2E-style tests for MinIO offload using a real MinIO testcontainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
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 laborious.activities.activities import Activities
|
||||||
|
from laborious.utils.models import minio_dataframe_payload as mdp
|
||||||
|
from laborious.workflows.predictions_batch import PredictionsBatch
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_load_query_with_minio_offload_writes_object_to_bucket(
|
||||||
|
postgres_engine,
|
||||||
|
minio_container,
|
||||||
|
test_activities_real_minio: Activities,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
With a tiny offload threshold, query results are uploaded as Parquet to MinIO.
|
||||||
|
|
||||||
|
Uses real MinioRepository against testcontainers MinIO (no MinIO mock).
|
||||||
|
"""
|
||||||
|
model_id = 501
|
||||||
|
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)
|
||||||
|
metadata = {'metadata': scenario_input['metadata']}
|
||||||
|
with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1):
|
||||||
|
payload = test_activities_real_minio.load_query_with_minio_offload(scenario_input)
|
||||||
|
assert payload.object_key, 'offloaded payload must reference a MinIO object'
|
||||||
|
assert payload.data is None or payload.data == {}, 'large payloads should not inline tabular dict'
|
||||||
|
|
||||||
|
df = payload.retrieve(test_activities_real_minio.minio_repository, metadata['metadata'])
|
||||||
|
assert len(df) >= 1
|
||||||
|
|
||||||
|
client = minio_container.get_client()
|
||||||
|
listed = list(client.list_objects('test-bucket', recursive=True))
|
||||||
|
names = [getattr(o, 'object_name', None) or getattr(o, '_object_name', '') for o in listed]
|
||||||
|
assert any(n and 'prediction_datasets' in n for n in names), f'unexpected object listing: {names!r}'
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_predictions_batch_with_minio_offload_path(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker_real_minio: Worker,
|
||||||
|
postgres_engine,
|
||||||
|
test_activities_real_minio: Activities,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Full PredictionsBatch run with offload: load step stores Parquet in MinIO; pipeline completes.
|
||||||
|
"""
|
||||||
|
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}'))
|
||||||
|
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)
|
||||||
|
|
||||||
|
with patch.object(mdp, 'OFFLOAD_THRESHOLD_BYTES', 1):
|
||||||
|
await start_and_await_workflow(
|
||||||
|
temporal_test_env.client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-batch-minio-offload'),
|
||||||
|
)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
count = conn.execute(
|
||||||
|
text(f'SELECT COUNT(*) FROM sientia_data.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
|
||||||
201
e2e/test_opc_real_server.py
Normal file
201
e2e/test_opc_real_server.py
Normal file
@@ -0,0 +1,201 @@
|
|||||||
|
"""
|
||||||
|
E2E tests for OPC export using an in-process asyncua server and real OpcRepository.
|
||||||
|
|
||||||
|
Covers scenarios 3.1.2, 3.2.2, 3.2.4, and 3.2.5 from e2e/scenarios.md.
|
||||||
|
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
|
||||||
|
from temporalio.worker import Worker
|
||||||
|
|
||||||
|
from e2e.helpers import assert_prediction, insert_sample_data, make_workflow_id, start_and_await_workflow
|
||||||
|
from e2e.opc_test_server import UNKNOWN_NODE_ID, OpcE2ETestServer, build_opc_output_config
|
||||||
|
from e2e.test_predictions_batch_format_export import get_base_input_data
|
||||||
|
from laborious.activities.activities import Activities
|
||||||
|
from laborious.activities.opc import OPC_RECONNECT_IN_PROGRESS_COMMENT
|
||||||
|
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:
|
||||||
|
"""
|
||||||
|
Hold the connection lock briefly so concurrent writes see reconnect_in_progress.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.opc
|
||||||
|
async def test_scenario_3_1_2_export_with_opc_only_real_server(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker_real_opc: Worker,
|
||||||
|
test_activities_real_opc: Activities,
|
||||||
|
opc_e2e_server: OpcE2ETestServer,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.2 (real OPC): connect, write prediction and confidence, verify server values.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 412
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['opc_output_config'] = build_opc_output_config(opc_e2e_server.node_ids)
|
||||||
|
input_data['pi_web_api_output_config'] = None
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-opc-real-happy'),
|
||||||
|
)
|
||||||
|
|
||||||
|
test_activities_real_opc.pi_web_api_client.write_value.assert_not_called()
|
||||||
|
assert await opc_e2e_server.read_prediction() == pytest.approx(0.5)
|
||||||
|
assert await opc_e2e_server.read_confidence() == pytest.approx(0.0)
|
||||||
|
assert_prediction(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.opc
|
||||||
|
async def test_scenario_3_2_2_opc_write_error_real_server(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker_real_opc: Worker,
|
||||||
|
test_activities_real_opc: Activities,
|
||||||
|
opc_e2e_server: OpcE2ETestServer,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.2 (real OPC): unknown NodeId yields generic write failure (confidence 12).
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 422
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
node_ids = opc_e2e_server.node_ids
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['opc_output_config'] = build_opc_output_config(
|
||||||
|
node_ids,
|
||||||
|
prediction_tag=UNKNOWN_NODE_ID,
|
||||||
|
confidence_tag=UNKNOWN_NODE_ID,
|
||||||
|
)
|
||||||
|
input_data['pi_web_api_output_config'] = None
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-opc-real-bad-node'),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine,
|
||||||
|
model_id,
|
||||||
|
prediction_confidence=12,
|
||||||
|
comments='Some data could not be written to OPC servers',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.opc
|
||||||
|
async def test_scenario_3_2_4_opc_session_bad_real_server(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker_real_opc: Worker,
|
||||||
|
test_activities_real_opc: Activities,
|
||||||
|
opc_e2e_server: OpcE2ETestServer,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.4 (real OPC): server PreWrite fault injects BadSessionIdInvalid (confidence 14).
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 424
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
opc_e2e_server.set_session_bad_on_write(True)
|
||||||
|
try:
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['opc_output_config'] = build_opc_output_config(
|
||||||
|
opc_e2e_server.node_ids,
|
||||||
|
prediction_only=True,
|
||||||
|
)
|
||||||
|
input_data['pi_web_api_output_config'] = None
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-opc-real-session-bad'),
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
opc_e2e_server.set_session_bad_on_write(False)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine,
|
||||||
|
model_id,
|
||||||
|
prediction_confidence=14,
|
||||||
|
comments_contains='OPC UA session/channel error: BadSessionIdInvalid',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.opc
|
||||||
|
async def test_scenario_3_2_5_opc_write_blocked_during_reconnect_real_server(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker_real_opc: Worker,
|
||||||
|
test_activities_real_opc: Activities,
|
||||||
|
opc_e2e_server: OpcE2ETestServer,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.5 (real OPC): writes rejected while reconnect holds the connection lock.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 425
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['opc_output_config'] = build_opc_output_config(opc_e2e_server.node_ids)
|
||||||
|
input_data['pi_web_api_output_config'] = None
|
||||||
|
|
||||||
|
try:
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-opc-real-reconnect-block'),
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
reconnect_thread.join(timeout=5.0)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine,
|
||||||
|
model_id,
|
||||||
|
prediction_confidence=14,
|
||||||
|
comments_contains=OPC_RECONNECT_IN_PROGRESS_COMMENT,
|
||||||
|
)
|
||||||
836
e2e/test_predictions_batch_format_export.py
Normal file
836
e2e/test_predictions_batch_format_export.py
Normal file
@@ -0,0 +1,836 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
|
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 laborious.activities.activities import Activities
|
||||||
|
from laborious.workflows.predictions_batch import PredictionsBatch
|
||||||
|
|
||||||
|
def get_base_input_data(model_id):
|
||||||
|
return load_scenario_input('format_export_base.json', model_id=model_id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_1_1_default_prediction_export(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.1: Default prediction export (non-None path_flag).
|
||||||
|
|
||||||
|
Triggers input_gate CONTINUE via SPECIFIC_VARIABLES_NULL_VALUES so
|
||||||
|
PredictionProcess calls FormatAndExportPrediction with path_flag set.
|
||||||
|
That workflow uses format_default_prediction (not format_prediction) and
|
||||||
|
skips format_transformed_data / transform Postgres export.
|
||||||
|
|
||||||
|
Optional PI Web API and OPC outputs still run when configured.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
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}"))
|
||||||
|
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
|
||||||
|
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['input_filters'] = {
|
||||||
|
'SPECIFIC_VARIABLES_NULL_VALUES': {
|
||||||
|
'POLICY': 'CONTINUE',
|
||||||
|
'CONFIG': {'variables': ['sensor_1']},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
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',
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
wid = make_workflow_id('test-default-prediction')
|
||||||
|
|
||||||
|
await start_and_await_workflow(client, PredictionsBatch.run, input_data, wid)
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_1'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 0,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 311,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_2'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 2,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 311,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
],
|
||||||
|
any_order=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
|
||||||
|
opc_write_data.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
'addr_1',
|
||||||
|
0,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 311,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
'addr_2',
|
||||||
|
2,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 311,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
tf_count = conn.execute(
|
||||||
|
text(f"SELECT COUNT(*) FROM sientia_data.transformed_data WHERE model_id = {model_id}")
|
||||||
|
).scalar()
|
||||||
|
assert tf_count == 0, 'transform export must be skipped when path_flag is set'
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine,
|
||||||
|
model_id,
|
||||||
|
prediction=0,
|
||||||
|
prediction_confidence=Decimal(2),
|
||||||
|
prediction_status='Bad',
|
||||||
|
comments='Input data with bad quality',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_1_2_export_with_opc_only(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.2: Export with OPC only
|
||||||
|
|
||||||
|
Description:
|
||||||
|
Export to PostgreSQL and OPC server only (no PI Web API).
|
||||||
|
|
||||||
|
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
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 312
|
||||||
|
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
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 # No PI Web API config
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-opc-only')
|
||||||
|
)
|
||||||
|
|
||||||
|
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
|
||||||
|
opc_write_data.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
'addr_1',
|
||||||
|
0.5,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 312,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
'addr_2',
|
||||||
|
0,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 312,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.assert_not_called()
|
||||||
|
|
||||||
|
assert_prediction(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_1_3_export_with_pi_web_api_only(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.3: Export with PI Web API only
|
||||||
|
|
||||||
|
Description:
|
||||||
|
Export to PostgreSQL and PI Web API only (no OPC).
|
||||||
|
|
||||||
|
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
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 313
|
||||||
|
|
||||||
|
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'] = None # No OPC config
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-pi-api-only')
|
||||||
|
)
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_1'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 0.5,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 313,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_2'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 0,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 313,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
],
|
||||||
|
any_order=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
|
||||||
|
opc_write_data.assert_not_called()
|
||||||
|
|
||||||
|
assert_prediction(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_1_4_export_without_optional_outputs(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.4: Export Without Optional Outputs
|
||||||
|
|
||||||
|
Description:
|
||||||
|
Export only to PostgreSQL (no OPC or PI Web API).
|
||||||
|
|
||||||
|
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
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 314
|
||||||
|
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['opc_output_config'] = None # No OPC config
|
||||||
|
input_data['pi_web_api_output_config'] = None # No PI Web API config
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-no-optional-outputs')
|
||||||
|
)
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.assert_not_called()
|
||||||
|
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
|
||||||
|
opc_write_data.assert_not_called()
|
||||||
|
|
||||||
|
assert_prediction(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_1_5_export_without_transformed_data(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.1.5: Export Without Transformed Data
|
||||||
|
|
||||||
|
Description:
|
||||||
|
Only prediction exported, no transform table.
|
||||||
|
|
||||||
|
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
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 315
|
||||||
|
|
||||||
|
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}"))
|
||||||
|
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['save_transform'] = False # Don't save transformed data
|
||||||
|
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-no-transform-export')
|
||||||
|
)
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_1'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 0.5,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 315,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_2'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01 12:00:00+0000',
|
||||||
|
'Value': 0,
|
||||||
|
},
|
||||||
|
metadata={
|
||||||
|
'model_id': 315,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
],
|
||||||
|
any_order=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
opc_write_data = cast(Any, test_activities.opc_repository['1'].write_data)
|
||||||
|
opc_write_data.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
'addr_1',
|
||||||
|
0.5,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 315,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
'addr_2',
|
||||||
|
0,
|
||||||
|
'float',
|
||||||
|
{
|
||||||
|
'model_id': 315,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'schedule_name': 'test-schedule',
|
||||||
|
'workflow_name': 'predictions_batch',
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
result_query = conn.execute(
|
||||||
|
text(f"SELECT COUNT(*) FROM sientia_data.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"
|
||||||
|
|
||||||
|
assert_prediction(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_2_1_pi_web_api_write_error(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
notification_inserts,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.1: PI Web API Write Error
|
||||||
|
|
||||||
|
Export failure is handled inside the activity; there is no retry loop. The
|
||||||
|
workflow completes and PostgreSQL stores prediction_confidence 13 and the
|
||||||
|
error message in comments.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 321
|
||||||
|
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.write_value.side_effect = Exception(
|
||||||
|
"PI Web API service unavailable")
|
||||||
|
|
||||||
|
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-api-error')
|
||||||
|
)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine, model_id,
|
||||||
|
prediction_confidence=13,
|
||||||
|
comments='PI Web API service unavailable',
|
||||||
|
)
|
||||||
|
assert notification_inserts.call_count >= 1
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_2_2_opc_write_error(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.2: OPC Write Error
|
||||||
|
|
||||||
|
OPC failure is reported without failing the workflow; there is no retry
|
||||||
|
loop. PostgreSQL stores prediction_confidence 12 and OPC error comments.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 322
|
||||||
|
|
||||||
|
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, {
|
||||||
|
'notification_id': 'OPC_WRITE_DATA_ERROR_1',
|
||||||
|
'message': 'OPC server unavailable',
|
||||||
|
'block': 'opc_repository',
|
||||||
|
'level': NotificationLevel.ERROR,
|
||||||
|
'attachment_content': 'OPC server unavailable',
|
||||||
|
})
|
||||||
|
|
||||||
|
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'] = {
|
||||||
|
'endpoint': 'test_endpoint',
|
||||||
|
'prediction_tags': {'tag_1': 'web_id_1'},
|
||||||
|
'confidence_tags': {'tag_2': 'web_id_2'},
|
||||||
|
}
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-opc-error')
|
||||||
|
)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine, model_id,
|
||||||
|
prediction_confidence=12,
|
||||||
|
comments='Some data could not be written to OPC servers',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_2_4_opc_session_bad_mock(
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 324
|
||||||
|
|
||||||
|
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': 'session_bad',
|
||||||
|
'opc_status': 'BadSessionIdInvalid',
|
||||||
|
'notification_id': 'OPC_WRITE_DATA_ERROR_1',
|
||||||
|
'message': 'OPC session invalid',
|
||||||
|
'block': 'opc_repository',
|
||||||
|
'level': NotificationLevel.ERROR,
|
||||||
|
'attachment_content': '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'}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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')
|
||||||
|
)
|
||||||
|
|
||||||
|
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',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_3_2_3_pi_web_api_partial_write_error(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Scenario 3.2.3: PI Web API Partial Write Error
|
||||||
|
|
||||||
|
Partial PI write: confidence 13, descriptive comments, workflow completes
|
||||||
|
without an activity retry loop.
|
||||||
|
"""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
model_id = 323
|
||||||
|
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
|
||||||
|
test_activities.pi_web_api_client.set_side_effect(
|
||||||
|
[
|
||||||
|
# Prediction batch: two web_ids requested, only one acknowledged.
|
||||||
|
[{'WebId': 'web_id_1', 'Errors': []}],
|
||||||
|
# Confidence write succeeds.
|
||||||
|
[{'WebId': 'web_id_2', 'Errors': []}],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
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', 'tag_3': 'web_id_3'},
|
||||||
|
'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-api-partial-error')
|
||||||
|
)
|
||||||
|
|
||||||
|
assert_prediction(
|
||||||
|
postgres_engine, model_id,
|
||||||
|
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)
|
||||||
|
|
||||||
181
e2e/test_predictions_batch_main_workflow.py
Normal file
181
e2e/test_predictions_batch_main_workflow.py
Normal file
@@ -0,0 +1,181 @@
|
|||||||
|
"""
|
||||||
|
End-to-end tests for PredictionsBatch workflow - Main workflow scenarios.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from sqlalchemy import text
|
||||||
|
from temporalio.testing import WorkflowEnvironment
|
||||||
|
from temporalio.worker import Worker
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from e2e.helpers import load_scenario_input, make_workflow_id, start_and_await_workflow
|
||||||
|
from laborious.activities.activities import Activities
|
||||||
|
from laborious.workflows.predictions_batch import PredictionsBatch
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_1_1_1_happy_path_complete_success(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""Scenario 1.1.1: Happy path with SQL load, MLflow mocks, Postgres predictions and transforms."""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
|
||||||
|
with postgres_engine.begin() as conn:
|
||||||
|
conn.execute(text('DELETE FROM sientia_data.laborious_data WHERE model_id = 123'))
|
||||||
|
insert_sql = """
|
||||||
|
INSERT INTO sientia_data.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'),
|
||||||
|
(123, 'sensor_3', 120.8, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00')
|
||||||
|
"""
|
||||||
|
conn.execute(text(insert_sql))
|
||||||
|
|
||||||
|
input_data = load_scenario_input('main_happy_path.json', model_id=123)
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-predictions-batch'),
|
||||||
|
)
|
||||||
|
|
||||||
|
schema_name = 'sientia_data'
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
result_query = conn.execute(
|
||||||
|
text(
|
||||||
|
f'SELECT model_id, prediction, prediction_confidence, response_time, prediction_status, comments '
|
||||||
|
f'FROM {schema_name}.predictions WHERE model_id = 123'
|
||||||
|
)
|
||||||
|
)
|
||||||
|
prediction_rows = result_query.fetchall()
|
||||||
|
assert len(prediction_rows) == 1
|
||||||
|
row = prediction_rows[0]
|
||||||
|
assert row[0] == 123
|
||||||
|
assert row[1] == 0.5
|
||||||
|
assert row[2] == 0, f'Expected prediction_confidence=0, got {row[2]}'
|
||||||
|
assert row[3] is not None
|
||||||
|
assert row[4] == 'Good'
|
||||||
|
assert row[5] == ''
|
||||||
|
|
||||||
|
result_query = conn.execute(
|
||||||
|
text(
|
||||||
|
f'SELECT model_id, variable, value FROM {schema_name}.transformed_data WHERE model_id = 123'
|
||||||
|
)
|
||||||
|
)
|
||||||
|
transformed_rows = result_query.fetchall()
|
||||||
|
assert len(transformed_rows) == 2
|
||||||
|
assert transformed_rows[0][0] == 123
|
||||||
|
assert transformed_rows[0][1] == 'feature_1'
|
||||||
|
assert float(transformed_rows[0][2]) == 0.234
|
||||||
|
assert transformed_rows[1][0] == 123
|
||||||
|
assert transformed_rows[1][1] == 'feature_2'
|
||||||
|
assert float(transformed_rows[1][2]) == 0.783
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_1_2_1_sql_query_execution_error(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""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)
|
||||||
|
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client,
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
make_workflow_id('test-sql-error'),
|
||||||
|
)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
count = conn.execute(
|
||||||
|
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 128')
|
||||||
|
).scalar()
|
||||||
|
assert count == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_1_2_2_missing_required_parameters(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""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)
|
||||||
|
|
||||||
|
handle = await client.start_workflow(
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
id=make_workflow_id('test-missing-param'),
|
||||||
|
task_queue='test-queue',
|
||||||
|
)
|
||||||
|
|
||||||
|
# Let Temporal process a few workflow tasks; for this case, result() can hang.
|
||||||
|
await asyncio.sleep(2.0)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
count = conn.execute(
|
||||||
|
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 129')
|
||||||
|
).scalar()
|
||||||
|
assert count == 0
|
||||||
|
|
||||||
|
await handle.terminate('expected failure path in e2e test (missing required parameters)')
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_1_2_3_invalid_datetime_column_specification(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""Invalid datetime column: no predictions persisted; workflow terminated after validation."""
|
||||||
|
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(
|
||||||
|
"""
|
||||||
|
INSERT INTO sientia_data.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)
|
||||||
|
|
||||||
|
handle = await client.start_workflow(
|
||||||
|
PredictionsBatch.run,
|
||||||
|
input_data,
|
||||||
|
id=make_workflow_id('test-invalid-datetime-col'),
|
||||||
|
task_queue='test-queue',
|
||||||
|
)
|
||||||
|
|
||||||
|
# Let Temporal process and surface the failure path internally.
|
||||||
|
await asyncio.sleep(2.0)
|
||||||
|
|
||||||
|
with postgres_engine.connect() as conn:
|
||||||
|
count = conn.execute(
|
||||||
|
text('SELECT COUNT(*) FROM sientia_data.predictions WHERE model_id = 130')
|
||||||
|
).scalar()
|
||||||
|
assert count == 0
|
||||||
|
|
||||||
|
await handle.terminate('expected failure path in e2e test (invalid datetime column)')
|
||||||
488
e2e/test_predictions_batch_prediction_process.py
Normal file
488
e2e/test_predictions_batch_prediction_process.py
Normal file
@@ -0,0 +1,488 @@
|
|||||||
|
"""
|
||||||
|
End-to-end tests for PredictionsBatch workflow - Prediction Process scenarios.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from decimal import Decimal
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import text
|
||||||
|
from temporalio.testing import WorkflowEnvironment
|
||||||
|
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'
|
||||||
|
|
||||||
|
|
||||||
|
def get_base_input_data(model_id):
|
||||||
|
return load_scenario_input('prediction_process_base.json', model_id=model_id)
|
||||||
|
|
||||||
|
|
||||||
|
@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
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def bad_predict_model(mlflow_repository_stub):
|
||||||
|
wrapper = mlflow_repository_stub.stub_wrapper
|
||||||
|
|
||||||
|
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, {}
|
||||||
|
|
||||||
|
wrapper.transform.side_effect = _good_transform
|
||||||
|
wrapper.predict = MagicMock(side_effect=Exception('Bad predict model'))
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_1_1_input_gate_triggers_continue(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
):
|
||||||
|
"""Input gate CONTINUE: export default prediction; MLflow transform/predict not used."""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 211
|
||||||
|
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
await start_and_await_workflow(
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_1_2_input_gate_triggers_stop(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
):
|
||||||
|
"""Input gate STOP: no export, no MLflow."""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 212
|
||||||
|
insert_sample_data(postgres_engine, model_id, ['NULL', 78.2])
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['input_filters']['SPECIFIC_VARIABLES_NULL_VALUES']['POLICY'] = 'STOP'
|
||||||
|
await start_and_await_workflow(
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_1_3_input_gate_repeat_batch_timestamp_equals_history_fails(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
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)
|
||||||
|
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')
|
||||||
|
)
|
||||||
|
assert_repeat(postgres_engine, model_id, data)
|
||||||
|
mlflow_repository_stub.stub_wrapper.transform.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_1_4_input_gate_repeat_without_prior_prediction(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""REPEAT when no prior row in predictions: repeat_last_prediction runs; still no new duplicate export path."""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
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}'))
|
||||||
|
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-no-history')
|
||||||
|
)
|
||||||
|
assert_stop(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_2_1_transform_gate_triggers_continue(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_data_model,
|
||||||
|
):
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 221
|
||||||
|
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'
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-transform-continue')
|
||||||
|
)
|
||||||
|
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_scenario_2_2_2_transform_gate_triggers_stop(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_data_model,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
):
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 222
|
||||||
|
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'] = 'STOP'
|
||||||
|
await start_and_await_workflow(
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_2_3_transform_gate_repeat_batch_timestamp_equals_history_fails(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_data_model,
|
||||||
|
):
|
||||||
|
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)
|
||||||
|
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')
|
||||||
|
)
|
||||||
|
assert_repeat(postgres_engine, model_id, data)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_2_4_transform_content_gate_nan_values_stop(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
mlflow_repository_stub,
|
||||||
|
):
|
||||||
|
"""mlflow_content_gate triggers STOP when transform output is all NaN (NAN_VALUES filter)."""
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 224
|
||||||
|
|
||||||
|
def all_nan_transform(data):
|
||||||
|
result = pd.DataFrame(
|
||||||
|
{'feature_1': [np.nan] * len(data), 'feature_2': [np.nan] * len(data)}
|
||||||
|
)
|
||||||
|
result.index = data.index
|
||||||
|
return result, {}
|
||||||
|
|
||||||
|
mlflow_repository_stub.stub_wrapper.transform = MagicMock(side_effect=all_nan_transform)
|
||||||
|
|
||||||
|
insert_sample_data(postgres_engine, model_id, [60.0, 78.2])
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['mlflow_transform_filters'] = {
|
||||||
|
'API_ERROR': {'POLICY': 'STOP', 'CONFIG': {}},
|
||||||
|
'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}},
|
||||||
|
}
|
||||||
|
await start_and_await_workflow(
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_3_1_predict_gate_triggers_continue(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_predict_model,
|
||||||
|
):
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 231
|
||||||
|
insert_sample_data(postgres_engine, model_id, [60.0, 78.2])
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'CONTINUE'
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-predict-continue')
|
||||||
|
)
|
||||||
|
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_scenario_2_3_2_predict_gate_triggers_stop(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_predict_model,
|
||||||
|
):
|
||||||
|
client = temporal_test_env.client
|
||||||
|
model_id = 232
|
||||||
|
insert_sample_data(postgres_engine, model_id, [23.5, 78.2])
|
||||||
|
input_data = get_base_input_data(model_id)
|
||||||
|
input_data['mlflow_predict_filters']['API_ERROR']['POLICY'] = 'STOP'
|
||||||
|
await start_and_await_workflow(
|
||||||
|
client, PredictionsBatch.run, input_data, make_workflow_id('test-predict-stop')
|
||||||
|
)
|
||||||
|
assert_stop(postgres_engine, model_id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_3_3_predict_gate_repeat_batch_timestamp_equals_history_fails(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
bad_predict_model,
|
||||||
|
):
|
||||||
|
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)
|
||||||
|
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')
|
||||||
|
)
|
||||||
|
assert_repeat(postgres_engine, model_id, data)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_scenario_2_4_1_input_empty_data_stop(
|
||||||
|
temporal_test_env: WorkflowEnvironment,
|
||||||
|
temporal_worker: Worker,
|
||||||
|
test_activities: Activities,
|
||||||
|
postgres_engine,
|
||||||
|
):
|
||||||
|
"""EMPTY_DATA filter with STOP when query returns no rows (offload payload empty)."""
|
||||||
|
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}'))
|
||||||
|
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()
|
||||||
333
e2e/test_simple_metrics.py
Normal file
333
e2e/test_simple_metrics.py
Normal file
@@ -0,0 +1,333 @@
|
|||||||
|
"""
|
||||||
|
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'
|
||||||
13
encode.sh
Executable file
13
encode.sh
Executable file
@@ -0,0 +1,13 @@
|
|||||||
|
source ./venv/bin/activate
|
||||||
|
|
||||||
|
pip install pathspec
|
||||||
|
pip install pyyaml
|
||||||
|
|
||||||
|
echo "
|
||||||
|
.git" >> .gitignore
|
||||||
|
|
||||||
|
python encrypt.py ./ code --ignore .gitignore --chunk-size 100000
|
||||||
|
|
||||||
|
sed -i '/.git/d' .gitignore
|
||||||
|
|
||||||
|
xdg-open .
|
||||||
113
encrypt.py
Normal file
113
encrypt.py
Normal file
@@ -0,0 +1,113 @@
|
|||||||
|
import os
|
||||||
|
import argparse
|
||||||
|
from pathspec import PathSpec
|
||||||
|
import yaml # type: ignore
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
'''
|
||||||
|
Usage:
|
||||||
|
python .\encrypt.py path_to_dir output_file --ignore ignore_file --chunk-size 100000
|
||||||
|
'''
|
||||||
|
|
||||||
|
|
||||||
|
def load_ignore_patterns(ignore_file, include_library):
|
||||||
|
# Ensure the .gitignore file exists
|
||||||
|
if not os.path.exists(ignore_file):
|
||||||
|
raise FileNotFoundError(f"Ignore file not found at {ignore_file}")
|
||||||
|
|
||||||
|
# Load and parse the .gitignore patterns
|
||||||
|
with open(ignore_file, 'r') as file:
|
||||||
|
patterns = file.readlines()
|
||||||
|
if not include_library:
|
||||||
|
patterns.append('**/deploy/library/')
|
||||||
|
|
||||||
|
spec = PathSpec.from_lines('gitwildmatch', patterns)
|
||||||
|
return spec
|
||||||
|
|
||||||
|
|
||||||
|
def is_ignored(file_path, spec):
|
||||||
|
"""Check if a file should be ignored based on the ignore patterns."""
|
||||||
|
return spec.match_file(file_path) if spec else False
|
||||||
|
|
||||||
|
|
||||||
|
def encode_file_tree_to_yaml(directory, ignore_file, include_library):
|
||||||
|
"""Encode the file tree into a single YAML file."""
|
||||||
|
ignore_patterns = load_ignore_patterns(
|
||||||
|
ignore_file, include_library) if ignore_file else None
|
||||||
|
file_tree: dict[str, Any] = {}
|
||||||
|
|
||||||
|
for root, dirs, files in os.walk(directory):
|
||||||
|
# Skip ignored directories
|
||||||
|
dirs[:] = [d for d in dirs if not is_ignored(
|
||||||
|
os.path.join(root, d), ignore_patterns)]
|
||||||
|
|
||||||
|
for file in files:
|
||||||
|
file_path = os.path.join(root, file)
|
||||||
|
|
||||||
|
# Skip ignored files
|
||||||
|
if is_ignored(file_path, ignore_patterns):
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Read file content
|
||||||
|
try:
|
||||||
|
with open(file_path, 'r', encoding='utf-8') as f:
|
||||||
|
content = f.read()
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error reading file {file_path}: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
# Create nested dictionary structure
|
||||||
|
path_parts = os.path.relpath(file_path, directory).split(os.sep)
|
||||||
|
current_level = file_tree
|
||||||
|
|
||||||
|
# all except the last part (the file name)
|
||||||
|
for part in path_parts[:-1]:
|
||||||
|
current_level = current_level.setdefault(part, {})
|
||||||
|
|
||||||
|
# Add the file and its content
|
||||||
|
current_level[path_parts[-1]] = content
|
||||||
|
return yaml.dump(file_tree, default_flow_style=False)
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_and_write_file_tree_to_yaml(yaml_content, output_file, chunk_size=None):
|
||||||
|
"""Chunk the YAML content and write it to the output file."""
|
||||||
|
|
||||||
|
chunks = [yaml_content] if chunk_size is None else [
|
||||||
|
yaml_content[i:i + chunk_size] for i in range(0, len(yaml_content), chunk_size)]
|
||||||
|
|
||||||
|
for i, chunk in enumerate(chunks):
|
||||||
|
chunk_file = f"{output_file}_{i}.yaml"
|
||||||
|
# Write the file tree to the output YAML file
|
||||||
|
with open(chunk_file, 'w', encoding='utf-8') as yaml_file:
|
||||||
|
yaml_file.write(chunk)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Encrypts file tree to yaml file")
|
||||||
|
parser.add_argument("input_directory", help="Directory to encode")
|
||||||
|
parser.add_argument("output_yaml_file", help="Output YAML file")
|
||||||
|
parser.add_argument("--ignore", default=None,
|
||||||
|
help="Path to the ignore file")
|
||||||
|
parser.add_argument("--chunk-size", type=int, default=None,
|
||||||
|
help="Chunk size for the output YAML file")
|
||||||
|
parser.add_argument("--library", type=bool, default=False,
|
||||||
|
help="Incude the library in the output YAML file")
|
||||||
|
|
||||||
|
# Parse arguments
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Example usage
|
||||||
|
directory_to_encode = args.input_directory
|
||||||
|
ignore_file_path = args.ignore
|
||||||
|
output_yaml_file = args.output_yaml_file
|
||||||
|
include_library = args.library
|
||||||
|
|
||||||
|
content = encode_file_tree_to_yaml(
|
||||||
|
directory_to_encode, ignore_file_path, include_library)
|
||||||
|
chunk_and_write_file_tree_to_yaml(
|
||||||
|
content, output_yaml_file, args.chunk_size)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
2
git-requirements-mapping.txt
Normal file
2
git-requirements-mapping.txt
Normal file
@@ -0,0 +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
|
||||||
143
input_sample.json
Normal file
143
input_sample.json
Normal file
@@ -0,0 +1,143 @@
|
|||||||
|
{
|
||||||
|
"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"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
25
inter_arrival.py
Normal file
25
inter_arrival.py
Normal file
@@ -0,0 +1,25 @@
|
|||||||
|
# %%
|
||||||
|
|
||||||
|
# Load logs.txt
|
||||||
|
with open('logs.txt', 'r') as file:
|
||||||
|
lines = file.readlines()
|
||||||
|
|
||||||
|
# %%
|
||||||
|
import re
|
||||||
|
# Grep "inter-arrival_s=number" with regex
|
||||||
|
intervals = []
|
||||||
|
for line in lines:
|
||||||
|
match = re.search(r'inter-arrival_s=([0-9.]+)', line)
|
||||||
|
if match:
|
||||||
|
intervals.append(float(match.group(1)))
|
||||||
|
# %%
|
||||||
|
|
||||||
|
print(intervals)
|
||||||
|
# %%
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
plt.plot(intervals)
|
||||||
|
plt.ylabel('Inter-arrival time (s)')
|
||||||
|
plt.xlabel('Sample')
|
||||||
|
plt.title('Inter-arrival time distribution')
|
||||||
|
plt.show()
|
||||||
|
# %%
|
||||||
305
laborious-temporal-plugin-store-migration-plan.md
Normal file
305
laborious-temporal-plugin-store-migration-plan.md
Normal file
@@ -0,0 +1,305 @@
|
|||||||
|
---
|
||||||
|
tags:
|
||||||
|
- engineering
|
||||||
|
- sientia
|
||||||
|
- runtime-system
|
||||||
|
- laborious-temporal
|
||||||
|
- plugin-store
|
||||||
|
- migration-plan
|
||||||
|
created: 2026-03-02
|
||||||
|
modified: 2026-03-02
|
||||||
|
created_by: Vitor Pimentel
|
||||||
|
modified_by: Vitor Pimentel
|
||||||
|
status: draft
|
||||||
|
---
|
||||||
|
|
||||||
|
# Sientia Laborious Temporal — PluginStore & Wrapper Migration Plan
|
||||||
|
|
||||||
|
> Migration plan for evolving `sientia-dataops-laborious_temporal` from direct MLflow model loading to a runtime-aware architecture that uses Sientia model wrappers (`SientiaModel`) via their public methods, aligned with the runtime strategy.
|
||||||
|
|
||||||
|
## Summary
|
||||||
|
|
||||||
|
1. [[#Objectives and Scope|Objectives and Scope]] — What this migration must achieve
|
||||||
|
2. [[#Existing State Overview (laborious_temporal)|Existing State Overview]] — Current responsibilities and coupling points
|
||||||
|
3. [[#Requirements Mapping|Requirements Mapping]] — Functional and non-functional requirements
|
||||||
|
4. [[#Target Architecture|Target Architecture]] — Desired runtime and model interaction architecture
|
||||||
|
5. [[#Implementation Plan|Implementation Plan]] — Phased, detailed changes to apply
|
||||||
|
6. [[#Testing Strategy|Testing Strategy]] — How to validate the new behavior
|
||||||
|
7. [[#Rollout and Migration Strategy|Rollout and Migration Strategy]] — How to safely roll out the changes
|
||||||
|
8. [[#Related Documents|Related Documents]] — Cross-links to supporting documents
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Objectives and Scope
|
||||||
|
|
||||||
|
This migration focuses on the `sientia-dataops-laborious_temporal` application and aims to:
|
||||||
|
|
||||||
|
- Keep the **runtime-aware deployment model** consistent with the rest of the runtime system (Helm + `RUNTIME` env var, runtime installation via PluginStore).
|
||||||
|
- Ensure that **all interactions with models use the public methods of the Sientia wrapper** (`SientiaModel`):
|
||||||
|
- Use `SientiaModel.train(...)` and `retrain(...)` for training and retraining flows.
|
||||||
|
- Use `SientiaModel.predict(...)` and `SientiaModel.transform(...)` for inference and preprocessing.
|
||||||
|
- **Use the shared MLflow repository** (`SientiaMLflowRepository`) for all MLflow operations (load, runs, artifacts, promotion, production lookup, metadata logging); do not implement these in Laborious.
|
||||||
|
|
||||||
|
Out of scope:
|
||||||
|
|
||||||
|
- Replacing MLflow as the tracking and registry backend.
|
||||||
|
- Redesigning Temporal workflows (queues, retry policies) beyond what is required for the new model interaction style.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Existing State Overview (laborious_temporal)
|
||||||
|
|
||||||
|
Key components in `sientia-dataops-laborious_temporal`:
|
||||||
|
|
||||||
|
- **MLflow activities** (`laborious/activities/mlflow.py`)
|
||||||
|
- `MLFlow` class exposes Temporal activities for:
|
||||||
|
- `request_transform` — loads transformation models from MLflow and applies them to input data.
|
||||||
|
- `request_predict` — loads predictive models from MLflow and generates predictions.
|
||||||
|
- `retrain_model` — orchestrates retraining using historical data stored in MinIO and MLflow registry.
|
||||||
|
- `update_production_model` — promotes new versions to production.
|
||||||
|
- `get_reference_data` — fetches evaluation/reference datasets from model artifacts.
|
||||||
|
- These activities delegate ML-specific work to `MLFlowRepository`.
|
||||||
|
|
||||||
|
- **MLflow repository** (`laborious/utils/repository/model_repository.py`)
|
||||||
|
- `MLFlowRepository` encapsulates the interaction with MLflow:
|
||||||
|
- Model discovery and run resolution (`get_model_run_id`, `get_model_uri`, `get_experiment`, etc.).
|
||||||
|
- Artifact download and loading for both transformer and prediction models.
|
||||||
|
- Model caching and retention (`get_model`, `get_cached_operation`).
|
||||||
|
- Transformation and prediction entry points:
|
||||||
|
- `transform(...)` wraps `get_cached_operation(..., operation='transform')`.
|
||||||
|
- `predict(...)` wraps `get_cached_operation(..., operation='predict')`.
|
||||||
|
- Retraining orchestration (`fit_models`, `create_new_experiment`, `retrain_model`, `update_production_model`).
|
||||||
|
- Today:
|
||||||
|
- Models are loaded via MLflow flavors: sklearn, pyfunc, pytorch.
|
||||||
|
- When `flavor == 'pyfunc'` and `load_wrapper=True`, the repository loads a wrapper via:
|
||||||
|
- `raw_model = mlflow.pyfunc.load_model(artifact_path)`
|
||||||
|
- `model = raw_model._model_impl.python_model`
|
||||||
|
- Production models are resolved using **stages** in the Model Registry (for example, selecting the latest version in stage `Production`); **aliases such as `@production` are not used yet**, and models are registered explicitly as part of the current retrain/promotion flows.
|
||||||
|
- The wrapper’s `_model_impl` class does not extend `SientiaModel`.
|
||||||
|
- All this MLflow-specific logic is local to `laborious_temporal` and partially duplicated in `sientia-dataops-model-manager`, which motivates the extraction of a shared MLflow repository in `sientia-dataops-library` (see `mlflow-shared-repository-migration-plan`).
|
||||||
|
|
||||||
|
**MLflow:** All MLflow ops → [[mlflow-shared-repository-migration-plan|shared repository]]. Laborious uses the interface; `SientiaModel` lifecycle is in sientia-model-library.
|
||||||
|
|
||||||
|
### Current vs Target — High-level Flow
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart LR
|
||||||
|
subgraph current [Current State — Laborious Temporal]
|
||||||
|
direction TB
|
||||||
|
TemporalWorker["Temporal Worker"]
|
||||||
|
MlflowActivities["MLFlow Activities\nrequest_transform / request_predict / retrain_model"]
|
||||||
|
MLFlowRepositoryNode["MLFlowRepository"]
|
||||||
|
MLflowRegistry["MLflow Tracking + Registry"]
|
||||||
|
RawModel["Loaded Model\n(sklearn / pyfunc / pytorch)"]
|
||||||
|
|
||||||
|
TemporalWorker -->|"start workflow\n(Temporal)"| MlflowActivities
|
||||||
|
MlflowActivities -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode
|
||||||
|
MLFlowRepositoryNode -->|"search_model_versions()\ncurrent_stage == 'Production'"| MLflowRegistry
|
||||||
|
MLFlowRepositoryNode -->|"mlflow.*.load_model(model_uri)"| RawModel
|
||||||
|
MLFlowRepositoryNode -->|"raw_model.predict(data)\nraw_model.fit(data)"| RawModel
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph target [Target State — Laborious Temporal]
|
||||||
|
direction TB
|
||||||
|
TemporalWorker2["Temporal Worker"]
|
||||||
|
MlflowActivities2["MLFlow Activities\n(same APIs)"]
|
||||||
|
MLFlowRepositoryNode2["MLFlowRepository\n(wrapper-aware)"]
|
||||||
|
MLflowRegistry2["MLflow Tracking + Registry\n(aliases enabled)"]
|
||||||
|
SientiaWrapperNode["Wrapper Instance\n(extends SientiaModel)"]
|
||||||
|
|
||||||
|
TemporalWorker2 -->|"start workflow\n(Temporal)"| MlflowActivities2
|
||||||
|
MlflowActivities2 -->|"call transform()/predict()/retrain_model()"| MLFlowRepositoryNode2
|
||||||
|
MLFlowRepositoryNode2 -->|"get_model_version_by_alias('production')\n& models:/name@production"| MLflowRegistry2
|
||||||
|
MLFlowRepositoryNode2 -->|"mlflow.pyfunc.load_model(...)"| SientiaWrapperNode
|
||||||
|
MLFlowRepositoryNode2 -->|"wrapper.transform(...)\nwrapper.predict(...)\nwrapper.train()/retrain()"| SientiaWrapperNode
|
||||||
|
SientiaWrapperNode -->|"store_model(...)\n(auto-register + update alias)"| MLflowRegistry2
|
||||||
|
end
|
||||||
|
```
|
||||||
|
|
||||||
|
### Current vs Target — Retrain Hot Path (Code Sketch)
|
||||||
|
|
||||||
|
Current retrain flow inside `MLFlowRepository.fit_models` / `retrain_model` (simplified):
|
||||||
|
|
||||||
|
```python
|
||||||
|
data_model, _ = await self.download_model(
|
||||||
|
model_name=model_name,
|
||||||
|
metadata=metadata,
|
||||||
|
model_type="transform",
|
||||||
|
flavor=transform_flavor,
|
||||||
|
load_wrapper=(transform_flavor == "pyfunc"),
|
||||||
|
)
|
||||||
|
|
||||||
|
prediction_model, _ = await self.download_model(
|
||||||
|
model_name=model_name,
|
||||||
|
metadata=metadata,
|
||||||
|
model_type="predict",
|
||||||
|
flavor=predict_flavor,
|
||||||
|
load_wrapper=(predict_flavor == "pyfunc"),
|
||||||
|
)
|
||||||
|
|
||||||
|
treated_data_candidate = data_model.fit(data) # or data_model.predict(data)
|
||||||
|
...
|
||||||
|
prediction_model.fit(retrain_dataset) # direct fit on underlying model
|
||||||
|
```
|
||||||
|
|
||||||
|
Target retrain flow when using wrappers that extend `SientiaModel`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
wrapper, _ = await self.download_model(
|
||||||
|
model_name=model_name,
|
||||||
|
metadata=metadata,
|
||||||
|
# model_type="predict", # wrapper owns both transformer + model stack
|
||||||
|
# flavor="pyfunc", flavor will always be pyfunc, this parameter will be removed
|
||||||
|
# load_wrapper=True, wrapper will always be loaded from _model_impl
|
||||||
|
)
|
||||||
|
|
||||||
|
# First-time training or full retrain using public API
|
||||||
|
wrapper.train(
|
||||||
|
train_data=train_df, # features + target
|
||||||
|
val_data=val_df,
|
||||||
|
target=target_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Incremental retrain (when appropriate)
|
||||||
|
wrapper.retrain(full_retrain_df)
|
||||||
|
|
||||||
|
transformed_df, trans_meta = wrapper.transform(raw_df)
|
||||||
|
pred_df, pred_meta = wrapper.predict({}, transformed_df, params={})
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Requirements Mapping
|
||||||
|
|
||||||
|
### Functional Requirements
|
||||||
|
|
||||||
|
| ID | Requirement | Description | Impacted Areas |
|
||||||
|
| ----- | -------------------------------------------------- | ------------------------------------------------------------------------------------------------ | ------------------------------------------------------ |
|
||||||
|
| FR-01 | Runtime detection and installation | Align worker startup with `RUNTIME` and runtime installation via PluginStore | Worker |
|
||||||
|
| FR-02 | Wrapper‑based training using public API | Retraining must call the public `retrain(...)` method of `SientiaModel` | `fit_models`, `retrain_model` paths |
|
||||||
|
| FR-03 | Wrapper‑based inference using public API | Inference must call `predict(...)` and `transform(...)` on the wrapper; obtain wrappers via `SientiaMLflowRepository` | Activities, shared repository |
|
||||||
|
| FR-04 | Use shared MLflow repository | Delegate all MLflow operations (load, runs, artifacts, promotion, production lookup) to `SientiaMLflowRepository`; do not implement in Laborious | [[mlflow-shared-repository-migration-plan]] |
|
||||||
|
| FR-05 | Model configuration for training/retraining | `model_config` must define `target` and `retention_minutes`; all models are pyfunc + wrapper (no flavor selection) | `model_config` structures |
|
||||||
|
|
||||||
|
### Non-Functional Requirements
|
||||||
|
|
||||||
|
| ID | Requirement | Description | Impacted Areas |
|
||||||
|
| ------ | -------------------------------- | ------------------------------------------------------------------------------------------------------------------------ | --------------------------------------- |
|
||||||
|
| NFR-01 | Consistent public API usage | All wrapper interactions must go through public `SientiaModel` methods; metadata logging is handled by the shared repository | Activities |
|
||||||
|
| NFR-02 | Fail fast when wrappers unavailable | Where wrappers are not yet available, raise an explicit error requesting model update | Shared repository, config |
|
||||||
|
| NFR-03 | Testability | Enable unit tests to validate wrapper-based flows and integration with shared repository | `tests/laborious` |
|
||||||
|
| NFR-04 | Operational safety | Retraining and promotion behavior must remain auditable and robust | Retrain & promotion flows |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Target Architecture
|
||||||
|
|
||||||
|
### Wrapper-centric model interactions
|
||||||
|
|
||||||
|
The target state for `laborious_temporal` is:
|
||||||
|
|
||||||
|
- All training and retraining logic goes through the wrapper’s public methods: `train(...)` and `retrain(...)`.
|
||||||
|
- Inference and transformation use `transform(df)` and `predict(context, df, params)`.
|
||||||
|
- All MLflow operations (load, runs, artifacts, promotion, production lookup) go through `SientiaMLflowRepository` (see [[mlflow-shared-repository-migration-plan]]).
|
||||||
|
|
||||||
|
### Use of Shared MLflow Repository
|
||||||
|
|
||||||
|
Laborious obtains wrappers and performs all MLflow operations via `SientiaMLflowRepository`. The shared repository (see [[mlflow-shared-repository-migration-plan]]) owns: alias-based URIs, pyfunc loading, wrapper extraction, promotion, runs, artifacts, and metadata logging. Laborious activities call `repo.load_wrapper(...)` and then the wrapper’s public methods (`transform`, `predict`, `train`, `retrain`); they do not implement MLflow logic.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Implementation Plan
|
||||||
|
|
||||||
|
### Phase 0 — Design and Configuration Alignment
|
||||||
|
|
||||||
|
- **P0-01**: Catalog model types and flavors used by `laborious_temporal`:
|
||||||
|
- For each active model:
|
||||||
|
- Flavor will be always pyfunc
|
||||||
|
- Whether a `_model_impl` wrapper is already present and extends `SientiaModel`.
|
||||||
|
- **P0-02**: Define configuration fields in `model_config` for wrapper usage:
|
||||||
|
- Example:
|
||||||
|
- `target` field for training/retraining (already partially present).
|
||||||
|
- `retention_minutes` field for model retention in minutes (unchanged).
|
||||||
|
|
||||||
|
### Phase 1 — Inference via SientiaModel Public API (using shared repository)
|
||||||
|
|
||||||
|
- **P1-01**: Replace `download_model` with `SientiaMLflowRepository.load_wrapper(...)`; remove local MLflow loading logic.
|
||||||
|
- **P1-02**: Update `get_cached_operation` to call `wrapper.transform(data)` and `wrapper.predict({}, data)`; unpack returned metadata; metadata logging is handled by the shared repository.
|
||||||
|
- **P1-03**: Ensure activities (`request_transform`, `request_predict`) remain unchanged externally (inputs/outputs unchanged).
|
||||||
|
|
||||||
|
### Phase 2 — Retraining via SientiaModel.retrain
|
||||||
|
|
||||||
|
- **P2-01**: Refactor `fit_models` to call `wrapper.retrain(data)` and `wrapper.train(...)`; use `SientiaMLflowRepository` for runs, metrics, artifacts, and promotion (no local MLflow logic).
|
||||||
|
|
||||||
|
### Phase 3 — MLflow Logging and Promotion (via shared repository)
|
||||||
|
|
||||||
|
- **P3-01**: In `create_new_experiment`, use the wrapper’s `store_model(...)` for model artifacts; use `SientiaMLflowRepository` for runs, metrics, and any additional MLflow operations.
|
||||||
|
- **P3-02**: Use `SientiaMLflowRepository.promote_to_alias(...)` for production promotion; model registration is auto-handled by wrappers.
|
||||||
|
|
||||||
|
### Phase 4 — Runtime Alignment
|
||||||
|
|
||||||
|
- **P4-01**:`laborious_temporal` is also deployed via the runtime-aware Helm chart:
|
||||||
|
- Read `RUNTIME` env var.
|
||||||
|
- Install runtime via PluginStore before starting Temporal workers.
|
||||||
|
- **P4-02**: Standardize worker queue name as `{project_name}-{runtime}-queue`.
|
||||||
|
- **P4-03**: Fix quality pipelines to a single runtime (to be decided).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Testing Strategy
|
||||||
|
|
||||||
|
- **T1 — Unit tests**
|
||||||
|
- Add tests in `tests/laborious/utils/repository/test_model_repository.py` to cover:
|
||||||
|
- Wrapper-based `get_cached_operation` for both `transform` and `predict` (using shared repository).
|
||||||
|
- Wrapper-based `fit_models` and `retrain_model` paths calling `train` and `retrain` respectively.
|
||||||
|
- Code coverage must be 100%.
|
||||||
|
- **T2 — Integration tests**
|
||||||
|
- Use existing end-to-end tests under `e2e/`:
|
||||||
|
- Configure a model with `SientiaModel` wrapper (loaded via shared repository).
|
||||||
|
- Run full prediction and retrain workflows; compare predictions, retrain outcomes, and MLflow artifacts.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Rollout and Migration Strategy
|
||||||
|
|
||||||
|
- **R1 — Create models with the new architecture**
|
||||||
|
- Update or create models in the modeling pipeline so they:
|
||||||
|
- Use wrappers that extend `SientiaModel`.
|
||||||
|
- Correctly implement `train`, `retrain`, `predict`, `transform`, and `store_model`.
|
||||||
|
- Produce structured metadata in `transform_meta` and `pred_meta`.
|
||||||
|
- Publish these models to a test store (or dedicated branch/experiment) for initial validation.
|
||||||
|
|
||||||
|
- **R2 — Provision runtimes for the new models**
|
||||||
|
- Configure and install dedicated runtimes for the new models:
|
||||||
|
- Ensure runtime dependencies (Python and system libraries) are available via PluginStore/runtime installer.
|
||||||
|
- Validate that each runtime can:
|
||||||
|
- Load the wrapper through MLflow.
|
||||||
|
- Execute `transform` and `predict` end-to-end on sample data.
|
||||||
|
|
||||||
|
- **R3 — Run new models in real workflows**
|
||||||
|
- Integrate the new wrapper-based models into real Laborious workflows, initially in non-critical environments:
|
||||||
|
- Route only a subset of flows or entities to the new models.
|
||||||
|
- Monitor logs (including metadata), business metrics, and retraining/promotion behavior.
|
||||||
|
- Promote these models to production using aliases (`@production`) in MLflow 3+.
|
||||||
|
|
||||||
|
- **R4 — Migrate remaining models progressively**
|
||||||
|
- Define migration waves by model family:
|
||||||
|
- For each existing model:
|
||||||
|
- Create or adapt a `SientiaModel` wrapper.
|
||||||
|
- Provision the corresponding runtime.
|
||||||
|
- Execute the test cycle (T1–T2) from the previous section.
|
||||||
|
- Update aliases so production traffic uses the new wrapper.
|
||||||
|
- After all models are migrated:
|
||||||
|
- Remove legacy stage-based paths (`Production`) and non-wrapper models.
|
||||||
|
- Simplify the codebase to assume wrappers + aliases only.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Related Documents
|
||||||
|
|
||||||
|
- [[mlflow-shared-repository-migration-plan|MLflow Shared Repository Migration Plan]] — Concepts implemented in the common interface (production lookup, wrapper loading, promotion, metadata logging, etc.)
|
||||||
|
- [[model-manager-plugin-store-migration-plan|Model Manager PluginStore Migration Plan]]
|
||||||
|
- [[analytics-implementation-plan|Runtime Analytics Helm Implementation Plan]]
|
||||||
|
- [[analytics|Runtime Analytics Architecture and Analysis]]
|
||||||
|
- [[../model-plugin-system/06-end-to-end-flow|Model Plugin System — End-to-End Flow]]
|
||||||
|
|
||||||
106
model_convert.ipynb
Normal file
106
model_convert.ipynb
Normal file
@@ -0,0 +1,106 @@
|
|||||||
|
{
|
||||||
|
"cells": [
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": 23,
|
||||||
|
"id": "e838ff21",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"import csv\n",
|
||||||
|
"\n",
|
||||||
|
"def csv_to_tag_lists(csv_path: str) -> dict:\n",
|
||||||
|
" read_tags = []\n",
|
||||||
|
" write_tags = []\n",
|
||||||
|
"\n",
|
||||||
|
" def to_float(val):\n",
|
||||||
|
" try:\n",
|
||||||
|
" return float(str(val).strip())\n",
|
||||||
|
" except Exception:\n",
|
||||||
|
" return None\n",
|
||||||
|
"\n",
|
||||||
|
" with open(csv_path, newline=\"\", encoding=\"utf-8\") as f:\n",
|
||||||
|
" reader = csv.DictReader(f)\n",
|
||||||
|
" for row in reader:\n",
|
||||||
|
" # Basic normalization\n",
|
||||||
|
" op = (row.get(\"operation\") or \"\").strip()\n",
|
||||||
|
"\n",
|
||||||
|
" if op == \"READ\":\n",
|
||||||
|
" # Build common tag payload with required mappings\n",
|
||||||
|
" tag = {\n",
|
||||||
|
" \"server_id\": \"1\",\n",
|
||||||
|
" \"tag_address\": row.get(\"opc_tag\"),\n",
|
||||||
|
" \"tag_name\": row.get(\"name\"),\n",
|
||||||
|
" \"data_range\": [to_float(row.get(\"min_value\")), to_float(row.get(\"max_value\"))],\n",
|
||||||
|
" \"aggr_func\": row.get(\"aggregation_func\").lower(),\n",
|
||||||
|
" # keep other fields with their original names\n",
|
||||||
|
" \"frequency\": row.get(\"frequency\"),\n",
|
||||||
|
" \"local\": row.get(\"local\"),\n",
|
||||||
|
" \"area\": row.get(\"area\"),\n",
|
||||||
|
" \"description\": row.get(\"description\"),\n",
|
||||||
|
" }\n",
|
||||||
|
"\n",
|
||||||
|
" read_tags.append(tag)\n",
|
||||||
|
"\n",
|
||||||
|
" else:\n",
|
||||||
|
" tag = {\n",
|
||||||
|
" \"server_id\": \"1\",\n",
|
||||||
|
" \"addr\": row.get(\"opc_tag\"),\n",
|
||||||
|
" \"tag_name\": row.get(\"name\"),\n",
|
||||||
|
" \"local\": row.get(\"local\"),\n",
|
||||||
|
" \"area\": row.get(\"area\"),\n",
|
||||||
|
" \"description\": row.get(\"description\"),\n",
|
||||||
|
" }\n",
|
||||||
|
" \n",
|
||||||
|
" if op == \"WRITE_PREDICTION\":\n",
|
||||||
|
" tag[\"type\"] = \"prediction\"\n",
|
||||||
|
" write_tags.append(tag)\n",
|
||||||
|
" elif op == \"WRITE_CONFIDENCE\":\n",
|
||||||
|
" tag[\"type\"] = \"confidence\"\n",
|
||||||
|
" write_tags.append(tag)\n",
|
||||||
|
" # ignore any other operation values silently\n",
|
||||||
|
"\n",
|
||||||
|
" return {\"read_tags\": read_tags, \"write_tags\": write_tags}"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"id": "4621cd43",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"import json\n",
|
||||||
|
"\n",
|
||||||
|
"file_names = [\"Courier - Página1.csv\"]\n",
|
||||||
|
"\n",
|
||||||
|
"for file_name in file_names:\n",
|
||||||
|
" write_file = file_name.replace(\".csv\", \".json\")\n",
|
||||||
|
"\n",
|
||||||
|
" with open(write_file, \"w\", encoding=\"utf-8\") as f:\n",
|
||||||
|
" json.dump(csv_to_tag_lists(file_name), f, indent=2, ensure_ascii=False)"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"kernelspec": {
|
||||||
|
"display_name": "venv",
|
||||||
|
"language": "python",
|
||||||
|
"name": "python3"
|
||||||
|
},
|
||||||
|
"language_info": {
|
||||||
|
"codemirror_mode": {
|
||||||
|
"name": "ipython",
|
||||||
|
"version": 3
|
||||||
|
},
|
||||||
|
"file_extension": ".py",
|
||||||
|
"mimetype": "text/x-python",
|
||||||
|
"name": "python",
|
||||||
|
"nbconvert_exporter": "python",
|
||||||
|
"pygments_lexer": "ipython3",
|
||||||
|
"version": "3.11.13"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"nbformat": 4,
|
||||||
|
"nbformat_minor": 5
|
||||||
|
}
|
||||||
158
pyproject.toml
Normal file
158
pyproject.toml
Normal file
@@ -0,0 +1,158 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["setuptools>=61.0"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
name = "laborious"
|
||||||
|
version = "0.0.0"
|
||||||
|
description = "Sientia DataOps Laborious - ML Model Orchestration System"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.11"
|
||||||
|
authors = [
|
||||||
|
{name = "Aignosi", email = "dev@aignosi.com"}
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
line-length = 100
|
||||||
|
target-version = "py311"
|
||||||
|
exclude = [
|
||||||
|
".git",
|
||||||
|
".venv",
|
||||||
|
"venv",
|
||||||
|
"__pycache__",
|
||||||
|
"*.pyc",
|
||||||
|
".pytest_cache",
|
||||||
|
"htmlcov",
|
||||||
|
"tests/laborious/workflows/subworkflows/test_prediction_process.py",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint]
|
||||||
|
select = [
|
||||||
|
"E", # pycodestyle errors
|
||||||
|
"W", # pycodestyle warnings
|
||||||
|
"F", # pyflakes
|
||||||
|
"I", # isort
|
||||||
|
"B", # flake8-bugbear
|
||||||
|
"C4", # flake8-comprehensions
|
||||||
|
"UP", # pyupgrade
|
||||||
|
"N", # pep8-naming
|
||||||
|
"YTT", # flake8-2020
|
||||||
|
"S", # flake8-bandit
|
||||||
|
"BLE", # flake8-blind-except
|
||||||
|
"A", # flake8-builtins
|
||||||
|
"C90", # mccabe complexity
|
||||||
|
]
|
||||||
|
|
||||||
|
ignore = [
|
||||||
|
"BLE001", # ignore blind except, we need to send notifications with any error
|
||||||
|
"E501", # line too long (handled by formatter)
|
||||||
|
"S101", # use of assert (needed for tests)
|
||||||
|
"S105", # possible hardcoded password (false positives)
|
||||||
|
"S106", # possible hardcoded password (false positives)
|
||||||
|
"S608", # potential sql injection (false positives)
|
||||||
|
"N802", # function name should be lowercase (temporal decorators)
|
||||||
|
"N806", # variable in function should be lowercase
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint.per-file-ignores]
|
||||||
|
"tests/**/*.py" = [
|
||||||
|
"S101", # assert allowed in tests
|
||||||
|
"S105", # hardcoded passwords ok in tests
|
||||||
|
"S106", # hardcoded passwords ok in tests
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint.mccabe]
|
||||||
|
max-complexity = 15
|
||||||
|
|
||||||
|
[tool.ruff.format]
|
||||||
|
quote-style = "single"
|
||||||
|
indent-style = "space"
|
||||||
|
line-ending = "auto"
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
python_version = "3.11"
|
||||||
|
warn_return_any = false
|
||||||
|
warn_unused_configs = true
|
||||||
|
disallow_untyped_defs = false
|
||||||
|
disallow_incomplete_defs = false
|
||||||
|
check_untyped_defs = true
|
||||||
|
no_implicit_optional = true
|
||||||
|
warn_redundant_casts = true
|
||||||
|
warn_unused_ignores = false
|
||||||
|
warn_no_return = true
|
||||||
|
strict_equality = true
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
# Ignore missing imports for external packages
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "temporalio.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "sientia_do.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "mlflow.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "prometheus_client.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "sientia.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "pandas.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
testpaths = ["tests"]
|
||||||
|
python_files = ["test_*.py"]
|
||||||
|
python_classes = ["Test*"]
|
||||||
|
python_functions = ["test_*"]
|
||||||
|
addopts = [
|
||||||
|
"-v",
|
||||||
|
"--strict-markers",
|
||||||
|
]
|
||||||
|
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",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.coverage.run]
|
||||||
|
source = ["laborious"]
|
||||||
|
omit = [
|
||||||
|
"*/tests/*",
|
||||||
|
"*/venv/*",
|
||||||
|
"*/__pycache__/*",
|
||||||
|
"*/site-packages/*",
|
||||||
|
]
|
||||||
|
branch = true
|
||||||
|
|
||||||
|
[tool.coverage.report]
|
||||||
|
precision = 2
|
||||||
|
show_missing = true
|
||||||
|
skip_covered = false
|
||||||
|
exclude_lines = [
|
||||||
|
"pragma: no cover",
|
||||||
|
"def __repr__",
|
||||||
|
"def __str__",
|
||||||
|
"raise AssertionError",
|
||||||
|
"raise NotImplementedError",
|
||||||
|
"if __name__ == .__main__.:",
|
||||||
|
"if TYPE_CHECKING:",
|
||||||
|
"class .*\\bProtocol\\):",
|
||||||
|
"@(abc\\.)?abstractmethod",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.coverage.html]
|
||||||
|
directory = "htmlcov"
|
||||||
|
|
||||||
|
[tool.bandit]
|
||||||
|
exclude_dirs = ["tests", "venv", ".venv"]
|
||||||
|
skips = ["B101", "B601", "B608"] # Skip assert, shell injection, and SQL injection (false positives)
|
||||||
21
requirements-dev.txt
Normal file
21
requirements-dev.txt
Normal file
@@ -0,0 +1,21 @@
|
|||||||
|
# Development and Testing Dependencies
|
||||||
|
# These packages are only needed for development, testing, and code quality checks
|
||||||
|
# Install with: pip install -r requirements-dev.txt
|
||||||
|
|
||||||
|
# Code Quality & Linting
|
||||||
|
ruff>=0.1.0 # Fast Python linter and formatter (replaces flake8, black, isort)
|
||||||
|
mypy>=1.7.0 # Static type checker
|
||||||
|
bandit>=1.7.5 # Security vulnerability scanner
|
||||||
|
pandas-stubs>=2.0.0 # Type stubs for pandas
|
||||||
|
types-requests>=2.31.0 # Type stubs for requests
|
||||||
|
|
||||||
|
# Testing
|
||||||
|
pytest>=7.4.0 # Testing framework
|
||||||
|
pytest-cov>=4.1.0 # Coverage plugin for pytest
|
||||||
|
pytest-asyncio>=0.21.0 # Async test support (already in main requirements)
|
||||||
|
testcontainers[postgres,minio] # PostgreSQL and MinIO containers for E2E tests
|
||||||
|
|
||||||
|
# Development Tools
|
||||||
|
ipython>=8.12.0 # Enhanced Python shell
|
||||||
|
ipdb>=0.13.13 # IPython debugger
|
||||||
|
ipykernel==6.30.1 # IPython kernel for Jupyter notebooks
|
||||||
18
requirements-local.txt
Normal file
18
requirements-local.txt
Normal file
@@ -0,0 +1,18 @@
|
|||||||
|
temporalio
|
||||||
|
psycopg2-binary
|
||||||
|
sqlalchemy
|
||||||
|
asyncua==1.0.6
|
||||||
|
redis
|
||||||
|
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.1
|
||||||
|
git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.10.0
|
||||||
|
prometheus-client
|
||||||
|
botocore
|
||||||
|
boto3
|
||||||
|
s3fs
|
||||||
|
pyarrow
|
||||||
|
kaleido
|
||||||
|
hyperopt
|
||||||
|
shap
|
||||||
|
pycurl
|
||||||
|
scipy<1.14.0
|
||||||
|
scikit-learn==1.5.2
|
||||||
@@ -1,8 +1,18 @@
|
|||||||
temporalio
|
temporalio
|
||||||
psycopg2-binary
|
psycopg2-binary
|
||||||
sqlalchemy
|
sqlalchemy
|
||||||
asyncua
|
asyncua==1.0.6
|
||||||
redis
|
redis
|
||||||
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.4.6
|
sientia_do>=1.12.1
|
||||||
git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.39.0
|
sientia_model>=0.8.2
|
||||||
prometheus-client
|
prometheus-client
|
||||||
|
botocore
|
||||||
|
boto3
|
||||||
|
s3fs
|
||||||
|
pyarrow
|
||||||
|
kaleido
|
||||||
|
hyperopt
|
||||||
|
shap
|
||||||
|
pycurl
|
||||||
|
scipy<1.14.0
|
||||||
|
scikit-learn==1.5.2
|
||||||
@@ -1,30 +0,0 @@
|
|||||||
# syntax=docker/dockerfile:1.4
|
|
||||||
|
|
||||||
FROM python:3.11-slim
|
|
||||||
|
|
||||||
# Enable use of SSH agent/socket
|
|
||||||
# This line enables SSH during build
|
|
||||||
# (don't forget the syntax header above)
|
|
||||||
RUN apt-get update && apt-get install -y git openssh-client && rm -rf /var/lib/apt/lists/*
|
|
||||||
|
|
||||||
# Use build-time SSH mount for Git clone
|
|
||||||
# The SSH key will NOT remain in the image
|
|
||||||
# IMPORTANT: this block requires BuildKit
|
|
||||||
# and the --ssh flag during docker build
|
|
||||||
|
|
||||||
# SSH config to skip host key check (safe in CI/local dev)
|
|
||||||
RUN mkdir -p /root/.ssh && echo "StrictHostKeyChecking no" > /root/.ssh/config
|
|
||||||
|
|
||||||
WORKDIR /app
|
|
||||||
|
|
||||||
# Clone using SSH
|
|
||||||
ARG GIT_REPO
|
|
||||||
ARG GIT_BRANCH=main
|
|
||||||
|
|
||||||
# Mount SSH key just for this RUN
|
|
||||||
RUN --mount=type=ssh git clone --branch ${GIT_BRANCH} ${GIT_REPO} .
|
|
||||||
|
|
||||||
# Install requirements if exists
|
|
||||||
RUN if [ -f requirements.txt ]; then pip install --no-cache-dir -r requirements.txt; fi
|
|
||||||
|
|
||||||
CMD ["python", "server.py"]
|
|
||||||
@@ -3,7 +3,7 @@ sonar.projectName=sientia-dataops-laborious_temporal
|
|||||||
sonar.sources=laborious
|
sonar.sources=laborious
|
||||||
sonar.tests=tests
|
sonar.tests=tests
|
||||||
sonar.projectVersion=1.0.0
|
sonar.projectVersion=1.0.0
|
||||||
sonar.coverage.exclusions=laborious/worker/worker.py
|
sonar.coverage.exclusions=laborious/worker/*
|
||||||
sonar.qualitygate.wait=true
|
sonar.qualitygate.wait=true
|
||||||
sonar.qualitygate.timeout=300
|
sonar.qualitygate.timeout=300
|
||||||
sonar.python.coverage.reportPaths=coverage.xml
|
sonar.python.coverage.reportPaths=coverage.xml
|
||||||
|
|||||||
1028
tests.ipynb
1028
tests.ipynb
File diff suppressed because one or more lines are too long
66
tests/conftest.py
Normal file
66
tests/conftest.py
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
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]
|
||||||
|
|
||||||
|
# The production code converts SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES to int at import-time.
|
||||||
|
# Tests must set it to a valid integer string to avoid import errors.
|
||||||
|
os.environ.setdefault('SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES', '1')
|
||||||
|
|
||||||
|
|
||||||
|
class DummyMinioDataFramePayload:
|
||||||
|
"""
|
||||||
|
Minimal payload double used by unit tests.
|
||||||
|
|
||||||
|
The production workflow/gates expect a MinioDataFramePayload-like object with:
|
||||||
|
- async retrieve(minio_repo, workflow_metadata) -> DataFrame | dict
|
||||||
|
- has_data() -> bool
|
||||||
|
- cleanup_prefix() -> str | None
|
||||||
|
- last_timestamp: attribute
|
||||||
|
- status: attribute
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
retrieve_return=None,
|
||||||
|
has_data: bool = True,
|
||||||
|
cleanup_prefix: str | None = None,
|
||||||
|
last_timestamp: str = '2024-01-01',
|
||||||
|
status: dict | None = None,
|
||||||
|
):
|
||||||
|
self._retrieve_return = retrieve_return
|
||||||
|
self._has_data = has_data
|
||||||
|
self._cleanup_prefix = cleanup_prefix
|
||||||
|
self.last_timestamp = last_timestamp
|
||||||
|
self.status = status
|
||||||
|
|
||||||
|
async def retrieve(self, _minio_repo, _workflow_metadata=None):
|
||||||
|
return self._retrieve_return
|
||||||
|
|
||||||
|
def has_data(self) -> bool:
|
||||||
|
return self._has_data
|
||||||
|
|
||||||
|
def cleanup_prefix(self) -> str | None:
|
||||||
|
return self._cleanup_prefix
|
||||||
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
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.
|
||||||
|
"""
|
||||||
@@ -1,18 +1,32 @@
|
|||||||
from pytest import mark
|
from unittest.mock import ANY, MagicMock, patch
|
||||||
from unittest.mock import patch, MagicMock, ANY
|
|
||||||
from sientia_do.temporal.activities.postgres import Postgres
|
|
||||||
from laborious.activities.activities import Activities
|
from laborious.activities.activities import Activities
|
||||||
from laborious.activities.mlflow import MLFlow
|
from laborious.activities.api import API
|
||||||
from laborious.activities.gates import Gates
|
from laborious.activities.gates import Gates
|
||||||
|
from laborious.activities.mlflow import MLFlow
|
||||||
|
from laborious.activities.model_metrics import ModelMetrics
|
||||||
from laborious.activities.opc import OPC
|
from laborious.activities.opc import OPC
|
||||||
|
from laborious.activities.storage import Storage
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.activities.activities.Postgres.__init__')
|
@patch('laborious.activities.activities.Storage.__init__')
|
||||||
@patch('laborious.activities.activities.MLFlow.__init__')
|
@patch('laborious.activities.activities.MLFlow.__init__')
|
||||||
@patch('laborious.activities.activities.OPC.__init__')
|
@patch('laborious.activities.activities.OPC.__init__')
|
||||||
@patch('laborious.activities.activities.Gates.__init__')
|
@patch('laborious.activities.activities.Gates.__init__')
|
||||||
def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_postgres_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__(
|
||||||
|
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,
|
||||||
|
):
|
||||||
postgres_config = {
|
postgres_config = {
|
||||||
'host': 'localhost',
|
'host': 'localhost',
|
||||||
'port': 5432,
|
'port': 5432,
|
||||||
@@ -20,20 +34,31 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_postgre
|
|||||||
'password': 'postgres',
|
'password': 'postgres',
|
||||||
'dbname': 'postgres',
|
'dbname': 'postgres',
|
||||||
'min_connections': 1,
|
'min_connections': 1,
|
||||||
'max_connections': 10
|
'max_connections': 10,
|
||||||
}
|
}
|
||||||
|
|
||||||
mlflow_config = {
|
minio_config = {
|
||||||
'host': 'localhost',
|
'endpoint_url': 'localhost:9000',
|
||||||
'port': 5000,
|
'access_key': 'minio',
|
||||||
'username': 'mlflow',
|
'secret_key': 'minio123',
|
||||||
'password': 'mlflow'
|
'default_bucket': 'test',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
mlflow_repository = MagicMock()
|
||||||
|
plugin_store = MagicMock()
|
||||||
|
|
||||||
opc_config = {
|
opc_config = {
|
||||||
'bootstrap_servers': 'localhost:9092',
|
'bootstrap_servers': 'localhost:9092',
|
||||||
'polling_time': 1000,
|
'polling_time': 1000,
|
||||||
'group_id': 'test-group'
|
'group_id': 'test-group',
|
||||||
|
}
|
||||||
|
|
||||||
|
pi_web_api_config = {
|
||||||
|
'base_url': 'https://test-pi-server.com',
|
||||||
|
'auth_type': 'bearer',
|
||||||
|
'auth_token': 'test_token',
|
||||||
}
|
}
|
||||||
|
|
||||||
logger = MagicMock()
|
logger = MagicMock()
|
||||||
@@ -41,19 +66,24 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_postgre
|
|||||||
|
|
||||||
activities = Activities(
|
activities = Activities(
|
||||||
postgres_config=postgres_config,
|
postgres_config=postgres_config,
|
||||||
mlflow_config=mlflow_config,
|
plugin_store=plugin_store,
|
||||||
|
minio_config=minio_config,
|
||||||
opc_config=opc_config,
|
opc_config=opc_config,
|
||||||
|
pi_web_api_config=pi_web_api_config,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
mlflow_repository=mlflow_repository,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert isinstance(activities, Activities)
|
assert isinstance(activities, Activities)
|
||||||
assert isinstance(activities, Postgres)
|
assert isinstance(activities, Storage)
|
||||||
assert isinstance(activities, MLFlow)
|
assert isinstance(activities, MLFlow)
|
||||||
assert isinstance(activities, OPC)
|
assert isinstance(activities, OPC)
|
||||||
assert isinstance(activities, Gates)
|
assert isinstance(activities, Gates)
|
||||||
|
assert isinstance(activities, ModelMetrics)
|
||||||
|
assert isinstance(activities, API)
|
||||||
|
|
||||||
mock_postgres_init.assert_called_once_with(
|
mock_storage_init.assert_called_once_with(
|
||||||
ANY,
|
ANY,
|
||||||
host=postgres_config['host'],
|
host=postgres_config['host'],
|
||||||
port=postgres_config['port'],
|
port=postgres_config['port'],
|
||||||
@@ -62,40 +92,85 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_postgre
|
|||||||
dbname=postgres_config['dbname'],
|
dbname=postgres_config['dbname'],
|
||||||
min_connections=postgres_config['min_connections'],
|
min_connections=postgres_config['min_connections'],
|
||||||
max_connections=postgres_config['max_connections'],
|
max_connections=postgres_config['max_connections'],
|
||||||
|
retention_hours=minio_config['retention_hours'],
|
||||||
|
minio_repository=mock_minio_repository.return_value,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_mlflow_init.assert_called_once_with(
|
mock_mlflow_init.assert_called_once_with(
|
||||||
ANY,
|
ANY,
|
||||||
mlflow_host=mlflow_config['host'],
|
mlflow_repository=mlflow_repository,
|
||||||
mlflow_port=mlflow_config['port'],
|
plugin_store=plugin_store,
|
||||||
mlflow_username=mlflow_config['username'],
|
minio_repository=mock_minio_repository.return_value,
|
||||||
mlflow_password=mlflow_config['password'],
|
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_opc_init.assert_called_once_with(
|
mock_opc_init.assert_called_once_with(
|
||||||
ANY,
|
ANY,
|
||||||
opc_servers=opc_config,
|
opc_servers=opc_config,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_gates_init.assert_called_once_with(
|
mock_gates_init.assert_called_once_with(
|
||||||
ANY,
|
ANY,
|
||||||
|
minio_repository=mock_minio_repository.return_value,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_model_metrics_init.assert_called_once_with(
|
||||||
|
ANY,
|
||||||
|
logger=logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_api_init.assert_called_once_with(
|
||||||
|
ANY,
|
||||||
|
base_url=pi_web_api_config['base_url'],
|
||||||
|
auth_type=pi_web_api_config['auth_type'],
|
||||||
|
auth_token=pi_web_api_config['auth_token'],
|
||||||
|
logger=logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_minio_repository.assert_called_once_with(
|
||||||
|
endpoint=minio_config['endpoint_url'],
|
||||||
|
access_key=minio_config['access_key'],
|
||||||
|
secret_key=minio_config['secret_key'],
|
||||||
|
bucket=minio_config['default_bucket'],
|
||||||
|
logger=logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller.return_value,
|
||||||
|
secure=minio_config['secure'],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@patch('laborious.activities.activities.Storage')
|
||||||
@patch('laborious.activities.activities.Postgres', return_value=MagicMock())
|
@patch('laborious.activities.activities.MLFlow')
|
||||||
@patch('laborious.activities.activities.MLFlow', return_value=MagicMock())
|
@patch('laborious.activities.activities.OPC')
|
||||||
@patch('laborious.activities.activities.OPC', return_value=MagicMock())
|
@patch('laborious.activities.activities.Gates')
|
||||||
async def test_shutdown(mock_opc_init,
|
@patch('laborious.activities.activities.ModelMetrics')
|
||||||
_mock_mlflow_init, mock_postgres_init):
|
@patch('laborious.activities.activities.API')
|
||||||
|
@patch('laborious.activities.activities.MinioRepository')
|
||||||
|
def test_shutdown(
|
||||||
|
_mock_minio_repository,
|
||||||
|
mock_api_init,
|
||||||
|
mock_model_metrics_init,
|
||||||
|
mock_gates_init,
|
||||||
|
mock_opc_init,
|
||||||
|
mock_mlflow_init,
|
||||||
|
mock_storage_init,
|
||||||
|
):
|
||||||
|
mock_opc_init.close = MagicMock()
|
||||||
postgres_config = {
|
postgres_config = {
|
||||||
'host': 'localhost',
|
'host': 'localhost',
|
||||||
'port': 5432,
|
'port': 5432,
|
||||||
@@ -103,20 +178,31 @@ async def test_shutdown(mock_opc_init,
|
|||||||
'password': 'postgres',
|
'password': 'postgres',
|
||||||
'dbname': 'postgres',
|
'dbname': 'postgres',
|
||||||
'min_connections': 1,
|
'min_connections': 1,
|
||||||
'max_connections': 10
|
'max_connections': 10,
|
||||||
}
|
}
|
||||||
|
|
||||||
mlflow_config = {
|
minio_config = {
|
||||||
'host': 'localhost',
|
'endpoint_url': 'localhost:9000',
|
||||||
'port': 5000,
|
'access_key': 'minio',
|
||||||
'username': 'mlflow',
|
'secret_key': 'minio123',
|
||||||
'password': 'mlflow'
|
'default_bucket': 'test',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
mlflow_repository = MagicMock()
|
||||||
|
plugin_store = MagicMock()
|
||||||
|
|
||||||
opc_config = {
|
opc_config = {
|
||||||
'bootstrap_servers': 'localhost:9092',
|
'bootstrap_servers': 'localhost:9092',
|
||||||
'polling_time': 1000,
|
'polling_time': 1000,
|
||||||
'group_id': 'test-group'
|
'group_id': 'test-group',
|
||||||
|
}
|
||||||
|
|
||||||
|
pi_web_api_config = {
|
||||||
|
'base_url': 'https://test-pi-server.com',
|
||||||
|
'auth_type': 'bearer',
|
||||||
|
'auth_token': 'test_token',
|
||||||
}
|
}
|
||||||
|
|
||||||
logger = MagicMock()
|
logger = MagicMock()
|
||||||
@@ -124,12 +210,90 @@ async def test_shutdown(mock_opc_init,
|
|||||||
|
|
||||||
activities = Activities(
|
activities = Activities(
|
||||||
postgres_config=postgres_config,
|
postgres_config=postgres_config,
|
||||||
mlflow_config=mlflow_config,
|
plugin_store=plugin_store,
|
||||||
|
minio_config=minio_config,
|
||||||
opc_config=opc_config,
|
opc_config=opc_config,
|
||||||
|
pi_web_api_config=pi_web_api_config,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler
|
notification_handler=notification_handler,
|
||||||
|
mlflow_repository=mlflow_repository,
|
||||||
)
|
)
|
||||||
|
|
||||||
await activities.shutdown()
|
activities.shutdown()
|
||||||
mock_opc_init.shutdown.assert_called_once()
|
mock_opc_init.close.assert_called_once()
|
||||||
mock_postgres_init.close.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,
|
||||||
|
)
|
||||||
|
|||||||
503
tests/laborious/activities/test_api.py
Normal file
503
tests/laborious/activities/test_api.py
Normal file
@@ -0,0 +1,503 @@
|
|||||||
|
from unittest.mock import ANY, MagicMock, call, patch
|
||||||
|
|
||||||
|
from pytest import fixture
|
||||||
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
|
|
||||||
|
from laborious.activities.api import API, PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
'metadata': {
|
||||||
|
'model_id': 'test_model',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'schema_name': 'test_schedule',
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _create_mock_dataframe(to_dict_return=None):
|
||||||
|
"""Helper function to create a mocked DataFrame for testing."""
|
||||||
|
mock_df = MagicMock()
|
||||||
|
mock_head = MagicMock()
|
||||||
|
|
||||||
|
def get_column_values(key):
|
||||||
|
if key == 'prediction':
|
||||||
|
return MagicMock(values=[0.75])
|
||||||
|
elif key == 'prediction_confidence':
|
||||||
|
return MagicMock(values=[0.95])
|
||||||
|
else:
|
||||||
|
return MagicMock(values=['2024-01-01T00:00:00+00:00'])
|
||||||
|
|
||||||
|
mock_head.__getitem__.side_effect = get_column_values
|
||||||
|
mock_df.head.return_value = mock_head
|
||||||
|
|
||||||
|
if to_dict_return is None:
|
||||||
|
to_dict_return = {
|
||||||
|
'prediction': [0.75],
|
||||||
|
'prediction_confidence': [0.95],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
}
|
||||||
|
mock_df.to_dict.return_value = to_dict_return
|
||||||
|
|
||||||
|
return mock_df
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def base_input_data():
|
||||||
|
"""Base input data for PI Web API tests."""
|
||||||
|
return {
|
||||||
|
**metadata,
|
||||||
|
'data': {
|
||||||
|
'prediction': [0.75],
|
||||||
|
'prediction_confidence': [0.95],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
},
|
||||||
|
'pi_web_api_output_config': {
|
||||||
|
'endpoint': 'https://test-pi-server.com/piwebapi',
|
||||||
|
'prediction_tags': {'tag1': 'web_id_1'},
|
||||||
|
'confidence_tags': {'tag2': 'web_id_2'},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.PIWebAPIClient')
|
||||||
|
def test_get_pi_web_api_core_labels_without_operation_type(mock_pi_web_api_client):
|
||||||
|
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||||
|
|
||||||
|
api_instance = API(
|
||||||
|
base_url='https://test-pi-server.com',
|
||||||
|
auth_type='bearer',
|
||||||
|
auth_token='test_token',
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
with patch.object(
|
||||||
|
SientiaMonitoring,
|
||||||
|
'get_core_labels',
|
||||||
|
return_value={
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'operation_type': 'write_pi_web_api_data',
|
||||||
|
},
|
||||||
|
):
|
||||||
|
labels = api_instance.get_pi_web_api_core_labels(metadata=metadata['metadata'])
|
||||||
|
assert labels['operation_type'] == 'write_pi_web_api_data'
|
||||||
|
assert labels == {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'operation_type': 'write_pi_web_api_data',
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.PIWebAPIClient')
|
||||||
|
def test_get_pi_web_api_core_labels_with_operation_type(mock_pi_web_api_client):
|
||||||
|
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||||
|
|
||||||
|
api_instance = API(
|
||||||
|
base_url='https://test-pi-server.com',
|
||||||
|
auth_type='bearer',
|
||||||
|
auth_token='test_token',
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
with patch.object(
|
||||||
|
SientiaMonitoring,
|
||||||
|
'get_core_labels',
|
||||||
|
return_value={
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'k8s',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'operation_type': 'write',
|
||||||
|
},
|
||||||
|
):
|
||||||
|
labels = api_instance.get_pi_web_api_core_labels(
|
||||||
|
metadata=metadata['metadata'], operation_type='write'
|
||||||
|
)
|
||||||
|
assert labels['operation_type'] == 'write'
|
||||||
|
assert labels['runtime'] == 'k8s'
|
||||||
|
|
||||||
|
|
||||||
|
def test__init__():
|
||||||
|
api = API(
|
||||||
|
base_url='https://test-pi-server.com',
|
||||||
|
auth_type='bearer',
|
||||||
|
auth_token='test_token',
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert api.pi_web_api_client is not None
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
@patch('laborious.activities.api.PIWebAPIClient')
|
||||||
|
def api(mock_pi_web_api_client):
|
||||||
|
mock_client = MagicMock()
|
||||||
|
mock_client.write_value = MagicMock()
|
||||||
|
mock_client.close = MagicMock()
|
||||||
|
mock_client.base_url = 'https://test-pi-server.com'
|
||||||
|
mock_pi_web_api_client.return_value = mock_client
|
||||||
|
|
||||||
|
api_instance = API(
|
||||||
|
base_url='https://test-pi-server.com',
|
||||||
|
auth_type='bearer',
|
||||||
|
auth_token='test_token',
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
api_instance.send_notification = MagicMock()
|
||||||
|
api_instance.info = MagicMock()
|
||||||
|
api_instance.error = MagicMock()
|
||||||
|
api_instance.emit_metric_sync = MagicMock()
|
||||||
|
api_instance.get_core_labels = MagicMock(
|
||||||
|
return_value={
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return api_instance
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.DataFrame')
|
||||||
|
def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
|
||||||
|
input_data = {
|
||||||
|
**base_input_data,
|
||||||
|
'pi_web_api_output_config': {
|
||||||
|
'endpoint': 'https://test-pi-server.com/piwebapi',
|
||||||
|
'prediction_tags': {'tag1': 'web_id_1', 'tag2': 'web_id_2'},
|
||||||
|
'confidence_tags': {'tag3': 'web_id_3', 'tag4': 'web_id_4'},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mock_dataframe.return_value = _create_mock_dataframe()
|
||||||
|
|
||||||
|
# Mock successful responses
|
||||||
|
api.pi_web_api_client.write_value.side_effect = [
|
||||||
|
[{'WebId': 'web_id_1', 'Errors': []}, {'WebId': 'web_id_2', 'Errors': []}],
|
||||||
|
[{'WebId': 'web_id_3', 'Errors': []}, {'WebId': 'web_id_4', 'Errors': []}],
|
||||||
|
]
|
||||||
|
|
||||||
|
result = api.write_pi_web_api_data(input_data)
|
||||||
|
|
||||||
|
api.pi_web_api_client.write_value.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_1', 'web_id_2'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01T00:00:00+00:00',
|
||||||
|
'Value': 0.75,
|
||||||
|
},
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
web_ids=['web_id_3', 'web_id_4'],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01T00:00:00+00:00',
|
||||||
|
'Value': 0.95,
|
||||||
|
},
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == {
|
||||||
|
'prediction': [0.75],
|
||||||
|
'prediction_confidence': [0.95],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.DataFrame')
|
||||||
|
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],
|
||||||
|
'prediction_confidence': [PI_WEB_API_PREDICTION_ERROR_CONFIDENCE],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
api.pi_web_api_client.write_value.side_effect = Exception('Prediction write failed')
|
||||||
|
|
||||||
|
result = api.write_pi_web_api_data(base_input_data)
|
||||||
|
|
||||||
|
api.send_notification.assert_called_once_with(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR',
|
||||||
|
message="Error writing prediction data to PI Web API: Prediction write failed\n Tags: {'tag1': 'web_id_1'}",
|
||||||
|
block='write_pi_web_api_data',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result['prediction_confidence'][0] == PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
assert api.pi_web_api_client.write_value.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.DataFrame')
|
||||||
|
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
|
||||||
|
api.pi_web_api_client.write_value.side_effect = [
|
||||||
|
[{'WebId': 'web_id_1', 'Errors': []}],
|
||||||
|
Exception('Confidence write failed'),
|
||||||
|
]
|
||||||
|
|
||||||
|
result = api.write_pi_web_api_data(base_input_data)
|
||||||
|
|
||||||
|
api.send_notification.assert_called_once_with(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR',
|
||||||
|
message="Error writing confidence data to PI Web API: Confidence write failed\n Tags: {'tag2': 'web_id_2'}",
|
||||||
|
block='write_pi_web_api_data',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == {
|
||||||
|
'prediction': [0.75],
|
||||||
|
'prediction_confidence': [0.95],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
}
|
||||||
|
assert api.pi_web_api_client.write_value.call_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.api.DataFrame')
|
||||||
|
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': {
|
||||||
|
'endpoint': 'https://test-pi-server.com/piwebapi',
|
||||||
|
'prediction_tags': {},
|
||||||
|
'confidence_tags': {},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mock_dataframe.return_value = _create_mock_dataframe()
|
||||||
|
|
||||||
|
# Mock empty responses
|
||||||
|
api.pi_web_api_client.write_value.side_effect = [
|
||||||
|
[],
|
||||||
|
[],
|
||||||
|
]
|
||||||
|
|
||||||
|
result = api.write_pi_web_api_data(input_data)
|
||||||
|
|
||||||
|
api.pi_web_api_client.write_value.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
web_ids=[],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01T00:00:00+00:00',
|
||||||
|
'Value': 0.75,
|
||||||
|
},
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
web_ids=[],
|
||||||
|
value={
|
||||||
|
'Timestamp': '2024-01-01T00:00:00+00:00',
|
||||||
|
'Value': 0.95,
|
||||||
|
},
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == {
|
||||||
|
'prediction': [0.75],
|
||||||
|
'prediction_confidence': [0.95],
|
||||||
|
'timestamp': ['2024-01-01T00:00:00+00:00'],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@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):
|
||||||
|
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):
|
||||||
|
"""Test successful processing of PI Web API response with all tags written."""
|
||||||
|
response_data = [
|
||||||
|
{'WebId': 'web_id_1', 'Errors': []},
|
||||||
|
{'WebId': 'web_id_2', 'Errors': []},
|
||||||
|
]
|
||||||
|
tags = {'tag1': 'web_id_1', 'tag2': 'web_id_2'}
|
||||||
|
core_labels = {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
|
||||||
|
confidence, message = api.process_pi_web_api_response(
|
||||||
|
response_data=response_data,
|
||||||
|
tags=tags,
|
||||||
|
core_labels=core_labels,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
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 len(call_args_list) == 2
|
||||||
|
# Check that all calls include core_labels and tag_name
|
||||||
|
for call_args in call_args_list:
|
||||||
|
assert 'tag_name' in call_args.kwargs['tags']
|
||||||
|
assert call_args.kwargs['tags']['tag_name'] in ['tag1', 'tag2']
|
||||||
|
|
||||||
|
|
||||||
|
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']},
|
||||||
|
{'WebId': 'web_id_2', 'Errors': []},
|
||||||
|
]
|
||||||
|
tags = {'tag1': 'web_id_1', 'tag2': 'web_id_2'}
|
||||||
|
core_labels = {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
|
||||||
|
confidence, message = api.process_pi_web_api_response(
|
||||||
|
response_data=response_data,
|
||||||
|
tags=tags,
|
||||||
|
core_labels=core_labels,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert confidence == PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
assert (
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
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': []},
|
||||||
|
]
|
||||||
|
tags = {'tag1': 'web_id_1', 'tag2': 'web_id_2'}
|
||||||
|
core_labels = {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
|
||||||
|
confidence, message = api.process_pi_web_api_response(
|
||||||
|
response_data=response_data,
|
||||||
|
tags=tags,
|
||||||
|
core_labels=core_labels,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert confidence == PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
assert (
|
||||||
|
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
|
||||||
|
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):
|
||||||
|
"""Test processing response when WebId is missing in response item."""
|
||||||
|
response_data = [
|
||||||
|
{'Errors': []},
|
||||||
|
{'WebId': 'web_id_2', 'Errors': []},
|
||||||
|
]
|
||||||
|
tags = {'tag1': 'web_id_1', 'tag2': 'web_id_2'}
|
||||||
|
core_labels = {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
|
||||||
|
confidence, message = api.process_pi_web_api_response(
|
||||||
|
response_data=response_data,
|
||||||
|
tags=tags,
|
||||||
|
core_labels=core_labels,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert confidence == PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
assert (
|
||||||
|
message
|
||||||
|
== "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag2'] tags were written."
|
||||||
|
)
|
||||||
|
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):
|
||||||
|
"""Test processing response when tag name is not found for WebId."""
|
||||||
|
response_data = [
|
||||||
|
{'WebId': 'unknown_web_id', 'Errors': []},
|
||||||
|
]
|
||||||
|
tags = {'tag1': 'web_id_1'}
|
||||||
|
core_labels = {
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'runtime': 'local',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
|
||||||
|
confidence, message = api.process_pi_web_api_response(
|
||||||
|
response_data=response_data,
|
||||||
|
tags=tags,
|
||||||
|
core_labels=core_labels,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert confidence == PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||||
|
assert (
|
||||||
|
message
|
||||||
|
== "The number of written tags does not match the number of tag names: Expected ['tag1'] tags, but [] tags were written."
|
||||||
|
)
|
||||||
|
api.error.assert_any_call(
|
||||||
|
'The response did not contain the tag name for WebId unknown_web_id', metadata['metadata']
|
||||||
|
)
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
954
tests/laborious/activities/test_model_metrics.py
Normal file
954
tests/laborious/activities/test_model_metrics.py
Normal file
@@ -0,0 +1,954 @@
|
|||||||
|
from unittest.mock import ANY, MagicMock, patch
|
||||||
|
|
||||||
|
from pandas import DataFrame, Timestamp
|
||||||
|
from pytest import fixture, raises
|
||||||
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
|
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||||
|
from sientia_model.analytics.drift_analysis import DriftInsufficientDataError
|
||||||
|
|
||||||
|
from laborious.activities.model_metrics import ModelMetrics
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def model_metrics_activity():
|
||||||
|
model_metrics = ModelMetrics(
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
model_metrics.error = MagicMock()
|
||||||
|
model_metrics.debug = MagicMock()
|
||||||
|
model_metrics.info = MagicMock()
|
||||||
|
model_metrics.warning = MagicMock()
|
||||||
|
model_metrics.critical = MagicMock()
|
||||||
|
model_metrics.send_notification = MagicMock()
|
||||||
|
model_metrics.emit_metric_sync = MagicMock()
|
||||||
|
model_metrics.get_core_labels = MagicMock(
|
||||||
|
return_value={
|
||||||
|
'pod_id': 'test_pod',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
}
|
||||||
|
)
|
||||||
|
model_metrics.observe_lag_sync = MagicMock()
|
||||||
|
model_metrics.pod_id = 'test_pod'
|
||||||
|
return model_metrics
|
||||||
|
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
'metadata': {
|
||||||
|
'model_id': 'test_model',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'schema_name': 'test_schedule',
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': None,
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'invalid',
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act & Assert
|
||||||
|
try:
|
||||||
|
model_metrics_activity.calculate_drift(input_data)
|
||||||
|
except ValueError as e:
|
||||||
|
assert str(e) == 'Invalid chunk period: invalid, must be "min" or "s"'
|
||||||
|
model_metrics_activity.error.assert_called_once_with(
|
||||||
|
'Invalid chunk period: invalid', metadata['metadata']
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise AssertionError('Expected ValueError')
|
||||||
|
|
||||||
|
|
||||||
|
def _sample_drift_metrics_df(ts: Timestamp) -> DataFrame:
|
||||||
|
"""Minimal analyzer-shaped dataframe (univariate row + columns the activity expects)."""
|
||||||
|
return DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': [ts],
|
||||||
|
'feature': ['feature1'],
|
||||||
|
'method': ['ks_test'],
|
||||||
|
'value': [0.5],
|
||||||
|
'alert': [False],
|
||||||
|
'chunk_index': [0],
|
||||||
|
'chunk_start_date': [ts],
|
||||||
|
'chunk_end_date': [ts],
|
||||||
|
'threshold': [0.1],
|
||||||
|
'drift_type': ['univariate'],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_with_reference_data(model_metrics_activity):
|
||||||
|
ts = Timestamp('2023-05-26 11:12:27')
|
||||||
|
drift_df = _sample_drift_metrics_df(ts)
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=drift_df)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict('list'),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
expected_timestamp = ts.tz_localize('UTC').strftime(DATETIME_FORMAT_WITH_TZ)
|
||||||
|
assert result == [
|
||||||
|
{
|
||||||
|
'timestamp': expected_timestamp,
|
||||||
|
'feature': 'feature1',
|
||||||
|
'method': 'ks_test',
|
||||||
|
'value': 0.5,
|
||||||
|
'alert': False,
|
||||||
|
'chunk_index': 0,
|
||||||
|
'chunk_start_date': ts.isoformat(),
|
||||||
|
'chunk_end_date': ts.isoformat(),
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'accurate': True,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
model_metrics_activity.info.assert_called()
|
||||||
|
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_without_reference_data(model_metrics_activity):
|
||||||
|
# Ten rows so int(len * 0.3) >= 1 for the built-in reference slice.
|
||||||
|
ts_last = Timestamp('2023-05-26 11:12:36')
|
||||||
|
drift_df = _sample_drift_metrics_df(ts_last)
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=drift_df)
|
||||||
|
|
||||||
|
timestamps = [f'2023-05-26 11:12:{27 + i:02d}' for i in range(10)]
|
||||||
|
target_data_dict = {
|
||||||
|
'timestamp': timestamps,
|
||||||
|
'variable': ['feature1'] * 10,
|
||||||
|
'value': [float(i) for i in range(10)],
|
||||||
|
}
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': None,
|
||||||
|
'target_data': target_data_dict,
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 's',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
expected_timestamp = ts_last.tz_localize('UTC').strftime(DATETIME_FORMAT_WITH_TZ)
|
||||||
|
assert result == [
|
||||||
|
{
|
||||||
|
'timestamp': expected_timestamp,
|
||||||
|
'feature': 'feature1',
|
||||||
|
'method': 'ks_test',
|
||||||
|
'value': 0.5,
|
||||||
|
'alert': False,
|
||||||
|
'chunk_index': 0,
|
||||||
|
'chunk_start_date': ts_last.isoformat(),
|
||||||
|
'chunk_end_date': ts_last.isoformat(),
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'accurate': False,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
model_metrics_activity.warning.assert_called()
|
||||||
|
model_metrics_activity.send_notification.assert_called_once_with(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING',
|
||||||
|
message='Using 30% first rows of target data as reference data',
|
||||||
|
block='model_metrics',
|
||||||
|
level=NotificationLevel.WARNING,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_empty_drift_df(model_metrics_activity):
|
||||||
|
"""Empty analyzer merge yields no rows and no insufficient-data alert (lib owns that failure mode)."""
|
||||||
|
ts = Timestamp('2023-05-26 11:12:27')
|
||||||
|
empty_df = _sample_drift_metrics_df(ts).iloc[0:0]
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=empty_df)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict(),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
assert result == []
|
||||||
|
model_metrics_activity.send_notification.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_empty_after_timestamp_filter(model_metrics_activity):
|
||||||
|
"""Rows dropped by target-window alignment yield an empty export list, not an insufficient-data error."""
|
||||||
|
drift_df = _sample_drift_metrics_df(Timestamp('2020-01-01'))
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=drift_df)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict(),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
assert result == []
|
||||||
|
model_metrics_activity.send_notification.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.DataFrame')
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
def test_calculate_drift_drift_insufficient_data_error_from_lib(
|
||||||
|
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||||
|
):
|
||||||
|
"""``DriftInsufficientDataError`` maps to MODEL_METRICS_DRIFT_INSUFFICIENT_DATA, not GET error."""
|
||||||
|
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||||
|
|
||||||
|
lib_msg = (
|
||||||
|
'[MODEL_METRICS_DRIFT_INSUFFICIENT_DATA] Drift analysis produced no time chunks '
|
||||||
|
"(chunk_period='min', analysis_rows=1)."
|
||||||
|
)
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(
|
||||||
|
side_effect=DriftInsufficientDataError(lib_msg, analysis_rows=1)
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_target_df = MagicMock()
|
||||||
|
mock_target_df.pivot.return_value = mock_target_df
|
||||||
|
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||||
|
mock_target_df.reset_index.return_value = mock_target_df
|
||||||
|
mock_target_df.dropna.return_value = mock_target_df
|
||||||
|
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||||
|
mock_target_df.drop.return_value.columns = ['feature1']
|
||||||
|
mock_dataframe.return_value = mock_target_df
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict(),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
with raises(DriftInsufficientDataError, match='MODEL_METRICS_DRIFT_INSUFFICIENT_DATA'):
|
||||||
|
model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
model_metrics_activity.error.assert_called_once_with(lib_msg, metadata['metadata'])
|
||||||
|
model_metrics_activity.send_notification.assert_called_once_with(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='MODEL_METRICS_DRIFT_INSUFFICIENT_DATA',
|
||||||
|
message=lib_msg,
|
||||||
|
block='model_metrics',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_success_min(model_metrics_activity):
|
||||||
|
ts = Timestamp('2023-05-26 11:12:27')
|
||||||
|
drift_df = _sample_drift_metrics_df(ts)
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=drift_df)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict('list'),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
expected_timestamp = ts.tz_localize('UTC').strftime(DATETIME_FORMAT_WITH_TZ)
|
||||||
|
assert result == [
|
||||||
|
{
|
||||||
|
'timestamp': expected_timestamp,
|
||||||
|
'feature': 'feature1',
|
||||||
|
'method': 'ks_test',
|
||||||
|
'value': 0.5,
|
||||||
|
'alert': False,
|
||||||
|
'chunk_index': 0,
|
||||||
|
'chunk_start_date': ts.isoformat(),
|
||||||
|
'chunk_end_date': ts.isoformat(),
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'accurate': True,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
model_metrics_activity.info.assert_called()
|
||||||
|
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_drift_success_s(model_metrics_activity):
|
||||||
|
ts = Timestamp('2023-05-26 11:12:27')
|
||||||
|
drift_df = _sample_drift_metrics_df(ts)
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(return_value=drift_df)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict('list'),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 's',
|
||||||
|
}
|
||||||
|
|
||||||
|
result = model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
expected_timestamp = ts.tz_localize('UTC').strftime(DATETIME_FORMAT_WITH_TZ)
|
||||||
|
assert result == [
|
||||||
|
{
|
||||||
|
'timestamp': expected_timestamp,
|
||||||
|
'feature': 'feature1',
|
||||||
|
'method': 'ks_test',
|
||||||
|
'value': 0.5,
|
||||||
|
'alert': False,
|
||||||
|
'chunk_index': 0,
|
||||||
|
'chunk_start_date': ts.isoformat(),
|
||||||
|
'chunk_end_date': ts.isoformat(),
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'accurate': True,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
model_metrics_activity.info.assert_called()
|
||||||
|
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.DataFrame')
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
def test_calculate_drift_get_drift_metrics_error(
|
||||||
|
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||||
|
):
|
||||||
|
# Arrange
|
||||||
|
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||||
|
|
||||||
|
model_metrics_activity.get_drift_metrics = MagicMock(
|
||||||
|
side_effect=Exception('Get drift metrics error')
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_target_df = MagicMock()
|
||||||
|
mock_target_df.pivot.return_value = mock_target_df
|
||||||
|
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||||
|
mock_target_df.reset_index.return_value = mock_target_df
|
||||||
|
mock_target_df.dropna.return_value = mock_target_df
|
||||||
|
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||||
|
mock_target_df.drop.return_value.columns = ['feature1']
|
||||||
|
mock_dataframe.return_value = mock_target_df
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'reference_data': reference_data.to_dict(),
|
||||||
|
'target_data': {
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'variable': ['feature1'],
|
||||||
|
'value': [1.0],
|
||||||
|
},
|
||||||
|
'target_name': 'target',
|
||||||
|
'drift_metrics': ['ks_test'],
|
||||||
|
'chunk_period': 'min',
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act / Assert
|
||||||
|
with raises(Exception, match='Get drift metrics error'):
|
||||||
|
model_metrics_activity.calculate_drift(input_data)
|
||||||
|
|
||||||
|
model_metrics_activity.error.assert_called_once_with(
|
||||||
|
'Error getting drift metrics: Get drift metrics error', metadata['metadata']
|
||||||
|
)
|
||||||
|
model_metrics_activity.send_notification.assert_called_once_with(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR',
|
||||||
|
message='Error getting drift metrics: Get drift metrics error',
|
||||||
|
block='model_metrics',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
@patch('laborious.activities.model_metrics.time.time')
|
||||||
|
@patch('laborious.activities.model_metrics.DriftAnalysis')
|
||||||
|
@patch('laborious.activities.model_metrics.metrics')
|
||||||
|
def test_get_drift_metrics_success(
|
||||||
|
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||||
|
):
|
||||||
|
# Arrange
|
||||||
|
mock_time.return_value = 1000.0
|
||||||
|
|
||||||
|
mock_drift_df = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'method': ['ks_test'],
|
||||||
|
'value': [0.5],
|
||||||
|
'feature': ['feature1'],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_model_analysis.return_value.detect_univariate_drift.return_value = MagicMock()
|
||||||
|
mock_model_analysis.return_value.detect_multivariate_drift.return_value = MagicMock()
|
||||||
|
mock_model_analysis.return_value.get_drift_metrics_dataframe.return_value = mock_drift_df
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_columns = reference_data.drop(
|
||||||
|
columns=['target', 'timestamp'], errors='ignore'
|
||||||
|
).columns
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = model_metrics_activity.get_drift_metrics(
|
||||||
|
reference_data=reference_data,
|
||||||
|
target_data=target_data,
|
||||||
|
target_name='target',
|
||||||
|
reference_columns=reference_columns,
|
||||||
|
drift_metrics=['ks_test'],
|
||||||
|
chunk_period='min',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert isinstance(result, DataFrame)
|
||||||
|
model_metrics_activity.debug.assert_called()
|
||||||
|
model_metrics_activity.observe_lag_sync.assert_called()
|
||||||
|
model_metrics_activity.emit_metric_sync.assert_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
@patch('laborious.activities.model_metrics.time.time')
|
||||||
|
@patch('laborious.activities.model_metrics.DriftAnalysis')
|
||||||
|
@patch('laborious.activities.model_metrics.metrics')
|
||||||
|
def test_get_drift_metrics_univariate_error(
|
||||||
|
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||||
|
):
|
||||||
|
# Arrange
|
||||||
|
mock_time.return_value = 1000.0
|
||||||
|
|
||||||
|
mock_model_analysis.return_value.detect_univariate_drift.side_effect = Exception(
|
||||||
|
'Univariate drift error'
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27'],
|
||||||
|
'target': [1.0],
|
||||||
|
'feature1': [1.0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_columns = reference_data.drop(
|
||||||
|
columns=['target', 'timestamp'], errors='ignore'
|
||||||
|
).columns
|
||||||
|
|
||||||
|
# Act & Assert
|
||||||
|
try:
|
||||||
|
model_metrics_activity.get_drift_metrics(
|
||||||
|
reference_data=reference_data,
|
||||||
|
target_data=target_data,
|
||||||
|
target_name='target',
|
||||||
|
reference_columns=reference_columns,
|
||||||
|
drift_metrics=['ks_test'],
|
||||||
|
chunk_period='min',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
assert str(e) == 'Univariate drift error'
|
||||||
|
model_metrics_activity.error.assert_called_once_with(
|
||||||
|
'Error detecting univariate drift: Univariate drift error', metadata['metadata']
|
||||||
|
)
|
||||||
|
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||||
|
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise AssertionError('Expected Exception')
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
@patch('laborious.activities.model_metrics.time.time')
|
||||||
|
@patch('laborious.activities.model_metrics.DriftAnalysis')
|
||||||
|
@patch('laborious.activities.model_metrics.metrics')
|
||||||
|
def test_get_drift_metrics_multivariate_error(
|
||||||
|
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||||
|
):
|
||||||
|
mock_time.return_value = 1000.0
|
||||||
|
|
||||||
|
mock_model_analysis.return_value.detect_univariate_drift.return_value = MagicMock()
|
||||||
|
mock_model_analysis.return_value.detect_multivariate_drift.side_effect = Exception(
|
||||||
|
'Multivariate drift error'
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{'timestamp': ['2023-05-26 11:12:27'], 'target': [1.0], 'feature1': [1.0]}
|
||||||
|
)
|
||||||
|
target_data = DataFrame(
|
||||||
|
{'timestamp': ['2023-05-26 11:12:27'], 'target': [1.0], 'feature1': [1.0]}
|
||||||
|
)
|
||||||
|
reference_columns = reference_data.drop(
|
||||||
|
columns=['target', 'timestamp'], errors='ignore'
|
||||||
|
).columns
|
||||||
|
|
||||||
|
try:
|
||||||
|
model_metrics_activity.get_drift_metrics(
|
||||||
|
reference_data=reference_data,
|
||||||
|
target_data=target_data,
|
||||||
|
target_name='target',
|
||||||
|
reference_columns=reference_columns,
|
||||||
|
drift_metrics=['ks_test'],
|
||||||
|
chunk_period='min',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
assert str(e) == 'Multivariate drift error'
|
||||||
|
model_metrics_activity.error.assert_called_once_with(
|
||||||
|
'Error detecting multivariate drift: Multivariate drift error', metadata['metadata']
|
||||||
|
)
|
||||||
|
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||||
|
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise AssertionError('Expected Exception')
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.model_metrics.to_datetime')
|
||||||
|
@patch('laborious.activities.model_metrics.time.time')
|
||||||
|
@patch('laborious.activities.model_metrics.DriftAnalysis')
|
||||||
|
@patch('laborious.activities.model_metrics.metrics')
|
||||||
|
def test_get_drift_metrics_dataframe_error(
|
||||||
|
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||||
|
):
|
||||||
|
mock_time.return_value = 1000.0
|
||||||
|
|
||||||
|
mock_model_analysis.return_value.detect_univariate_drift.return_value = MagicMock()
|
||||||
|
mock_model_analysis.return_value.detect_multivariate_drift.return_value = MagicMock()
|
||||||
|
mock_model_analysis.return_value.get_drift_metrics_dataframe.side_effect = Exception(
|
||||||
|
'Dataframe error'
|
||||||
|
)
|
||||||
|
|
||||||
|
reference_data = DataFrame(
|
||||||
|
{'timestamp': ['2023-05-26 11:12:27'], 'target': [1.0], 'feature1': [1.0]}
|
||||||
|
)
|
||||||
|
target_data = DataFrame(
|
||||||
|
{'timestamp': ['2023-05-26 11:12:27'], 'target': [1.0], 'feature1': [1.0]}
|
||||||
|
)
|
||||||
|
reference_columns = reference_data.drop(
|
||||||
|
columns=['target', 'timestamp'], errors='ignore'
|
||||||
|
).columns
|
||||||
|
|
||||||
|
try:
|
||||||
|
model_metrics_activity.get_drift_metrics(
|
||||||
|
reference_data=reference_data,
|
||||||
|
target_data=target_data,
|
||||||
|
target_name='target',
|
||||||
|
reference_columns=reference_columns,
|
||||||
|
drift_metrics=['ks_test'],
|
||||||
|
chunk_period='min',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
assert str(e) == 'Dataframe error'
|
||||||
|
model_metrics_activity.error.assert_called_once_with(
|
||||||
|
'Error building drift metrics dataframe: Dataframe error', metadata['metadata']
|
||||||
|
)
|
||||||
|
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||||
|
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise AssertionError('Expected Exception')
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||||
|
'target': [1.0, 2.0, 3.0],
|
||||||
|
'prediction': [1.1, 2.1, 2.9],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['rmse', 'mse', 'mae', 'r2'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 4
|
||||||
|
assert 'rmse' in result['metric'].values
|
||||||
|
assert 'mse' in result['metric'].values
|
||||||
|
assert 'mae' in result['metric'].values
|
||||||
|
assert 'r2' in result['metric'].values
|
||||||
|
assert all(model_id == 'test_model_id' for model_id in result['model_id'].values)
|
||||||
|
assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values)
|
||||||
|
assert all(data_size == 3 for data_size in result['data_size'].values)
|
||||||
|
assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values)
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['rmse', 'mse', 'mae', 'r2']",
|
||||||
|
metadata['metadata'],
|
||||||
|
)
|
||||||
|
model_metrics_activity.debug.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['rmse'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'rmse'
|
||||||
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
|
assert result['data_size'].values[0] == 2
|
||||||
|
assert result['interval_minutes'].values[0] == 5
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['rmse']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['mse'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'mse'
|
||||||
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
|
assert result['data_size'].values[0] == 2
|
||||||
|
assert result['interval_minutes'].values[0] == 5
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['mse']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['mae'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'mae'
|
||||||
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
|
assert result['data_size'].values[0] == 2
|
||||||
|
assert result['interval_minutes'].values[0] == 5
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['mae']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['r2'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'r2'
|
||||||
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
|
assert result['data_size'].values[0] == 2
|
||||||
|
assert result['interval_minutes'].values[0] == 5
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
# All target values are the same, so ss_tot will be 0
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 1.0],
|
||||||
|
'prediction': [1.1, 1.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['r2'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'r2'
|
||||||
|
assert result['value'].values[0] == 0.0 # Should return 0.0 when ss_tot == 0
|
||||||
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
|
assert result['data_size'].values[0] == 2
|
||||||
|
assert result['interval_minutes'].values[0] == 5
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity):
|
||||||
|
# Arrange
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||||
|
'target': [1.0, 2.0, 3.0],
|
||||||
|
'prediction': [1.1, 2.1, 2.9],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['rmse', 'mae'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert len(result['metric']) == 2
|
||||||
|
assert 'rmse' in result['metric'].values
|
||||||
|
assert 'mae' in result['metric'].values
|
||||||
|
assert all(model_id == 'test_model_id' for model_id in result['model_id'].values)
|
||||||
|
assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values)
|
||||||
|
assert all(data_size == 3 for data_size in result['data_size'].values)
|
||||||
|
assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values)
|
||||||
|
model_metrics_activity.info.assert_called_once_with(
|
||||||
|
"Calculating simple metrics for model test_model_id: ['rmse', 'mae']", metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity):
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['unknown_metric', 'rmse'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result['metric']) == 1
|
||||||
|
assert result['metric'].values[0] == 'rmse'
|
||||||
@@ -1,62 +1,75 @@
|
|||||||
from unittest.mock import patch, MagicMock, ANY, call, AsyncMock
|
from unittest.mock import ANY, MagicMock, call, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
from pandas import DataFrame
|
from pandas import DataFrame
|
||||||
from pytest import fixture, mark
|
from pytest import mark
|
||||||
import pytest_asyncio
|
|
||||||
from sientia_do.notifications.models import NotificationLevel
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
|
|
||||||
from laborious.activities.opc import OPC
|
from laborious.activities.opc import (
|
||||||
|
OPC,
|
||||||
|
OPC_COMMENT_SEPARATOR,
|
||||||
|
OPC_RECONNECT_IN_PROGRESS_COMMENT,
|
||||||
|
OPC_SESSION_BAD_COMMENT_PREFIX,
|
||||||
|
OPC_SESSION_BAD_CONFIDENCE,
|
||||||
|
OPC_WRITTING_ERROR_CONFIDENCE,
|
||||||
|
OPC_WRITTING_ERROR_MESSAGE,
|
||||||
|
_apply_opc_write_error,
|
||||||
|
)
|
||||||
|
|
||||||
metadata = {
|
metadata = {
|
||||||
"metadata": {
|
'metadata': {
|
||||||
"model_id": "test_model",
|
'model_id': 'test_model',
|
||||||
"model_name": "test_model",
|
'model_name': 'test_model',
|
||||||
"workflow_name": "test_workflow",
|
'workflow_name': 'test_workflow',
|
||||||
"schema_name": "test_schedule",
|
'schema_name': 'test_schedule',
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def test__init__():
|
def test__init__():
|
||||||
servers = {
|
servers = {'server1': {'id': 'server1'}}
|
||||||
'server1': 'config'
|
|
||||||
}
|
|
||||||
opc = OPC(
|
opc = OPC(
|
||||||
opc_servers=servers,
|
opc_servers=servers,
|
||||||
logger=MagicMock(),
|
logger=MagicMock(),
|
||||||
notification_handler=MagicMock()
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert opc.opc_servers == servers
|
assert opc.opc_servers == servers
|
||||||
assert opc.opc_repository == {}
|
assert opc.opc_repository == {}
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@patch('laborious.activities.opc.OpcRepository')
|
||||||
@patch("laborious.activities.opc.OpcRepository")
|
@patch('laborious.activities.opc.OPC.send_notification')
|
||||||
@patch("laborious.activities.opc.OPC.send_notification")
|
def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||||
async def test_init_opc(mock_send_notification, mock_opc_repository):
|
|
||||||
mock_logger = MagicMock()
|
mock_logger = MagicMock()
|
||||||
|
mock_metrics_controller = MagicMock()
|
||||||
server1 = MagicMock(
|
server1 = MagicMock(
|
||||||
connect=AsyncMock(return_value=(True, {})),
|
connect=MagicMock(return_value=(True, {})), write_data=MagicMock(return_value=(True, {}))
|
||||||
write_data=AsyncMock(return_value=(True, {}))
|
|
||||||
)
|
)
|
||||||
server2 = MagicMock(
|
server2 = MagicMock(
|
||||||
connect=AsyncMock(return_value=(True, {})),
|
connect=MagicMock(return_value=(True, {})), write_data=MagicMock(return_value=(True, {}))
|
||||||
write_data=AsyncMock(return_value=(True, {}))
|
|
||||||
)
|
)
|
||||||
server3 = MagicMock(
|
server3 = MagicMock(
|
||||||
connect=AsyncMock(return_value=(False, {
|
connect=MagicMock(
|
||||||
|
return_value=(
|
||||||
|
False,
|
||||||
|
{
|
||||||
'notification_id': 'OPC_CONNECTION_ERROR_server3',
|
'notification_id': 'OPC_CONNECTION_ERROR_server3',
|
||||||
'message': 'Failed to connect to OPC server: Test error',
|
'message': 'Failed to connect to OPC server: Test error',
|
||||||
'block': 'opc_repository',
|
'block': 'opc_repository',
|
||||||
'level': NotificationLevel.ERROR,
|
'level': NotificationLevel.ERROR,
|
||||||
'attachment_content': 'Test error'
|
'attachment_content': 'Test error',
|
||||||
})),
|
},
|
||||||
write_data=AsyncMock(return_value=(True, {}))
|
)
|
||||||
|
),
|
||||||
|
write_data=MagicMock(return_value=(True, {})),
|
||||||
)
|
)
|
||||||
mock_opc_repository.side_effect = [server1, server2, server3]
|
mock_opc_repository.side_effect = [server1, server2, server3]
|
||||||
mock_notification_handler = MagicMock()
|
mock_notification_handler = MagicMock()
|
||||||
servers = {
|
servers = {
|
||||||
'server1': {
|
'server1': {
|
||||||
|
'server_name': 'server1',
|
||||||
'id': 'server1',
|
'id': 'server1',
|
||||||
'url': 'http://localhost:8080',
|
'url': 'http://localhost:8080',
|
||||||
'server_uri': 'opc.tcp://localhost:4840',
|
'server_uri': 'opc.tcp://localhost:4840',
|
||||||
@@ -66,6 +79,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
|||||||
'reconnection_interval': 60,
|
'reconnection_interval': 60,
|
||||||
},
|
},
|
||||||
'server2': {
|
'server2': {
|
||||||
|
'server_name': 'server2',
|
||||||
'id': 'server2',
|
'id': 'server2',
|
||||||
'url': 'http://localhost:8080',
|
'url': 'http://localhost:8080',
|
||||||
'server_uri': 'opc.tcp://localhost:4840',
|
'server_uri': 'opc.tcp://localhost:4840',
|
||||||
@@ -75,6 +89,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
|||||||
'reconnection_interval': 60,
|
'reconnection_interval': 60,
|
||||||
},
|
},
|
||||||
'server3': {
|
'server3': {
|
||||||
|
'server_name': 'server3',
|
||||||
'id': 'server3',
|
'id': 'server3',
|
||||||
'url': 'http://localhost:8080',
|
'url': 'http://localhost:8080',
|
||||||
'server_uri': 'opc.tcp://localhost:4840',
|
'server_uri': 'opc.tcp://localhost:4840',
|
||||||
@@ -82,14 +97,15 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
|||||||
'private_key_path': '',
|
'private_key_path': '',
|
||||||
'server_cert_path': '',
|
'server_cert_path': '',
|
||||||
'reconnection_interval': 60,
|
'reconnection_interval': 60,
|
||||||
}
|
},
|
||||||
}
|
}
|
||||||
opc = OPC(
|
opc = OPC(
|
||||||
opc_servers=servers,
|
opc_servers=servers,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
await opc.init_opc()
|
opc.init_opc()
|
||||||
|
|
||||||
assert opc.opc_servers == servers
|
assert opc.opc_servers == servers
|
||||||
assert opc.logger == mock_logger
|
assert opc.logger == mock_logger
|
||||||
@@ -97,61 +113,70 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
|||||||
assert opc.opc_repository['server1'] == server1
|
assert opc.opc_repository['server1'] == server1
|
||||||
assert opc.opc_repository['server2'] == server2
|
assert opc.opc_repository['server2'] == server2
|
||||||
|
|
||||||
mock_opc_repository.assert_has_calls([
|
mock_opc_repository.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
id="server1",
|
opc_id='server1',
|
||||||
url="http://localhost:8080",
|
server_name='server1',
|
||||||
|
url='http://localhost:8080',
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
server_uri="opc.tcp://localhost:4840",
|
server_uri='opc.tcp://localhost:4840',
|
||||||
cert_path="",
|
cert_path='',
|
||||||
private_key_path="",
|
private_key_path='',
|
||||||
server_cert_path="",
|
server_cert_path='',
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
reconnection_interval=60,
|
reconnection_interval=60,
|
||||||
pod_id='localhost'
|
metrics_controller=mock_metrics_controller,
|
||||||
),
|
),
|
||||||
])
|
]
|
||||||
mock_opc_repository.assert_has_calls([
|
)
|
||||||
|
mock_opc_repository.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
id="server2",
|
opc_id='server2',
|
||||||
url="http://localhost:8080",
|
server_name='server2',
|
||||||
|
url='http://localhost:8080',
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
server_uri="opc.tcp://localhost:4840",
|
server_uri='opc.tcp://localhost:4840',
|
||||||
cert_path="",
|
cert_path='',
|
||||||
private_key_path="",
|
private_key_path='',
|
||||||
server_cert_path="",
|
server_cert_path='',
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
reconnection_interval=60,
|
reconnection_interval=60,
|
||||||
pod_id='localhost'
|
metrics_controller=mock_metrics_controller,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
server1.connect.assert_called_once()
|
server1.connect.assert_called_once()
|
||||||
server2.connect.assert_called_once()
|
server2.connect.assert_called_once()
|
||||||
|
|
||||||
mock_send_notification.assert_has_calls([
|
mock_send_notification.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
metadata={
|
metadata={
|
||||||
'model_id': '-',
|
'model_id': '-',
|
||||||
'model_name': '-',
|
'model_name': '-',
|
||||||
'workflow_name': '-',
|
'workflow_name': '-',
|
||||||
'schedule_name': 'INITIALIZATION'
|
'schedule_name': 'INITIALIZATION',
|
||||||
},
|
},
|
||||||
notification_id="OPC_CONNECTION_ERROR_server3",
|
notification_id='OPC_CONNECTION_ERROR_server3',
|
||||||
message="Failed to connect to OPC server: Test error",
|
message='Failed to connect to OPC server: Test error',
|
||||||
block="opc_repository",
|
block='opc_repository',
|
||||||
level=NotificationLevel.ERROR,
|
level=NotificationLevel.ERROR,
|
||||||
attachment_content=ANY
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture
|
@pytest.fixture
|
||||||
@patch("laborious.activities.opc.OpcRepository")
|
@patch('laborious.activities.opc.OpcRepository')
|
||||||
async def opc(mock_opc_repository):
|
def opc(mock_opc_repository):
|
||||||
servers = {
|
servers = {
|
||||||
'server1': {
|
'server1': {
|
||||||
'id': 'server1',
|
'id': 'server1',
|
||||||
|
'server_name': 'server1',
|
||||||
'url': 'http://localhost:8080',
|
'url': 'http://localhost:8080',
|
||||||
'server_uri': 'opc.tcp://localhost:4840',
|
'server_uri': 'opc.tcp://localhost:4840',
|
||||||
'cert_path': '',
|
'cert_path': '',
|
||||||
@@ -161,19 +186,17 @@ async def opc(mock_opc_repository):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
mock_opc_repository.return_value.write_data = AsyncMock(
|
mock_opc_repository.return_value.write_data = MagicMock(return_value=(True, {}))
|
||||||
return_value=(True, {})
|
mock_opc_repository.return_value.connect = MagicMock(return_value=(True, {}))
|
||||||
)
|
|
||||||
mock_opc_repository.return_value.connect = AsyncMock(
|
|
||||||
return_value=(True, {})
|
|
||||||
)
|
|
||||||
opc = OPC(
|
opc = OPC(
|
||||||
opc_servers=servers,
|
opc_servers=servers,
|
||||||
logger=MagicMock(),
|
logger=MagicMock(),
|
||||||
notification_handler=MagicMock()
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
)
|
)
|
||||||
await opc.init_opc()
|
opc.init_opc()
|
||||||
opc.send_notification = MagicMock()
|
opc.send_notification = MagicMock()
|
||||||
|
opc.emit_metric_sync = MagicMock()
|
||||||
return opc
|
return opc
|
||||||
|
|
||||||
|
|
||||||
@@ -186,184 +209,512 @@ WRITE_DATA_CASES = [
|
|||||||
|
|
||||||
|
|
||||||
@mark.parametrize('tag,data_type,data', WRITE_DATA_CASES)
|
@mark.parametrize('tag,data_type,data', WRITE_DATA_CASES)
|
||||||
@mark.asyncio
|
def test_write_data_success(opc, tag, data_type, data):
|
||||||
async def test_write_data_success(opc, tag, data_type, data):
|
opc.opc_repository['server1'].write_data.return_value = (True, {'response_time': 0.1})
|
||||||
result = await opc.write_data(server_id='server1', tag=tag, data=data,
|
|
||||||
data_type=data_type, tag_type='prediction', metadata=metadata)
|
response_time, error_info = opc.write_data(
|
||||||
assert result is True
|
server_id='server1',
|
||||||
opc.opc_repository['server1'].write_data.assert_called_once_with(
|
tag=tag,
|
||||||
tag, data, data_type, opc.logger, metadata)
|
data=data,
|
||||||
|
data_type=data_type,
|
||||||
|
tag_type='prediction',
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
assert response_time == 0.1
|
||||||
|
assert error_info is None
|
||||||
|
opc.opc_repository['server1'].write_data.assert_called_once_with(tag, data, data_type, metadata)
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
def test_write_data_failed(opc):
|
||||||
async def test_write_data_failed(opc):
|
opc.opc_repository['server1'].write_data.return_value = (
|
||||||
opc.opc_repository['server1'].write_data.return_value = (False, {
|
False,
|
||||||
|
{
|
||||||
'notification_id': 'OPC_WRITE_DATA_ERROR_server1',
|
'notification_id': 'OPC_WRITE_DATA_ERROR_server1',
|
||||||
'message': 'Failed to write data to OPC server: Test error',
|
'message': 'Failed to write data to OPC server: Test error',
|
||||||
'block': 'opc_repository',
|
'block': 'opc_repository',
|
||||||
'level': NotificationLevel.ERROR,
|
'level': NotificationLevel.ERROR,
|
||||||
'attachment_content': 'Test error'
|
'attachment_content': 'Test error',
|
||||||
})
|
},
|
||||||
|
)
|
||||||
|
|
||||||
result = await opc.write_data(server_id='server1', tag='tag1', data=50,
|
response_time, error_info = opc.write_data(
|
||||||
data_type='int', tag_type='prediction', metadata=metadata)
|
server_id='server1',
|
||||||
assert result is False
|
tag='tag1',
|
||||||
|
data=50,
|
||||||
|
data_type='int',
|
||||||
|
tag_type='prediction',
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
assert response_time is None
|
||||||
|
assert error_info is not None
|
||||||
|
|
||||||
opc.send_notification.assert_called_once_with(
|
opc.send_notification.assert_called_once_with(
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
notification_id="OPC_WRITE_DATA_ERROR_server1",
|
notification_id='OPC_WRITE_DATA_ERROR_server1',
|
||||||
message="Failed to write data to OPC server: Test error",
|
message='Failed to write data to OPC server: Test error',
|
||||||
block="opc_repository",
|
block='opc_repository',
|
||||||
level=NotificationLevel.ERROR,
|
level=NotificationLevel.ERROR,
|
||||||
attachment_content=ANY
|
attachment_content=ANY,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
def test_write_data_exception(opc):
|
||||||
async def test_write_data_exception(opc):
|
opc.opc_repository['server1'].write_data.side_effect = Exception('Test error')
|
||||||
opc.opc_repository['server1'].write_data.side_effect = Exception(
|
|
||||||
"Test error")
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await opc.write_data(server_id='server1', tag='tag1', data=50,
|
opc.write_data(
|
||||||
data_type='int', tag_type='prediction', metadata=metadata)
|
server_id='server1',
|
||||||
|
tag='tag1',
|
||||||
|
data=50,
|
||||||
|
data_type='int',
|
||||||
|
tag_type='prediction',
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
|
||||||
except Exception:
|
except Exception:
|
||||||
opc.send_notification.assert_called_once_with(
|
opc.send_notification.assert_called_once_with(
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
notification_id="WRITE_OPC_PREDICTION_ERROR",
|
notification_id='WRITE_OPC_PREDICTION_ERROR',
|
||||||
message="Error writing data to OPC server: Test error",
|
message='Error writing data to OPC server: Test error',
|
||||||
block="write_opc_data",
|
block='write_opc_data',
|
||||||
level=NotificationLevel.ERROR,
|
level=NotificationLevel.ERROR,
|
||||||
attachment_content=ANY
|
attachment_content=ANY,
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
assert False, "Expected an exception to be raised"
|
raise AssertionError('Expected an exception to be raised')
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@mark.parametrize(
|
||||||
async def test_write_opc_data_success(opc):
|
'error_info,initial_seen,initial_status,initial_reconnect,expected',
|
||||||
# Arrange
|
[
|
||||||
input_data = {
|
(None, False, None, False, (False, None, False)),
|
||||||
**metadata,
|
({}, False, None, False, (False, None, False)),
|
||||||
'data': {
|
(
|
||||||
'prediction': [0.75],
|
{'opc_error_kind': 'session_bad', 'opc_status': 'BadSessionIdInvalid'},
|
||||||
'prediction_confidence': [0.95]
|
False,
|
||||||
},
|
None,
|
||||||
'opc_output_config': {
|
False,
|
||||||
'server1': {
|
(True, 'BadSessionIdInvalid', False),
|
||||||
'prediction_tags': {
|
),
|
||||||
'tag1': {'data_type': 'float'}
|
(
|
||||||
},
|
{'opc_error_kind': 'session_bad', 'opc_status': 'NewStatus'},
|
||||||
'confidence_tags': {
|
True,
|
||||||
'tag2': {'data_type': 'float'}
|
'OldStatus',
|
||||||
}
|
False,
|
||||||
}
|
(True, 'NewStatus', False),
|
||||||
}
|
),
|
||||||
}
|
(
|
||||||
|
{'opc_error_kind': 'session_bad'},
|
||||||
|
True,
|
||||||
|
'KeptStatus',
|
||||||
|
False,
|
||||||
|
(True, 'KeptStatus', False),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{'opc_error_kind': 'reconnect_in_progress'},
|
||||||
|
False,
|
||||||
|
None,
|
||||||
|
False,
|
||||||
|
(False, None, True),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{'opc_error_kind': 'other'},
|
||||||
|
True,
|
||||||
|
'Status',
|
||||||
|
True,
|
||||||
|
(True, 'Status', True),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_apply_opc_write_error(
|
||||||
|
error_info, initial_seen, initial_status, initial_reconnect, expected
|
||||||
|
):
|
||||||
|
result = _apply_opc_write_error(
|
||||||
|
error_info,
|
||||||
|
initial_seen,
|
||||||
|
initial_status,
|
||||||
|
initial_reconnect,
|
||||||
|
)
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
# Act
|
|
||||||
opc.write_data = AsyncMock(return_value=True)
|
|
||||||
opc.process_confidence = MagicMock(return_value={'data': 'data'})
|
|
||||||
output = await opc.write_opc_data(input_data)
|
|
||||||
|
|
||||||
# Assert
|
def test_write_tags_from_config_prediction_success(opc):
|
||||||
assert output == {'data': 'data'}
|
opc.write_data = MagicMock(return_value=(0.1, None))
|
||||||
opc.write_data.assert_has_calls([
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
call(
|
tags_config = {'tag1': {'data_type': 'float'}}
|
||||||
|
|
||||||
|
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
|
||||||
|
server_id='server1',
|
||||||
|
tags_config=tags_config,
|
||||||
|
data=data,
|
||||||
|
data_column='prediction',
|
||||||
|
tag_type='prediction',
|
||||||
|
log_label='Prediction data',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response_times == {'tag1': 0.1}
|
||||||
|
assert session_bad is False
|
||||||
|
assert opc_status is None
|
||||||
|
assert reconnect is False
|
||||||
|
opc.write_data.assert_called_once_with(
|
||||||
server_id='server1',
|
server_id='server1',
|
||||||
tag='tag1',
|
tag='tag1',
|
||||||
data=0.75,
|
data=0.75,
|
||||||
data_type='float',
|
data_type='float',
|
||||||
tag_type='prediction',
|
tag_type='prediction',
|
||||||
metadata=metadata['metadata']
|
metadata=metadata['metadata'],
|
||||||
)])
|
)
|
||||||
opc.write_data.assert_has_calls([
|
|
||||||
call(
|
|
||||||
|
def test_write_tags_from_config_confidence_success(opc):
|
||||||
|
opc.write_data = MagicMock(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(
|
||||||
|
server_id='server1',
|
||||||
|
tags_config=tags_config,
|
||||||
|
data=data,
|
||||||
|
data_column='prediction_confidence',
|
||||||
|
tag_type='confidence',
|
||||||
|
log_label='Confidence data',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response_times == {'tag2': 0.2}
|
||||||
|
assert session_bad is False
|
||||||
|
assert opc_status is None
|
||||||
|
assert reconnect is False
|
||||||
|
opc.write_data.assert_called_once_with(
|
||||||
server_id='server1',
|
server_id='server1',
|
||||||
tag='tag2',
|
tag='tag2',
|
||||||
data=0.95,
|
data=0.95,
|
||||||
data_type='float',
|
data_type='float',
|
||||||
tag_type='confidence',
|
tag_type='confidence',
|
||||||
metadata=metadata['metadata']
|
metadata=metadata['metadata'],
|
||||||
)
|
)
|
||||||
])
|
|
||||||
assert opc.write_data.call_count == 2
|
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
def test_write_tags_from_config_write_failure(opc):
|
||||||
async def test_write_opc_data_empty_config(opc):
|
opc.write_data = MagicMock(return_value=(None, {}))
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
|
||||||
|
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
|
||||||
|
server_id='server1',
|
||||||
|
tags_config={'tag1': {'data_type': 'float'}},
|
||||||
|
data=data,
|
||||||
|
data_column='prediction',
|
||||||
|
tag_type='prediction',
|
||||||
|
log_label='Prediction data',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response_times == {'tag1': None}
|
||||||
|
assert session_bad is False
|
||||||
|
assert opc_status is None
|
||||||
|
assert reconnect is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_write_tags_from_config_session_bad(opc):
|
||||||
|
opc.write_data = MagicMock(
|
||||||
|
return_value=(
|
||||||
|
None,
|
||||||
|
{
|
||||||
|
'opc_error_kind': 'session_bad',
|
||||||
|
'opc_status': 'BadSessionIdInvalid',
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
|
||||||
|
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
|
||||||
|
server_id='server1',
|
||||||
|
tags_config={'tag1': {'data_type': 'float'}},
|
||||||
|
data=data,
|
||||||
|
data_column='prediction',
|
||||||
|
tag_type='prediction',
|
||||||
|
log_label='Prediction data',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response_times == {'tag1': None}
|
||||||
|
assert session_bad is True
|
||||||
|
assert opc_status == 'BadSessionIdInvalid'
|
||||||
|
assert reconnect is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_write_tags_from_config_reconnect_in_progress(opc):
|
||||||
|
opc.write_data = MagicMock(
|
||||||
|
return_value=(
|
||||||
|
None,
|
||||||
|
{'opc_error_kind': 'reconnect_in_progress'},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
|
||||||
|
response_times, session_bad, opc_status, reconnect = opc._write_tags_from_config(
|
||||||
|
server_id='server1',
|
||||||
|
tags_config={'tag1': {'data_type': 'float'}},
|
||||||
|
data=data,
|
||||||
|
data_column='prediction',
|
||||||
|
tag_type='prediction',
|
||||||
|
log_label='Prediction data',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response_times == {'tag1': None}
|
||||||
|
assert session_bad is False
|
||||||
|
assert opc_status is None
|
||||||
|
assert reconnect is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_manage_output_tags_success(opc):
|
||||||
|
opc._write_tags_from_config = MagicMock(
|
||||||
|
side_effect=[
|
||||||
|
({'tag1': 0.1}, False, None, False),
|
||||||
|
({'tag2': 0.1}, False, None, False),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
config = {
|
||||||
|
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||||
|
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||||
|
}
|
||||||
|
|
||||||
|
output_data, opc_metrics, session_bad, opc_status, reconnect = opc.manage_output_tags(
|
||||||
|
server_id='server1',
|
||||||
|
config=config,
|
||||||
|
data=data,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert output_data is True
|
||||||
|
assert opc_metrics == {'tag1': 0.1, 'tag2': 0.1}
|
||||||
|
assert session_bad is False
|
||||||
|
assert opc_status is None
|
||||||
|
assert reconnect is False
|
||||||
|
assert opc._write_tags_from_config.call_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_manage_output_tags_failed(opc):
|
||||||
|
opc._write_tags_from_config = MagicMock(
|
||||||
|
side_effect=[
|
||||||
|
({'tag1': 0.1}, False, None, False),
|
||||||
|
({'tag2': None}, False, None, False),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
config = {
|
||||||
|
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||||
|
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||||
|
}
|
||||||
|
|
||||||
|
output_data, opc_metrics, _, _, _ = opc.manage_output_tags(
|
||||||
|
server_id='server1',
|
||||||
|
config=config,
|
||||||
|
data=data,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert output_data is False
|
||||||
|
assert opc_metrics == {'tag1': 0.1, 'tag2': None}
|
||||||
|
|
||||||
|
|
||||||
|
def test_manage_output_tags_do_nothing(opc):
|
||||||
|
opc._write_tags_from_config = MagicMock()
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
config = {'_invalid_key': {'tag1': {'data_type': 'float'}}}
|
||||||
|
|
||||||
|
output_data, opc_metrics, _, _, _ = opc.manage_output_tags(
|
||||||
|
server_id='server1',
|
||||||
|
config=config,
|
||||||
|
data=data,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert output_data is True
|
||||||
|
assert opc_metrics == {}
|
||||||
|
opc._write_tags_from_config.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.opc.DataFrame')
|
||||||
|
def test_write_opc_data_success(mock_dataframe, opc):
|
||||||
# Arrange
|
# Arrange
|
||||||
input_data = {
|
input_data = {
|
||||||
**metadata,
|
**metadata,
|
||||||
'data': {
|
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||||
'prediction': [0.75],
|
|
||||||
'prediction_confidence': [0.95]
|
|
||||||
},
|
|
||||||
'opc_servers': ['server1'],
|
|
||||||
'opc_output_config': {
|
'opc_output_config': {
|
||||||
'server1': {
|
'server1': {
|
||||||
'prediction_tags': {},
|
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||||
'confidence_tags': {}
|
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
# Act
|
||||||
await opc.write_opc_data(input_data)
|
opc.manage_output_tags = MagicMock(
|
||||||
|
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)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert output_data == {'data': 'data'}
|
||||||
|
assert opc_metrics == {'server1': {'tag1': 0.1, 'tag2': 0.2}}
|
||||||
|
opc.manage_output_tags.assert_called_once_with(
|
||||||
|
'server1',
|
||||||
|
input_data['opc_output_config']['server1'],
|
||||||
|
mock_dataframe.return_value,
|
||||||
|
metadata['metadata'],
|
||||||
|
)
|
||||||
|
opc.process_confidence.assert_called_once_with(
|
||||||
|
mock_dataframe.return_value,
|
||||||
|
True,
|
||||||
|
metadata['metadata'],
|
||||||
|
session_bad=False,
|
||||||
|
opc_status=None,
|
||||||
|
reconnect_in_progress=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_write_opc_data_empty_config(opc):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||||
|
'opc_servers': ['server1'],
|
||||||
|
'opc_output_config': {'server1': {'prediction_tags': {}, 'confidence_tags': {}}},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Act
|
||||||
|
opc.write_opc_data(input_data)
|
||||||
|
|
||||||
# Assert
|
# Assert
|
||||||
opc.opc_repository['server1'].write_data.assert_not_called()
|
opc.opc_repository['server1'].write_data.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
def test_write_opc_data_no_validate_server(opc):
|
||||||
async def test_write_opc_data_no_validate_server(opc):
|
|
||||||
opc.validate_server = MagicMock(return_value=False)
|
opc.validate_server = MagicMock(return_value=False)
|
||||||
input_data = {
|
input_data = {
|
||||||
**metadata,
|
**metadata,
|
||||||
'data': {
|
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||||
'prediction': [0.75],
|
|
||||||
'prediction_confidence': [0.95]
|
|
||||||
},
|
|
||||||
'opc_output_config': {
|
'opc_output_config': {
|
||||||
'server1': {
|
'server1': {
|
||||||
'prediction_tags': {
|
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||||
'tag1': {'data_type': 'float'}
|
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||||
|
}
|
||||||
},
|
},
|
||||||
'confidence_tags': {
|
|
||||||
'tag2': {'data_type': 'float'}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
# Act
|
||||||
await opc.write_opc_data(input_data)
|
opc.write_opc_data(input_data)
|
||||||
|
|
||||||
# Assert
|
# Assert
|
||||||
opc.opc_repository['server1'].write_data.assert_not_called()
|
opc.opc_repository['server1'].write_data.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
@mark.parametrize('data,success,expected', [
|
@mark.parametrize(
|
||||||
|
'data,success,expected',
|
||||||
|
[
|
||||||
(DataFrame({'prediction_confidence': [0]}), True, 0),
|
(DataFrame({'prediction_confidence': [0]}), True, 0),
|
||||||
(DataFrame({'prediction_confidence': [0]}), False, 12),
|
(DataFrame({'prediction_confidence': [0]}), False, 12),
|
||||||
])
|
],
|
||||||
|
)
|
||||||
def test_process_confidence(opc, data, success, expected):
|
def test_process_confidence(opc, data, success, expected):
|
||||||
# Act
|
result = opc.process_confidence(data, success, metadata['metadata'])
|
||||||
result = opc.process_confidence(data, success, metadata)
|
|
||||||
|
|
||||||
# Assert
|
|
||||||
assert result['prediction_confidence'][0] == expected
|
assert result['prediction_confidence'][0] == expected
|
||||||
|
|
||||||
|
|
||||||
|
def test_process_confidence_session_bad(opc):
|
||||||
|
data = DataFrame({'prediction_confidence': [0.9]})
|
||||||
|
result = opc.process_confidence(
|
||||||
|
data,
|
||||||
|
False,
|
||||||
|
metadata['metadata'],
|
||||||
|
session_bad=True,
|
||||||
|
opc_status='BadSessionIdInvalid',
|
||||||
|
)
|
||||||
|
assert result['prediction_confidence'][0] == OPC_SESSION_BAD_CONFIDENCE
|
||||||
|
assert result['comments'][0].startswith(OPC_SESSION_BAD_COMMENT_PREFIX)
|
||||||
|
assert 'BadSessionIdInvalid' in result['comments'][0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_process_confidence_generic_failure(opc):
|
||||||
|
data = DataFrame({'prediction_confidence': [0.9]})
|
||||||
|
result = opc.process_confidence(data, False, metadata['metadata'])
|
||||||
|
assert result['prediction_confidence'][0] == OPC_WRITTING_ERROR_CONFIDENCE
|
||||||
|
assert result['comments'][0] == OPC_WRITTING_ERROR_MESSAGE
|
||||||
|
|
||||||
|
|
||||||
|
def test_manage_output_tags_merges_error_flags(opc):
|
||||||
|
opc._write_tags_from_config = MagicMock(
|
||||||
|
side_effect=[
|
||||||
|
({'tag1': None}, True, 'BadSessionIdInvalid', False),
|
||||||
|
({'tag2': 0.2}, False, None, True),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||||
|
config = {
|
||||||
|
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||||
|
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||||
|
}
|
||||||
|
|
||||||
|
(
|
||||||
|
success,
|
||||||
|
metrics,
|
||||||
|
session_bad_seen,
|
||||||
|
opc_status,
|
||||||
|
reconnect_in_progress,
|
||||||
|
) = opc.manage_output_tags('server1', config, data, metadata['metadata'])
|
||||||
|
|
||||||
|
assert success is False
|
||||||
|
assert session_bad_seen is True
|
||||||
|
assert reconnect_in_progress is True
|
||||||
|
assert opc_status == 'BadSessionIdInvalid'
|
||||||
|
assert metrics == {'tag1': None, 'tag2': 0.2}
|
||||||
|
|
||||||
|
|
||||||
|
def test_process_confidence_reconnect_in_progress(opc):
|
||||||
|
data = DataFrame({'prediction_confidence': [0.9]})
|
||||||
|
result = opc.process_confidence(
|
||||||
|
data,
|
||||||
|
False,
|
||||||
|
metadata['metadata'],
|
||||||
|
reconnect_in_progress=True,
|
||||||
|
)
|
||||||
|
assert result['prediction_confidence'][0] == OPC_SESSION_BAD_CONFIDENCE
|
||||||
|
assert result['comments'][0] == OPC_RECONNECT_IN_PROGRESS_COMMENT
|
||||||
|
|
||||||
|
|
||||||
|
def test_process_confidence_concatenates_multiple_comments(opc):
|
||||||
|
data = DataFrame({'prediction_confidence': [0.9]})
|
||||||
|
session_comment = f'{OPC_SESSION_BAD_COMMENT_PREFIX} BadSessionIdInvalid'
|
||||||
|
|
||||||
|
result = opc.process_confidence(
|
||||||
|
data,
|
||||||
|
False,
|
||||||
|
metadata['metadata'],
|
||||||
|
session_bad=True,
|
||||||
|
opc_status='BadSessionIdInvalid',
|
||||||
|
reconnect_in_progress=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result['prediction_confidence'][0] == OPC_SESSION_BAD_CONFIDENCE
|
||||||
|
assert result['comments'][0] == OPC_COMMENT_SEPARATOR.join(
|
||||||
|
[session_comment, OPC_RECONNECT_IN_PROGRESS_COMMENT]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_validate_server(opc):
|
def test_validate_server(opc):
|
||||||
assert opc.validate_server('server1', metadata) is True
|
assert opc.validate_server('server1', metadata) is True
|
||||||
assert opc.validate_server('server2', metadata) is False
|
assert opc.validate_server('server2', metadata) is False
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
def test_close(opc):
|
||||||
async def test_shutdown(opc):
|
repo = opc.opc_repository['server1']
|
||||||
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
|
repo.disconnect = MagicMock(return_value=True)
|
||||||
await opc.shutdown()
|
opc.close()
|
||||||
opc.opc_repository['server1'].disconnect.assert_called_once()
|
repo.disconnect.assert_called_once()
|
||||||
|
|||||||
324
tests/laborious/activities/test_storage.py
Normal file
324
tests/laborious/activities/test_storage.py
Normal file
@@ -0,0 +1,324 @@
|
|||||||
|
import datetime
|
||||||
|
import os
|
||||||
|
from unittest.mock import ANY, MagicMock, patch
|
||||||
|
|
||||||
|
from pytest import fixture, 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 laborious.activities.storage import Storage
|
||||||
|
|
||||||
|
|
||||||
|
@fixture(autouse=True)
|
||||||
|
def _passthrough_from_dict():
|
||||||
|
with patch(
|
||||||
|
'laborious.activities.storage.MinioDataFramePayload.from_dict', side_effect=lambda x: x
|
||||||
|
):
|
||||||
|
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',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'test_workflow',
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
@patch('laborious.activities.storage.MinioRepository')
|
||||||
|
def storage(mock_minio_repository):
|
||||||
|
return Storage(
|
||||||
|
host='localhost',
|
||||||
|
port=5432,
|
||||||
|
user='postgres',
|
||||||
|
password='postgres',
|
||||||
|
dbname='postgres',
|
||||||
|
min_connections=1,
|
||||||
|
max_connections=10,
|
||||||
|
retention_hours=24,
|
||||||
|
minio_repository=mock_minio_repository.return_value,
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.MinioRepository')
|
||||||
|
def test___init___not_hasattr(mock_minio_repository):
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
metrics_controller = MagicMock()
|
||||||
|
minio_repo = mock_minio_repository.return_value
|
||||||
|
storage = Storage(
|
||||||
|
host='localhost',
|
||||||
|
port=5432,
|
||||||
|
user='postgres',
|
||||||
|
password='postgres',
|
||||||
|
dbname='postgres',
|
||||||
|
min_connections=1,
|
||||||
|
max_connections=10,
|
||||||
|
retention_hours=24,
|
||||||
|
minio_repository=minio_repo,
|
||||||
|
logger=logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
|
)
|
||||||
|
assert isinstance(storage, Postgres)
|
||||||
|
|
||||||
|
assert storage.minio_repository is minio_repo
|
||||||
|
mock_minio_repository.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.MinioRepository')
|
||||||
|
def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||||
|
storage.minio_repository = None
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
metrics_controller = MagicMock()
|
||||||
|
storage.__init__(
|
||||||
|
host='localhost',
|
||||||
|
port=5432,
|
||||||
|
user='postgres',
|
||||||
|
password='postgres',
|
||||||
|
dbname='postgres',
|
||||||
|
min_connections=1,
|
||||||
|
max_connections=10,
|
||||||
|
retention_hours=24,
|
||||||
|
minio_repository=None,
|
||||||
|
logger=logger,
|
||||||
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert storage.minio_repository is None
|
||||||
|
mock_minio_repository.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.MinioRepository')
|
||||||
|
def test___init___done_repository(mock_minio_repository, storage):
|
||||||
|
storage.__init__(
|
||||||
|
host='localhost',
|
||||||
|
port=5432,
|
||||||
|
user='postgres',
|
||||||
|
password='postgres',
|
||||||
|
dbname='postgres',
|
||||||
|
min_connections=1,
|
||||||
|
max_connections=10,
|
||||||
|
retention_hours=24,
|
||||||
|
minio_repository=mock_minio_repository.return_value,
|
||||||
|
logger=MagicMock(),
|
||||||
|
notification_handler=MagicMock(),
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
|
)
|
||||||
|
mock_minio_repository.assert_not_called()
|
||||||
|
assert storage.minio_repository is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_close(storage, _patch_monitoring_shutdown):
|
||||||
|
storage.minio_repository = MagicMock()
|
||||||
|
|
||||||
|
storage.close()
|
||||||
|
|
||||||
|
assert storage.minio_repository is None
|
||||||
|
_patch_monitoring_shutdown.assert_called_once_with(storage)
|
||||||
|
|
||||||
|
|
||||||
|
def test_close_when_minio_repository_already_none(storage, _patch_monitoring_shutdown):
|
||||||
|
"""Closing without an initialized MinIO repository skips MinIO teardown."""
|
||||||
|
storage.minio_repository = None
|
||||||
|
|
||||||
|
storage.close()
|
||||||
|
|
||||||
|
assert storage.minio_repository is None
|
||||||
|
_patch_monitoring_shutdown.assert_called_once_with(storage)
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_query_with_minio_offload_no_rows(storage):
|
||||||
|
storage.load_custom_query = MagicMock(return_value=None)
|
||||||
|
storage_result = {'success': False}
|
||||||
|
with patch(
|
||||||
|
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||||
|
new_callable=MagicMock,
|
||||||
|
return_value=storage_result,
|
||||||
|
) as mock_from_dataframe:
|
||||||
|
result = storage.load_query_with_minio_offload(
|
||||||
|
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
|
||||||
|
)
|
||||||
|
assert result == storage_result
|
||||||
|
mock_from_dataframe.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_query_with_minio_offload_inline(storage):
|
||||||
|
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
|
||||||
|
storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None}
|
||||||
|
with patch(
|
||||||
|
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||||
|
new_callable=MagicMock,
|
||||||
|
return_value=storage_result,
|
||||||
|
) as mock_from_dataframe:
|
||||||
|
result = storage.load_query_with_minio_offload(
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'query': 'SELECT 1',
|
||||||
|
'model_name': 'my-model',
|
||||||
|
'key_prefix': 'predictions/s',
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert result == storage_result
|
||||||
|
mock_from_dataframe.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_query_with_minio_offload_minio(storage):
|
||||||
|
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
|
||||||
|
storage_result = {'success': True, 'data': None, 'object_key': 'object-key'}
|
||||||
|
with patch(
|
||||||
|
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||||
|
new_callable=MagicMock,
|
||||||
|
return_value=storage_result,
|
||||||
|
) as mock_from_dataframe:
|
||||||
|
result = storage.load_query_with_minio_offload(
|
||||||
|
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == storage_result
|
||||||
|
mock_from_dataframe.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@patch.dict(os.environ, {'SIENTIA_MINIO_RETENTION_HOURS': '1'})
|
||||||
|
@patch('laborious.activities.storage.now')
|
||||||
|
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(
|
||||||
|
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()
|
||||||
|
|
||||||
|
data_mock = MagicMock()
|
||||||
|
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||||
|
|
||||||
|
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||||
|
|
||||||
|
assert result['deleted_count'] == 1
|
||||||
|
assert result['failed_count'] == 0
|
||||||
|
deleted_key = (
|
||||||
|
'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet'
|
||||||
|
)
|
||||||
|
assert deleted_key in result['deleted']
|
||||||
|
assert result['deleted'][deleted_key]['success'] is True
|
||||||
|
storage.minio_repository.list_objects.assert_called_once_with(
|
||||||
|
prefix='training_datasets/m',
|
||||||
|
recursive=True,
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
storage.minio_repository.delete_file.assert_called_once_with(
|
||||||
|
object_name='sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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'})
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_payload_to_postgres(storage):
|
||||||
|
payload = MagicMock()
|
||||||
|
payload.retrieve = MagicMock(return_value=MagicMock())
|
||||||
|
storage.export_data_to_postgres = MagicMock(return_value={'success': True})
|
||||||
|
|
||||||
|
result = 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()
|
||||||
|
assert result == {'success': True}
|
||||||
|
|
||||||
|
|
||||||
|
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})
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.now')
|
||||||
|
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(
|
||||||
|
return_value=['some/random/key-without-timestamp.parquet']
|
||||||
|
)
|
||||||
|
storage.minio_repository.delete_file = MagicMock()
|
||||||
|
storage.send_notification = MagicMock()
|
||||||
|
|
||||||
|
data_mock = MagicMock()
|
||||||
|
data_mock.cleanup_prefix.return_value = 'test'
|
||||||
|
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||||
|
|
||||||
|
assert result['deleted_count'] == 0
|
||||||
|
assert result['failed_count'] == 0
|
||||||
|
storage.minio_repository.delete_file.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.now')
|
||||||
|
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()
|
||||||
|
|
||||||
|
data_mock = MagicMock()
|
||||||
|
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||||
|
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||||
|
|
||||||
|
assert result['deleted_count'] == 0
|
||||||
|
assert result['failed_count'] == 1
|
||||||
|
assert old_key in result['failed']
|
||||||
|
assert result['failed'][old_key]['success'] is False
|
||||||
|
assert result['failed'][old_key]['message'] == 'delete error'
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.activities.storage.now')
|
||||||
|
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.error = MagicMock()
|
||||||
|
|
||||||
|
data_mock = MagicMock()
|
||||||
|
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||||
|
result = 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(
|
||||||
|
metadata=metadata['metadata'],
|
||||||
|
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
|
||||||
|
message='Error cleaning up MinIO objects: list error',
|
||||||
|
block='cleanup_minio_objects_expired',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=ANY,
|
||||||
|
)
|
||||||
|
storage.error.assert_called_once()
|
||||||
@@ -1,23 +1,36 @@
|
|||||||
from pandas import DataFrame
|
from pandas import DataFrame
|
||||||
|
|
||||||
from laborious.utils.filters.conditional_filters import (
|
from laborious.utils.filters.conditional_filters import (
|
||||||
|
filter_empty_data,
|
||||||
filter_specific_variables_null_values,
|
filter_specific_variables_null_values,
|
||||||
filter_empty_data
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_filter_specific_variables_null_values():
|
def test_filter_specific_variables_null_values():
|
||||||
assert filter_specific_variables_null_values(
|
assert (
|
||||||
DataFrame(
|
filter_specific_variables_null_values(
|
||||||
{'variable': ['variable1', 'variable2'], 'value': [1, 2]}),
|
DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}),
|
||||||
config={'variables': ['variable2']}) is False
|
config={'variables': ['variable2']},
|
||||||
|
)
|
||||||
|
is False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_filter_specific_variables_null_values_with_empty_data():
|
||||||
|
assert (
|
||||||
|
filter_specific_variables_null_values(DataFrame(), config={'variables': ['variable2']})
|
||||||
|
is False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_filter_specific_variables_null_values_with_null_values():
|
def test_filter_specific_variables_null_values_with_null_values():
|
||||||
assert filter_specific_variables_null_values(
|
assert (
|
||||||
DataFrame(
|
filter_specific_variables_null_values(
|
||||||
{'variable': ['variable1', 'variable2'], 'value': [1, None]}),
|
DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, None]}),
|
||||||
config={'variables': ['variable2']}) is True
|
config={'variables': ['variable2']},
|
||||||
|
)
|
||||||
|
is True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_filter_empty_data():
|
def test_filter_empty_data():
|
||||||
@@ -25,6 +38,7 @@ def test_filter_empty_data():
|
|||||||
|
|
||||||
|
|
||||||
def test_filter_empty_data_with_data():
|
def test_filter_empty_data_with_data():
|
||||||
assert filter_empty_data(
|
assert (
|
||||||
DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}),
|
filter_empty_data(DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}), {})
|
||||||
{}) is False
|
is False
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,22 +1,23 @@
|
|||||||
from pandas import DataFrame
|
from pandas import DataFrame
|
||||||
|
|
||||||
from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter
|
from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter
|
||||||
|
|
||||||
|
|
||||||
def test_api_error_filter_invalid_response():
|
def test_api_error_filter_invalid_response():
|
||||||
assert api_error_filter(None, {}) == True # NOSONAR
|
assert api_error_filter(None, {}) is True # NOSONAR
|
||||||
|
|
||||||
|
|
||||||
def test_api_error_filter_valid_response_fail():
|
def test_api_error_filter_valid_response_fail():
|
||||||
assert api_error_filter({'success': False}, {}) == True
|
assert api_error_filter({'success': False}, {}) is True
|
||||||
|
|
||||||
|
|
||||||
def test_api_error_filter_valid_response_success():
|
def test_api_error_filter_valid_response_success():
|
||||||
assert api_error_filter({'success': True}, {}) == False
|
assert api_error_filter({'success': True}, {}) is False
|
||||||
|
|
||||||
|
|
||||||
def test_nan_values_filter_all_nan_values():
|
def test_nan_values_filter_all_nan_values():
|
||||||
assert nan_values_filter(DataFrame({'variable': [None, None]}), {}) == True
|
assert nan_values_filter(DataFrame({'variable': [None, None]}), {}) is True
|
||||||
|
|
||||||
|
|
||||||
def test_nan_values_filter_no_nan_values():
|
def test_nan_values_filter_no_nan_values():
|
||||||
assert nan_values_filter(DataFrame({'variable': [1, 2]}), {}) == False
|
assert nan_values_filter(DataFrame({'variable': [1, 2]}), {}) is False
|
||||||
|
|||||||
278
tests/laborious/utils/models/test_minio_dataframe_payload.py
Normal file
278
tests/laborious/utils/models/test_minio_dataframe_payload.py
Normal file
@@ -0,0 +1,278 @@
|
|||||||
|
from datetime import datetime
|
||||||
|
from io import BytesIO
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from pandas import DataFrame
|
||||||
|
|
||||||
|
from laborious.utils.models.minio_dataframe_payload import (
|
||||||
|
MinioDataFramePayload,
|
||||||
|
_build_object_key,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_object_timestamp_hyphenated_model():
|
||||||
|
key = 'predictions/sched/my-long-model-initial-2024-06-15_10-30-45.parquet'
|
||||||
|
ts = MinioDataFramePayload.parse_object_timestamp(key)
|
||||||
|
assert ts == datetime(2024, 6, 15, 10, 30, 45)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_object_timestamp_transform():
|
||||||
|
key = 'p/m-transform-2024-01-02_03-04-05.parquet'
|
||||||
|
ts = MinioDataFramePayload.parse_object_timestamp(key)
|
||||||
|
assert ts == datetime(2024, 1, 2, 3, 4, 5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_object_timestamp_invalid():
|
||||||
|
assert MinioDataFramePayload.parse_object_timestamp('bad.parquet') is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_estimate_size_bytes_returns_positive_for_nonempty_frame():
|
||||||
|
df = DataFrame({'a': [1, 2]})
|
||||||
|
size = MinioDataFramePayload.estimate_size_bytes(df)
|
||||||
|
assert isinstance(size, int)
|
||||||
|
assert size > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_cleanup_prefix_when_offloaded_returns_object_prefix():
|
||||||
|
payload = MinioDataFramePayload(
|
||||||
|
last_timestamp='t',
|
||||||
|
data=None,
|
||||||
|
object_key='training_datasets/m/m-initial-2024-01-01_00-00-00.parquet',
|
||||||
|
object_prefix='training_datasets/m',
|
||||||
|
)
|
||||||
|
assert MinioDataFramePayload.cleanup_prefix(payload) == 'training_datasets/m'
|
||||||
|
|
||||||
|
|
||||||
|
def test_cleanup_prefix_when_inline_returns_none():
|
||||||
|
payload = MinioDataFramePayload(last_timestamp='t', data={'x': [1]}, object_key=None)
|
||||||
|
assert MinioDataFramePayload.cleanup_prefix(payload) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_has_data_true_when_object_key_set():
|
||||||
|
payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key='k')
|
||||||
|
assert payload.has_data() is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_retrieve_inline_dict_as_dataframe():
|
||||||
|
payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]})
|
||||||
|
minio = MagicMock()
|
||||||
|
out = payload.retrieve(minio, {'metadata': {}})
|
||||||
|
assert list(out.columns) == ['a']
|
||||||
|
minio.download_file.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_retrieve_downloads_parquet_when_offloaded():
|
||||||
|
source = DataFrame({'a': [1, 2]})
|
||||||
|
buf = BytesIO()
|
||||||
|
source.to_parquet(buf, engine='pyarrow', index=True)
|
||||||
|
file_bytes = buf.getvalue()
|
||||||
|
|
||||||
|
payload = MinioDataFramePayload(
|
||||||
|
last_timestamp='t',
|
||||||
|
data=None,
|
||||||
|
object_key='training_datasets/m/f.parquet',
|
||||||
|
object_prefix='training_datasets/m',
|
||||||
|
)
|
||||||
|
minio = MagicMock()
|
||||||
|
minio.download_file = MagicMock(return_value=file_bytes)
|
||||||
|
|
||||||
|
out = payload.retrieve(minio, {'metadata': {}})
|
||||||
|
|
||||||
|
minio.download_file.assert_called_once_with(
|
||||||
|
object_name='training_datasets/m/f.parquet',
|
||||||
|
metadata={'metadata': {}},
|
||||||
|
)
|
||||||
|
assert list(out.columns) == ['a']
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_object_key():
|
||||||
|
key, prefix = _build_object_key('my-model', 'initial', '2024-01-01_00-00-00')
|
||||||
|
assert key == 'prediction_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet'
|
||||||
|
assert prefix == 'prediction_datasets/my-model'
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_object_key_strips_slashes():
|
||||||
|
key, prefix = _build_object_key(' /my-model/ ', 'transform', '2024-06-15_10-30-45')
|
||||||
|
assert prefix == 'prediction_datasets/my-model'
|
||||||
|
assert key.startswith('prediction_datasets/my-model/')
|
||||||
|
|
||||||
|
|
||||||
|
def test_estimate_size_bytes_fallback():
|
||||||
|
df = DataFrame({'a': [1, 2]})
|
||||||
|
with patch.object(df, 'to_dict', side_effect=RuntimeError('to_dict failed')):
|
||||||
|
size = MinioDataFramePayload.estimate_size_bytes(df)
|
||||||
|
assert isinstance(size, int)
|
||||||
|
assert size > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_object_timestamp_bad_datetime():
|
||||||
|
key = 'p/m-initial-9999-99-99_99-99-99.parquet'
|
||||||
|
assert MinioDataFramePayload.parse_object_timestamp(key) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_retrieve_empty_when_no_data():
|
||||||
|
payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key=None)
|
||||||
|
minio = MagicMock()
|
||||||
|
out = payload.retrieve(minio, {})
|
||||||
|
assert out.empty
|
||||||
|
minio.download_file.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.utils.models.minio_dataframe_payload.now')
|
||||||
|
def test_from_dataframe_none(mock_now):
|
||||||
|
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
|
||||||
|
minio = MagicMock()
|
||||||
|
result = MinioDataFramePayload.from_dataframe(
|
||||||
|
dataframe=None,
|
||||||
|
minio_repo=minio,
|
||||||
|
model_name='m',
|
||||||
|
operation='initial',
|
||||||
|
status={'success': False, 'message': 'no data'},
|
||||||
|
)
|
||||||
|
assert result.data is None
|
||||||
|
assert result.status == {'success': False, 'message': 'no data'}
|
||||||
|
assert result.object_key is None
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.utils.models.minio_dataframe_payload.now')
|
||||||
|
def test_from_dataframe_empty(mock_now):
|
||||||
|
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
|
||||||
|
minio = MagicMock()
|
||||||
|
mock_df = MagicMock()
|
||||||
|
mock_df.__bool__ = MagicMock(return_value=True)
|
||||||
|
mock_df.empty = True
|
||||||
|
result = MinioDataFramePayload.from_dataframe(
|
||||||
|
dataframe=mock_df,
|
||||||
|
minio_repo=minio,
|
||||||
|
model_name='m',
|
||||||
|
operation='initial',
|
||||||
|
)
|
||||||
|
assert result.data is None
|
||||||
|
assert result.object_key is None
|
||||||
|
|
||||||
|
|
||||||
|
def _mock_dataframe(data_dict, timestamp_values=None):
|
||||||
|
"""Build a MagicMock that behaves enough like a DataFrame for from_dataframe."""
|
||||||
|
mock_df = MagicMock()
|
||||||
|
mock_df.__bool__ = MagicMock(return_value=True)
|
||||||
|
mock_df.empty = False
|
||||||
|
if timestamp_values is None:
|
||||||
|
timestamp_values = data_dict.get('timestamp', ['2024-01-01'])
|
||||||
|
ts_col = MagicMock()
|
||||||
|
ts_col.values.tolist.return_value = timestamp_values
|
||||||
|
mock_df.__getitem__ = MagicMock(return_value=ts_col)
|
||||||
|
mock_df.to_dict.return_value = data_dict
|
||||||
|
buf = BytesIO()
|
||||||
|
DataFrame(data_dict).to_parquet(buf, engine='pyarrow', index=True)
|
||||||
|
mock_df.to_parquet = MagicMock(side_effect=lambda b, **kw: b.write(buf.getvalue()))
|
||||||
|
return mock_df
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
|
||||||
|
def test_from_dataframe_inline():
|
||||||
|
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',
|
||||||
|
)
|
||||||
|
assert result.data is not None
|
||||||
|
assert result.object_key is None
|
||||||
|
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'
|
||||||
|
|
||||||
|
|
||||||
|
@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):
|
||||||
|
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.bucket = 'test-bucket'
|
||||||
|
|
||||||
|
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
|
||||||
|
result = MinioDataFramePayload.from_dataframe(
|
||||||
|
dataframe=df,
|
||||||
|
minio_repo=minio,
|
||||||
|
model_name='m',
|
||||||
|
operation='initial',
|
||||||
|
workflow_metadata={'wf': 'data'},
|
||||||
|
)
|
||||||
|
assert result.data is None
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_dict_inline():
|
||||||
|
raw = {
|
||||||
|
'last_timestamp': '2024-01-01T00:00:00+00:00',
|
||||||
|
'status': None,
|
||||||
|
'data': {'col1': {0: 'val1'}},
|
||||||
|
'bucket': None,
|
||||||
|
'object_key': None,
|
||||||
|
'object_prefix': None,
|
||||||
|
'uri': None,
|
||||||
|
}
|
||||||
|
payload = MinioDataFramePayload.from_dict(raw)
|
||||||
|
assert isinstance(payload, MinioDataFramePayload)
|
||||||
|
assert payload.last_timestamp == '2024-01-01T00:00:00+00:00'
|
||||||
|
assert payload.data == {'col1': {0: 'val1'}}
|
||||||
|
assert payload.object_key is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_dict_offloaded():
|
||||||
|
raw = {
|
||||||
|
'last_timestamp': '2024-06-15T10:30:45+00:00',
|
||||||
|
'status': {'success': True},
|
||||||
|
'data': None,
|
||||||
|
'bucket': 'my-bucket',
|
||||||
|
'object_key': 'training_datasets/model/model-initial-2024-06-15_10-30-45.parquet',
|
||||||
|
'object_prefix': 'training_datasets/model',
|
||||||
|
'uri': 's3://my-bucket/training_datasets/model/model-initial-2024-06-15_10-30-45.parquet',
|
||||||
|
}
|
||||||
|
payload = MinioDataFramePayload.from_dict(raw)
|
||||||
|
assert isinstance(payload, MinioDataFramePayload)
|
||||||
|
assert payload.data is None
|
||||||
|
assert payload.bucket == 'my-bucket'
|
||||||
|
assert payload.object_key == raw['object_key']
|
||||||
|
assert payload.object_prefix == 'training_datasets/model'
|
||||||
|
assert payload.uri == raw['uri']
|
||||||
|
assert payload.status == {'success': True}
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_dict_minimal_keys():
|
||||||
|
raw = {'last_timestamp': '2024-01-01'}
|
||||||
|
payload = MinioDataFramePayload.from_dict(raw)
|
||||||
|
assert payload.last_timestamp == '2024-01-01'
|
||||||
|
assert payload.data is None
|
||||||
|
assert payload.bucket is None
|
||||||
|
assert payload.object_key is None
|
||||||
|
|
||||||
|
|
||||||
|
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})
|
||||||
@@ -1,502 +0,0 @@
|
|||||||
from unittest.mock import ANY, MagicMock, call, patch
|
|
||||||
import numpy as np
|
|
||||||
from pandas import DataFrame
|
|
||||||
import pytest
|
|
||||||
from datetime import datetime, timezone
|
|
||||||
from pandas import Timestamp
|
|
||||||
from laborious.utils.repository.model_repository import MLFlowRepository
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def mlflow_repository():
|
|
||||||
with patch('laborious.utils.repository.model_repository.ModelServing',
|
|
||||||
autospec=True) as mock_model_serving:
|
|
||||||
mock_instance = mock_model_serving.return_value
|
|
||||||
mock_instance.get_transformed_data = MagicMock()
|
|
||||||
|
|
||||||
repo = MLFlowRepository(
|
|
||||||
host='http://localhost:5000',
|
|
||||||
username='admin',
|
|
||||||
password='admin',
|
|
||||||
logger=MagicMock()
|
|
||||||
)
|
|
||||||
return repo
|
|
||||||
|
|
||||||
|
|
||||||
metadata = {
|
|
||||||
"metadata": {
|
|
||||||
"model_id": "test_model",
|
|
||||||
"model_name": "test_model",
|
|
||||||
"workflow_name": "test_workflow",
|
|
||||||
"schema_name": "test_schedule",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class Any:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
invalid_cases = [
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
'2024-01-01 12:00:00': 1,
|
|
||||||
2024: 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
),
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
'2024-01-01': 1,
|
|
||||||
'2024-01-02': 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
),
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
Any(): 1,
|
|
||||||
Any(): 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
)
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("data", invalid_cases)
|
|
||||||
def test_detect_and_parse_datetime_index_error_cases(mlflow_repository, data):
|
|
||||||
input_data = DataFrame(
|
|
||||||
data
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(ValueError) as e:
|
|
||||||
mlflow_repository.detect_and_parse_datetime_index(
|
|
||||||
input_data, metadata['metadata'])
|
|
||||||
|
|
||||||
assert str(e) == "Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S"
|
|
||||||
|
|
||||||
|
|
||||||
valid_cases = [
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
'2024-01-01 12:00:00+0000': 1,
|
|
||||||
'2024-01-02 12:00:00+0000': 2
|
|
||||||
}
|
|
||||||
}, ['2024-01-01 12:00:00+0000', '2024-01-02 12:00:00+0000']
|
|
||||||
),
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
datetime(2025, 1, 1, 12, 0, 0, tzinfo=timezone.utc): 1,
|
|
||||||
datetime(2025, 1, 2, 12, 0, 0, tzinfo=timezone.utc): 2
|
|
||||||
}
|
|
||||||
}, ['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000']
|
|
||||||
),
|
|
||||||
(
|
|
||||||
{
|
|
||||||
'value': {
|
|
||||||
Timestamp(2026, 1, 1, 12, 0, 0, tzinfo=timezone.utc): 1,
|
|
||||||
Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=timezone.utc): 2
|
|
||||||
}
|
|
||||||
}, ['2026-01-01 12:00:00+0000', '2026-01-02 12:00:00+0000']
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("data,expected", valid_cases)
|
|
||||||
def test_detect_and_parse_datetime_index_valid_format(mlflow_repository, data, expected):
|
|
||||||
input_data = DataFrame(data)
|
|
||||||
|
|
||||||
response = mlflow_repository.detect_and_parse_datetime_index(
|
|
||||||
input_data, metadata['metadata'])
|
|
||||||
|
|
||||||
assert response.index.tolist() == expected
|
|
||||||
|
|
||||||
|
|
||||||
def test_transform_success(mlflow_repository):
|
|
||||||
data = MagicMock()
|
|
||||||
model_name = 'model'
|
|
||||||
|
|
||||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
|
|
||||||
|
|
||||||
output = mlflow_repository.transform(
|
|
||||||
model_name, data, {}, metadata['metadata'])
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
|
|
||||||
model_name, data, 0, 'sklearn', False, 'model', 'predict')
|
|
||||||
|
|
||||||
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
|
|
||||||
mlflow_repository.model_serving.get_cached_transform.return_value, metadata['metadata'])
|
|
||||||
|
|
||||||
assert output == {
|
|
||||||
'success': True,
|
|
||||||
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_transform_error(mlflow_repository):
|
|
||||||
data = MagicMock()
|
|
||||||
model_name = 'model'
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_transform.side_effect = Exception(
|
|
||||||
'error')
|
|
||||||
|
|
||||||
output = mlflow_repository.transform(
|
|
||||||
model_name, data, {}, metadata['metadata'])
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
|
|
||||||
model_name, data, 0, 'sklearn', False, 'model', 'predict')
|
|
||||||
|
|
||||||
assert output == {
|
|
||||||
'success': False,
|
|
||||||
'content': {
|
|
||||||
'message': 'error',
|
|
||||||
'traceback': ANY
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_predict_success(mlflow_repository):
|
|
||||||
data = DataFrame({
|
|
||||||
'feat_1': {
|
|
||||||
'index_1': 2,
|
|
||||||
'index_2': 3
|
|
||||||
}
|
|
||||||
})
|
|
||||||
model_name = 'model'
|
|
||||||
mlflow_repository.model_serving.get_cached_predict.return_value = np.array(
|
|
||||||
[2, 3]
|
|
||||||
)
|
|
||||||
|
|
||||||
output = mlflow_repository.predict(
|
|
||||||
model_name, data, {}, metadata['metadata'])
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
|
|
||||||
model_name, data, 0, 'pyfunc', False, 'model')
|
|
||||||
|
|
||||||
assert output['success'] is True
|
|
||||||
assert output['content'] == {
|
|
||||||
'prediction': {
|
|
||||||
'index_1': 2,
|
|
||||||
'index_2': 3
|
|
||||||
}, 'response_time': {
|
|
||||||
'index_1': ANY,
|
|
||||||
'index_2': ANY
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_predict_error(mlflow_repository):
|
|
||||||
data = DataFrame({
|
|
||||||
'feat_1': {
|
|
||||||
'index_1': 2,
|
|
||||||
'index_2': 3
|
|
||||||
}
|
|
||||||
})
|
|
||||||
model_name = 'model'
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_predict = MagicMock(
|
|
||||||
side_effect=Exception('error')
|
|
||||||
)
|
|
||||||
|
|
||||||
output = mlflow_repository.predict(
|
|
||||||
model_name, data, {}, metadata['metadata'])
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
|
|
||||||
model_name, data, 0, 'pyfunc', False, 'model')
|
|
||||||
|
|
||||||
assert output == {
|
|
||||||
'success': False,
|
|
||||||
'content': {
|
|
||||||
'message': 'error',
|
|
||||||
'traceback': ANY
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_experiment_by_run_id(mlflow, mlflow_repository):
|
|
||||||
mlflow.get_run.return_value = MagicMock(
|
|
||||||
info=MagicMock(
|
|
||||||
experiment_id='0',
|
|
||||||
)
|
|
||||||
)
|
|
||||||
mlflow.get_experiment.return_value = MagicMock()
|
|
||||||
mlflow.get_experiment.return_value.name = 'test'
|
|
||||||
|
|
||||||
output = mlflow_repository.get_experiment_by_run_id('0')
|
|
||||||
assert output == 'test'
|
|
||||||
mlflow.get_run.assert_called_once_with('0')
|
|
||||||
mlflow.get_experiment.assert_called_once_with('0')
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_next_run_name(mlflow, mlflow_repository):
|
|
||||||
mlflow.search_runs.return_value = [1, 2, 3]
|
|
||||||
output = mlflow_repository.get_next_run_name('run')
|
|
||||||
assert output == 'run-4'
|
|
||||||
mlflow.search_runs.assert_called_once_with(
|
|
||||||
experiment_names=['run'],
|
|
||||||
order_by=['start_time desc'],
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_experiment_success(mlflow, mlflow_repository):
|
|
||||||
mlflow.get_experiment_by_name.return_value = MagicMock(
|
|
||||||
experiment_id='0')
|
|
||||||
|
|
||||||
output = mlflow_repository.get_experiment('test')
|
|
||||||
|
|
||||||
assert output == 0
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_experiment_error(mlflow, mlflow_repository):
|
|
||||||
mlflow.get_experiment_by_name.return_value = None
|
|
||||||
|
|
||||||
try:
|
|
||||||
mlflow_repository.get_experiment('test')
|
|
||||||
except ValueError as e:
|
|
||||||
assert str(e) == 'Experiment test not found'
|
|
||||||
else:
|
|
||||||
assert False
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_experiment_last_run(mlflow, mlflow_repository):
|
|
||||||
mlflow.search_runs.return_value = DataFrame({
|
|
||||||
'params.retrain': ['True', 'False', 'True', 'False'],
|
|
||||||
'end_time': ['2021-01-01', '2021-01-02', '2021-01-03', '2021-01-04'],
|
|
||||||
'run_id': ['0', '1', '2', '3'],
|
|
||||||
})
|
|
||||||
|
|
||||||
output = mlflow_repository.get_experiment_last_run(0)
|
|
||||||
|
|
||||||
mlflow.search_runs.assert_called_once_with(
|
|
||||||
experiment_ids=[0],
|
|
||||||
filter_string="",
|
|
||||||
output_format="pandas",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert output == '2'
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_get_experiment_last_run_error(mlflow, mlflow_repository):
|
|
||||||
mlflow.search_runs.return_value = []
|
|
||||||
|
|
||||||
try:
|
|
||||||
mlflow_repository.get_experiment_last_run(0)
|
|
||||||
except ValueError as e:
|
|
||||||
assert str(e) == 'Runs is not a pandas DataFrame'
|
|
||||||
else:
|
|
||||||
assert False
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.sklearn')
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.set_experiment')
|
|
||||||
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_model_run_id = MagicMock(
|
|
||||||
return_value='0')
|
|
||||||
mlflow_repository.model_serving.get_model_uri = MagicMock(
|
|
||||||
return_value='test')
|
|
||||||
mlflow_repository.get_experiment_by_run_id = MagicMock()
|
|
||||||
|
|
||||||
data_model_mock = MagicMock()
|
|
||||||
prediction_model_mock = MagicMock()
|
|
||||||
|
|
||||||
sklearn.load_model.side_effect = [data_model_mock, prediction_model_mock]
|
|
||||||
|
|
||||||
data_model_mock.fit.return_value = data_model_mock
|
|
||||||
data_model_mock.predict.return_value = DataFrame({
|
|
||||||
'x': [10, 20, 30],
|
|
||||||
})
|
|
||||||
data_model_mock.target_variable = 'y'
|
|
||||||
|
|
||||||
prediction_model_mock.fit.return_value = prediction_model_mock
|
|
||||||
|
|
||||||
data = DataFrame({
|
|
||||||
'x': [1, 2, 3],
|
|
||||||
'y': [4, 5, 6]
|
|
||||||
})
|
|
||||||
|
|
||||||
output = mlflow_repository.create_model_experiment('test', data)
|
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_model_run_id.assert_called_once_with(
|
|
||||||
'test', stage='Production')
|
|
||||||
mlflow_repository.model_serving.get_model_uri.assert_called_once_with(
|
|
||||||
'0', prediction=False)
|
|
||||||
|
|
||||||
sklearn.load_model.assert_has_calls([
|
|
||||||
call(mlflow_repository.model_serving.get_model_uri.return_value),
|
|
||||||
call("models:/test/production"),
|
|
||||||
])
|
|
||||||
assert sklearn.load_model.call_count == 2
|
|
||||||
|
|
||||||
data_model_mock.fit.assert_called_once_with(data)
|
|
||||||
data_model_mock.predict.assert_called_once_with(data)
|
|
||||||
|
|
||||||
fit_args = prediction_model_mock.fit.call_args[0][0]
|
|
||||||
assert fit_args.equals(
|
|
||||||
DataFrame({
|
|
||||||
'x': [10, 20, 30],
|
|
||||||
'y': [4, 5, 6],
|
|
||||||
})
|
|
||||||
)
|
|
||||||
|
|
||||||
mlflow_repository.get_experiment_by_run_id.assert_called_once_with('0')
|
|
||||||
|
|
||||||
set_experiment.assert_called_once_with(
|
|
||||||
mlflow_repository.get_experiment_by_run_id.return_value
|
|
||||||
)
|
|
||||||
|
|
||||||
assert output == (prediction_model_mock,
|
|
||||||
data_model_mock,
|
|
||||||
mlflow_repository.get_experiment_by_run_id.return_value)
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.start_run')
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.log_param')
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.sklearn.log_model')
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow.log_artifact')
|
|
||||||
def test_perform_model_retrain(log_artifact, log_model, log_param, start_run, mlflow_repository):
|
|
||||||
|
|
||||||
prediction_model_mock = MagicMock()
|
|
||||||
data_model_mock = MagicMock()
|
|
||||||
experiment = 'test'
|
|
||||||
model_name = 'test'
|
|
||||||
data = MagicMock()
|
|
||||||
|
|
||||||
mlflow_repository.get_next_run_name = MagicMock(
|
|
||||||
return_value='test-1')
|
|
||||||
run = MagicMock()
|
|
||||||
start_run.__enter__.return_value = run
|
|
||||||
|
|
||||||
output = mlflow_repository.perform_model_retrain(
|
|
||||||
prediction_model_mock, data_model_mock, experiment, model_name, data)
|
|
||||||
|
|
||||||
mlflow_repository.get_next_run_name.assert_called_once_with(experiment)
|
|
||||||
start_run.assert_called_once_with(
|
|
||||||
run_name='test-1', description='Retrain model test with new data')
|
|
||||||
|
|
||||||
log_model.assert_has_calls([
|
|
||||||
call(data_model_mock, "data_model"),
|
|
||||||
call(prediction_model_mock, "prediction_model"),
|
|
||||||
])
|
|
||||||
|
|
||||||
data.to_csv.assert_called_once_with(
|
|
||||||
"temp/raw_data_test.csv", index=True)
|
|
||||||
|
|
||||||
log_artifact.assert_called_once_with(
|
|
||||||
"temp/raw_data_test.csv")
|
|
||||||
|
|
||||||
log_param.assert_has_calls([
|
|
||||||
call("retrain", True),
|
|
||||||
])
|
|
||||||
|
|
||||||
assert output == ("Model retrained successfully", experiment)
|
|
||||||
|
|
||||||
|
|
||||||
def test_retrain_model(mlflow_repository):
|
|
||||||
data = MagicMock()
|
|
||||||
model_name = 'test'
|
|
||||||
|
|
||||||
mlflow_repository.create_model_experiment = MagicMock(
|
|
||||||
return_value=('data_model', 'prediction_model', '0'))
|
|
||||||
|
|
||||||
mlflow_repository.perform_model_retrain = MagicMock(
|
|
||||||
return_value='Model retrained successfully')
|
|
||||||
|
|
||||||
output = mlflow_repository.retrain_model(data, model_name)
|
|
||||||
|
|
||||||
mlflow_repository.create_model_experiment.assert_called_once_with(
|
|
||||||
model_name, data)
|
|
||||||
|
|
||||||
mlflow_repository.perform_model_retrain.assert_called_once_with(
|
|
||||||
'data_model', 'prediction_model', '0', model_name, data)
|
|
||||||
|
|
||||||
assert output == 'Model retrained successfully'
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
|
||||||
client_mock = MagicMock()
|
|
||||||
mlflow.tracking.MlflowClient.return_value = client_mock
|
|
||||||
|
|
||||||
client_mock.get_registered_model.return_value = MagicMock(
|
|
||||||
latest_versions=[
|
|
||||||
MagicMock(version='1'),
|
|
||||||
MagicMock(version='2'),
|
|
||||||
MagicMock(version='3'),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
output = mlflow_repository.update_production_model_by_run_id('0', 'test')
|
|
||||||
|
|
||||||
mlflow.register_model.assert_called_once_with(
|
|
||||||
"runs:/0/prediction_model",
|
|
||||||
'test',
|
|
||||||
)
|
|
||||||
|
|
||||||
mlflow.tracking.MlflowClient.assert_called_once()
|
|
||||||
client_mock.get_registered_model.assert_called_once_with('test')
|
|
||||||
client_mock.transition_model_version_stage.assert_called_once_with(
|
|
||||||
name='test',
|
|
||||||
version='3',
|
|
||||||
stage='Production',
|
|
||||||
archive_existing_versions=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert output == {
|
|
||||||
'model_name': 'test',
|
|
||||||
'version': '3',
|
|
||||||
'mlflow_run_id': '0',
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
|
||||||
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
|
||||||
mlflow.tracking.MlflowClient.return_value = MagicMock(
|
|
||||||
get_registered_model=MagicMock(
|
|
||||||
return_value=MagicMock(
|
|
||||||
latest_versions={}
|
|
||||||
)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
mlflow_repository.update_production_model_by_run_id('0', 'test')
|
|
||||||
except Exception as e:
|
|
||||||
assert str(e) == 'Model versions is not a list'
|
|
||||||
else:
|
|
||||||
assert False
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_production_model(mlflow_repository):
|
|
||||||
connector = mlflow_repository
|
|
||||||
|
|
||||||
with patch.object(connector, 'get_experiment',
|
|
||||||
return_value='0') as get_experiment:
|
|
||||||
with patch.object(connector, 'get_experiment_last_run',
|
|
||||||
return_value='2') as get_experiment_last_run:
|
|
||||||
with patch.object(connector, 'update_production_model_by_run_id',
|
|
||||||
return_value={'model_name': 'test', 'version': '3',
|
|
||||||
'mlflow_run_id': '0'}) as update_production_model_by_run_id:
|
|
||||||
|
|
||||||
output = connector.update_production_model('0', 'test')
|
|
||||||
|
|
||||||
get_experiment.assert_called_once_with('0')
|
|
||||||
get_experiment_last_run.assert_called_once_with('0')
|
|
||||||
update_production_model_by_run_id.assert_called_once_with(
|
|
||||||
'2', 'test')
|
|
||||||
|
|
||||||
assert output == {
|
|
||||||
'model_name': 'test',
|
|
||||||
'version': '3',
|
|
||||||
'mlflow_run_id': '0',
|
|
||||||
'mlflow_experiment_id': '0',
|
|
||||||
}
|
|
||||||
@@ -1,9 +1,20 @@
|
|||||||
import pytest
|
import concurrent.futures
|
||||||
from unittest.mock import AsyncMock, Mock, patch, MagicMock, ANY, call
|
import json
|
||||||
from asyncua.crypto.security_policies import SecurityPolicyBasic256
|
|
||||||
from laborious.utils.repository.opc_repository import OpcRepository
|
|
||||||
from sientia_do.notifications.models import NotificationLevel
|
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
from unittest.mock import ANY, MagicMock, Mock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from asyncua.crypto import security_policies
|
||||||
|
from asyncua.ua.uaerrors import BadNodeIdUnknown, BadSessionIdInvalid
|
||||||
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
|
|
||||||
|
from laborious.utils.repository.opc_repository import (
|
||||||
|
OpcClientAlreadyExistsError,
|
||||||
|
OpcClientNotInitializedError,
|
||||||
|
OpcRepository,
|
||||||
|
OpcSessionAlreadyConnectedError,
|
||||||
|
is_reconnectable_opcua_bad,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -13,376 +24,561 @@ def mock_logger():
|
|||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def opc_repository(mock_logger):
|
def opc_repository(mock_logger):
|
||||||
return OpcRepository(
|
repository = OpcRepository(
|
||||||
id="test_repo",
|
opc_id='test_repo',
|
||||||
url="opc.tcp://localhost:4840",
|
server_name='test_server',
|
||||||
|
url='opc.tcp://localhost:4840',
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=Mock(),
|
notification_handler=Mock(),
|
||||||
reconnection_interval=60,
|
reconnection_interval=60,
|
||||||
server_uri="urn:test:server",
|
server_uri='urn:test:server',
|
||||||
cert_path="/path/to/cert.pem",
|
cert_path='/path/to/cert.pem',
|
||||||
private_key_path="/path/to/key.pem",
|
private_key_path='/path/to/key.pem',
|
||||||
server_cert_path="/path/to/server_cert.pem"
|
server_cert_path='/path/to/server_cert.pem',
|
||||||
|
metrics_controller=MagicMock(),
|
||||||
)
|
)
|
||||||
|
repository.disconnection_interval = 0.1
|
||||||
|
repository.send_notification = MagicMock()
|
||||||
|
repository.send_notification = MagicMock()
|
||||||
|
repository.emit_metric_sync = MagicMock()
|
||||||
|
repository.info = MagicMock()
|
||||||
|
repository.error = MagicMock()
|
||||||
|
repository.warning = MagicMock()
|
||||||
|
repository.debug = MagicMock()
|
||||||
|
repository._session_ready.set()
|
||||||
|
return repository
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_client():
|
def mock_client():
|
||||||
with patch('laborious.utils.repository.opc_repository.Client') as mock:
|
with patch('laborious.utils.repository.opc_repository.Client') as mock:
|
||||||
client_instance = AsyncMock()
|
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
|
||||||
mock.return_value = client_instance
|
mock.return_value = client_instance
|
||||||
yield client_instance
|
yield client_instance
|
||||||
|
|
||||||
|
|
||||||
metadata = {
|
metadata = {
|
||||||
"metadata": {
|
'metadata': {
|
||||||
"model_id": "test_model",
|
'model_id': 'test_model',
|
||||||
"model_name": "test_model",
|
'model_name': 'test_model',
|
||||||
"workflow_name": "test_workflow",
|
'workflow_name': 'test_workflow',
|
||||||
"schema_name": "test_schedule",
|
'schema_name': 'test_schedule',
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def test_init(opc_repository):
|
def test_init(opc_repository):
|
||||||
assert opc_repository.id == "test_repo"
|
assert opc_repository.id == 'test_repo'
|
||||||
assert opc_repository.url == "opc.tcp://localhost:4840"
|
assert opc_repository.server_name == 'test_server'
|
||||||
assert opc_repository.server_uri == "urn:test:server"
|
assert opc_repository.url == 'opc.tcp://localhost:4840'
|
||||||
assert opc_repository.cert_path == "/path/to/cert.pem"
|
assert opc_repository.server_uri == 'urn:test:server'
|
||||||
assert opc_repository.private_key_path == "/path/to/key.pem"
|
assert opc_repository.cert_path == '/path/to/cert.pem'
|
||||||
assert opc_repository.server_cert_path == "/path/to/server_cert.pem"
|
assert opc_repository.private_key_path == '/path/to/key.pem'
|
||||||
|
assert opc_repository.server_cert_path == '/path/to/server_cert.pem'
|
||||||
assert opc_repository.reconnection_interval == 60
|
assert opc_repository.reconnection_interval == 60
|
||||||
assert opc_repository.client is None
|
assert opc_repository.client is None
|
||||||
assert opc_repository.last_reconnection_time is None
|
assert opc_repository.last_reconnection_time is None
|
||||||
assert opc_repository.error_count == 0
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_set_security(opc_repository, mock_client):
|
||||||
async def test_set_security(opc_repository, mock_client):
|
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
await opc_repository.set_security()
|
opc_repository.set_security()
|
||||||
|
|
||||||
mock_client.application_uri = "urn:test:server"
|
mock_client.application_uri = 'urn:test:server'
|
||||||
mock_client.set_security.assert_called_once_with(
|
mock_client.set_security.assert_called_once_with(
|
||||||
SecurityPolicyBasic256,
|
security_policies.SecurityPolicyBasic256,
|
||||||
certificate="/path/to/cert.pem",
|
'/path/to/cert.pem',
|
||||||
private_key="/path/to/key.pem",
|
'/path/to/key.pem',
|
||||||
server_certificate="/path/to/server_cert.pem"
|
None,
|
||||||
|
'/path/to/server_cert.pem',
|
||||||
)
|
)
|
||||||
assert mock_client.secure_channel_timeout == 10000000
|
assert mock_client.aio_obj.secure_channel_timeout == 600_000
|
||||||
assert mock_client.session_timeout == 10000000
|
assert mock_client.aio_obj.session_timeout == 600_000
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_set_security_missing_certificates(opc_repository):
|
||||||
async def test_set_security_missing_certificates(opc_repository):
|
|
||||||
opc_repository.cert_path = None
|
opc_repository.cert_path = None
|
||||||
opc_repository.private_key_path = None
|
opc_repository.private_key_path = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await opc_repository.set_security()
|
opc_repository.set_security()
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
assert str(
|
assert str(e) == 'Certificate and private key paths must be provided for secure connection.'
|
||||||
e) == "Certificate and private key paths must be provided for secure connection."
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_set_security_missing_client(opc_repository):
|
||||||
async def test_connect_with_security(opc_repository, mock_client):
|
opc_repository.client = None
|
||||||
opc_repository.try_connect = AsyncMock(return_value=(True, {}))
|
try:
|
||||||
result = await opc_repository.connect()
|
opc_repository.set_security()
|
||||||
|
except ValueError as e:
|
||||||
|
assert str(e) == 'Client must be initialized before setting security'
|
||||||
|
|
||||||
opc_repository.try_connect.assert_called_once()
|
|
||||||
assert opc_repository.client == mock_client
|
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()
|
||||||
|
|
||||||
|
opc_repository._create_client.assert_called_once()
|
||||||
|
opc_repository._open_session.assert_called_once()
|
||||||
assert result == (True, {})
|
assert result == (True, {})
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_connect_without_security(opc_repository, mock_client):
|
||||||
async def test_connect_without_security(opc_repository, mock_client):
|
|
||||||
opc_repository.cert_path = None
|
opc_repository.cert_path = None
|
||||||
opc_repository.try_connect = AsyncMock(return_value=(True, {}))
|
opc_repository._create_client = MagicMock()
|
||||||
opc_repository.set_security = AsyncMock()
|
opc_repository._open_session = MagicMock(return_value=(True, {}))
|
||||||
result = await opc_repository.connect()
|
opc_repository.set_security = MagicMock()
|
||||||
|
result = opc_repository.connect()
|
||||||
|
|
||||||
opc_repository.try_connect.assert_called_once()
|
opc_repository._create_client.assert_called_once()
|
||||||
|
opc_repository._open_session.assert_called_once()
|
||||||
opc_repository.set_security.assert_not_called()
|
opc_repository.set_security.assert_not_called()
|
||||||
assert opc_repository.client == mock_client
|
|
||||||
assert result == (True, {})
|
assert result == (True, {})
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_connect_raises_when_session_already_open(opc_repository, mock_client):
|
||||||
async def test_try_connect_success(opc_repository):
|
opc_repository.client = mock_client
|
||||||
opc_repository.last_reconnection_time = None
|
proto = MagicMock()
|
||||||
opc_repository.client = AsyncMock()
|
proto.state = 'open'
|
||||||
result = await opc_repository.try_connect()
|
mock_client.aio_obj.uaclient.protocol = proto
|
||||||
|
|
||||||
|
with pytest.raises(OpcSessionAlreadyConnectedError, match='disconnect'):
|
||||||
|
opc_repository.connect()
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
open_proto = MagicMock()
|
||||||
|
open_proto.state = 'open'
|
||||||
|
open_proto.authentication_token = 'tok'
|
||||||
|
|
||||||
|
def connect_side_effect():
|
||||||
|
aio.uaclient.protocol = open_proto
|
||||||
|
|
||||||
|
opc_repository.client.connect = MagicMock(side_effect=connect_side_effect)
|
||||||
|
|
||||||
|
result = opc_repository._open_session()
|
||||||
|
|
||||||
opc_repository.client.connect.assert_called_once()
|
opc_repository.client.connect.assert_called_once()
|
||||||
assert opc_repository.last_reconnection_time is not None
|
|
||||||
assert result == (True, {})
|
assert result == (True, {})
|
||||||
|
assert opc_repository._session_ready.is_set()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_open_session_raises_when_already_connected(opc_repository, mock_client):
|
||||||
async def test_try_connect_fail(opc_repository):
|
opc_repository.client = mock_client
|
||||||
opc_repository.last_reconnection_time = None
|
proto = MagicMock()
|
||||||
|
proto.state = 'open'
|
||||||
|
mock_client.aio_obj.uaclient.protocol = proto
|
||||||
|
|
||||||
|
with pytest.raises(OpcSessionAlreadyConnectedError, match='disconnect'):
|
||||||
|
opc_repository._open_session()
|
||||||
|
|
||||||
|
|
||||||
|
def test_open_session_fail(opc_repository):
|
||||||
|
opc_repository._disconnect_locked = MagicMock()
|
||||||
opc_repository.client = MagicMock()
|
opc_repository.client = MagicMock()
|
||||||
opc_repository.client.connect.side_effect = Exception("Test error")
|
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'))
|
||||||
|
|
||||||
is_connected, error_data = await opc_repository.try_connect()
|
is_connected, error_data = opc_repository._open_session()
|
||||||
|
|
||||||
|
opc_repository._disconnect_locked.assert_called_once()
|
||||||
opc_repository.client.connect.assert_called_once()
|
opc_repository.client.connect.assert_called_once()
|
||||||
assert is_connected is False
|
assert is_connected is False
|
||||||
assert error_data['notification_id'] == f"OPC_CONNECTION_ERROR_{opc_repository.id}"
|
assert error_data['notification_id'] == f'OPC_CONNECTION_ERROR_{opc_repository.id}'
|
||||||
assert error_data['message'] == "Failed to connect to OPC server: Test error"
|
assert error_data['message'] == 'Failed to connect to OPC server: Test error'
|
||||||
assert error_data['block'] == "opc_repository"
|
assert error_data['block'] == 'opc_repository'
|
||||||
assert error_data['level'] == NotificationLevel.ERROR
|
assert error_data['level'] == NotificationLevel.ERROR
|
||||||
assert error_data['attachment_content'] is not None
|
assert error_data['attachment_content'] is not None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_open_session_raises_when_no_client(opc_repository):
|
||||||
async def test_disconnect(opc_repository, mock_client):
|
opc_repository.client = None
|
||||||
|
|
||||||
|
with pytest.raises(OpcClientNotInitializedError, match='not initialized'):
|
||||||
|
opc_repository._open_session()
|
||||||
|
|
||||||
|
|
||||||
|
def test_disconnection_fallback_success(opc_repository, mock_client):
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
await opc_repository.disconnect()
|
mock_client.disconnect.return_value = True
|
||||||
|
result = opc_repository._disconnection_fallback()
|
||||||
|
|
||||||
mock_client.disconnect.assert_called_once()
|
mock_client.disconnect.assert_called_once()
|
||||||
assert opc_repository.client is None
|
assert result == []
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_disconnection_fallback_fail(opc_repository, mock_client):
|
||||||
async def test_disconnect_no_client(opc_repository):
|
|
||||||
opc_repository.client = None
|
|
||||||
assert await opc_repository.disconnect() is None
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_disconnect_error(opc_repository, mock_client):
|
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
mock_client.disconnect.side_effect = Exception("Test error")
|
mock_client.disconnect.side_effect = Exception('Test error')
|
||||||
await opc_repository.disconnect()
|
result = opc_repository._disconnection_fallback()
|
||||||
|
assert result == [
|
||||||
|
{'attempt': 1, 'error': 'Test error', 'traceback': ANY},
|
||||||
|
{'attempt': 2, 'error': 'Test error', 'traceback': ANY},
|
||||||
|
{'attempt': 3, 'error': 'Test error', 'traceback': ANY},
|
||||||
|
{'attempt': 4, 'error': 'Test error', 'traceback': ANY},
|
||||||
|
{'attempt': 5, 'error': 'Test error', 'traceback': ANY},
|
||||||
|
]
|
||||||
|
assert mock_client.disconnect.call_count == 5
|
||||||
|
|
||||||
opc_repository.logger.custom_error.assert_called_once_with(
|
|
||||||
"Failed to disconnect from OPC server: Test error",
|
def test_disconnect(opc_repository, mock_client):
|
||||||
ANY
|
opc_repository.client = mock_client
|
||||||
|
opc_repository._disconnection_fallback = MagicMock(return_value=[])
|
||||||
|
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):
|
||||||
|
opc_repository.client = None
|
||||||
|
assert opc_repository.disconnect() is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_disconnect_error(opc_repository, mock_client):
|
||||||
|
opc_repository.client = mock_client
|
||||||
|
opc_repository._disconnection_fallback = MagicMock(
|
||||||
|
return_value=[{'attempt': 1, 'error': 'Test error', 'traceback': 'text'}]
|
||||||
|
)
|
||||||
|
opc_repository.disconnect()
|
||||||
|
|
||||||
|
opc_repository._disconnection_fallback.assert_called_once()
|
||||||
|
opc_repository.send_notification.assert_called_once_with(
|
||||||
|
metadata=opc_repository.metadata,
|
||||||
|
notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}',
|
||||||
|
message='Failed to disconnect from OPC server in 5 attempts.',
|
||||||
|
block='opc_repository',
|
||||||
|
level=NotificationLevel.ERROR,
|
||||||
|
attachment_content=json.dumps(
|
||||||
|
[{'attempt': 1, 'error': 'Test error', 'traceback': 'text'}], indent=4
|
||||||
|
),
|
||||||
)
|
)
|
||||||
assert opc_repository.client is None
|
assert opc_repository.client is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_validate_connection_none_client(opc_repository):
|
||||||
async def test_validate_connection_none_client(opc_repository):
|
|
||||||
opc_repository.client = None
|
opc_repository.client = None
|
||||||
opc_repository.connect = AsyncMock(return_value=(True, {}))
|
response = opc_repository.validate_connection()
|
||||||
response = await opc_repository.validate_connection()
|
assert response == (False, opc_repository._not_connected_error())
|
||||||
assert response == (True, {})
|
|
||||||
opc_repository.connect.assert_called_once()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_validate_connection_session_not_open(opc_repository):
|
||||||
async def test_validate_connection_error_count_disconnect_error(opc_repository):
|
|
||||||
opc_repository.error_count = 6
|
|
||||||
opc_repository.client = AsyncMock()
|
|
||||||
opc_repository.disconnect = AsyncMock(
|
|
||||||
side_effect=Exception("Test error")
|
|
||||||
)
|
|
||||||
opc_repository.connect = AsyncMock(return_value=(True, {}))
|
|
||||||
|
|
||||||
response = await opc_repository.validate_connection()
|
|
||||||
assert response == opc_repository.connect.return_value
|
|
||||||
opc_repository.disconnect.assert_called_once()
|
|
||||||
opc_repository.connect.assert_called_once()
|
|
||||||
opc_repository.logger.custom_error.assert_has_calls(
|
|
||||||
[
|
|
||||||
call("Failed to disconnect from OPC server: Test error", ANY),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_validate_connection_error_validate_connection_error(opc_repository):
|
|
||||||
opc_repository.client = MagicMock(
|
|
||||||
uaclient=Exception("Test error")
|
|
||||||
)
|
|
||||||
opc_repository.error_count = 0
|
|
||||||
|
|
||||||
response = await opc_repository.validate_connection()
|
|
||||||
|
|
||||||
assert response == (False, {
|
|
||||||
"notification_id": f"OPC_CONNECTION_CHECK_ERROR_{opc_repository.id}",
|
|
||||||
"message": "Failed to validate connection to OPC server: 'Exception' object has no attribute 'protocol'",
|
|
||||||
"block": "opc_repository",
|
|
||||||
"level": NotificationLevel.ERROR,
|
|
||||||
"attachment_content": ANY
|
|
||||||
})
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
@patch('laborious.utils.repository.opc_repository.datetime')
|
|
||||||
async def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, opc_repository):
|
|
||||||
_mock_datetime.now = MagicMock(
|
|
||||||
return_value=datetime(2025, 1, 1, 0, 0, 0))
|
|
||||||
opc_repository.error_count = 0
|
|
||||||
opc_repository.client = MagicMock()
|
opc_repository.client = MagicMock()
|
||||||
opc_repository.client.uaclient.protocol = None
|
opc_repository.client.aio_obj.uaclient.protocol = None
|
||||||
opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0)
|
|
||||||
opc_repository.connect = MagicMock(return_value=(True, {}))
|
|
||||||
|
|
||||||
response = await opc_repository.validate_connection()
|
response = opc_repository.validate_connection()
|
||||||
opc_repository.connect.assert_not_called()
|
|
||||||
assert response == (False, {
|
assert response == (False, opc_repository._not_connected_error())
|
||||||
"notification_id": f"OPC_CONNECTION_AWAITING_RECONNECTION_WINDOW_{opc_repository.id}",
|
opc_repository.error.assert_called_once()
|
||||||
"message": f"OPC server {opc_repository.id} is not connected, waiting for next reconnection window...",
|
|
||||||
"block": "opc_repository",
|
|
||||||
"level": NotificationLevel.WARNING
|
|
||||||
})
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_validate_connection_success(opc_repository):
|
||||||
@patch('laborious.utils.repository.opc_repository.datetime')
|
|
||||||
async def test_validate_connection_lost_time_to_reconnect(mock_datetime, opc_repository):
|
|
||||||
mock_datetime.now = MagicMock(
|
|
||||||
return_value=datetime(2025, 1, 1, 1, 0, 0))
|
|
||||||
opc_repository.error_count = 0
|
|
||||||
opc_repository.client = AsyncMock()
|
|
||||||
opc_repository.client.uaclient.protocol = None
|
|
||||||
opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0)
|
|
||||||
opc_repository.connect = AsyncMock(return_value=(True, {}))
|
|
||||||
|
|
||||||
response = await opc_repository.validate_connection()
|
|
||||||
opc_repository.connect.assert_called_once()
|
|
||||||
assert response == opc_repository.connect.return_value
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_validate_connection_success(opc_repository):
|
|
||||||
opc_repository.client = MagicMock()
|
opc_repository.client = MagicMock()
|
||||||
opc_repository.error_count = 0
|
proto = MagicMock()
|
||||||
opc_repository.client.uaclient.protocol = MagicMock()
|
proto.state = 'open'
|
||||||
opc_repository.client.uaclient.protocol.state = "open"
|
opc_repository.client.aio_obj.uaclient.protocol = proto
|
||||||
|
|
||||||
output = await opc_repository.validate_connection()
|
output = opc_repository.validate_connection()
|
||||||
assert output == (True, {})
|
assert output == (True, {})
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data_validate_connection_do_nothing(opc_repository):
|
||||||
async def test_write_data_validate_connection_do_nothing(opc_repository):
|
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
opc_repository.client = MagicMock(get_node=MagicMock())
|
||||||
opc_repository.client = AsyncMock(
|
mock_node = MagicMock()
|
||||||
get_node=MagicMock()
|
|
||||||
)
|
|
||||||
mock_node = AsyncMock()
|
|
||||||
opc_repository.client.get_node.return_value = mock_node
|
opc_repository.client.get_node.return_value = mock_node
|
||||||
|
|
||||||
result = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
result = opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
|
||||||
"float", opc_repository.logger, metadata['metadata'])
|
|
||||||
|
|
||||||
opc_repository.validate_connection.assert_called_once()
|
opc_repository.validate_connection.assert_called_once()
|
||||||
opc_repository.client.get_node.assert_called_once_with("ns=2;s=TestNode")
|
opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||||
assert result == (True, {})
|
assert result == (True, {'response_time': ANY})
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data_validate_connection_failed(opc_repository):
|
||||||
async def test_write_data_validate_connection_failed(opc_repository):
|
opc_repository.validate_connection = MagicMock(return_value=(False, {}))
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(False, {}))
|
opc_repository.client = MagicMock()
|
||||||
opc_repository.client = AsyncMock()
|
opc_repository._start_reconnect = MagicMock()
|
||||||
opc_repository.error_count = 0
|
|
||||||
|
|
||||||
result = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
is_success, error_data = opc_repository.write_data(
|
||||||
"float", opc_repository.logger, metadata['metadata'])
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
opc_repository.validate_connection.assert_called_once()
|
opc_repository.validate_connection.assert_called_once()
|
||||||
opc_repository.client.get_node.assert_not_called()
|
opc_repository.client.get_node.assert_not_called()
|
||||||
assert result == (False, {})
|
opc_repository._start_reconnect.assert_called_once()
|
||||||
|
assert opc_repository._start_reconnect.call_args.args[0] == 'ProtocolClosed'
|
||||||
|
assert is_success is False
|
||||||
|
assert error_data['opc_error_kind'] == 'connection_lost'
|
||||||
|
assert error_data['opc_status'] == 'ProtocolClosed'
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data_get_node_failed(opc_repository):
|
||||||
async def test_write_data_get_node_failed(opc_repository):
|
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
opc_repository.client = MagicMock()
|
||||||
opc_repository.client = AsyncMock()
|
opc_repository.client.get_node = MagicMock(side_effect=Exception('Test error'))
|
||||||
opc_repository.error_count = 0
|
|
||||||
opc_repository.client.get_node = MagicMock(
|
|
||||||
side_effect=Exception("Test error"))
|
|
||||||
|
|
||||||
is_success, error_data = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
is_success, error_data = opc_repository.write_data(
|
||||||
"float", opc_repository.logger, metadata['metadata'])
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
opc_repository.validate_connection.assert_called_once()
|
opc_repository.validate_connection.assert_called_once()
|
||||||
opc_repository.client.get_node.assert_called_once_with("ns=2;s=TestNode")
|
opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||||
assert is_success is False
|
assert is_success is False
|
||||||
assert error_data['notification_id'] == f"OPC_WRITE_GET_NODE_ERROR_{opc_repository.id}"
|
assert error_data['notification_id'] == f'OPC_WRITE_GET_NODE_ERROR_{opc_repository.id}'
|
||||||
assert error_data['message'] == "Failed to get node from OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
assert (
|
||||||
assert error_data['block'] == "opc_repository"
|
error_data['message']
|
||||||
|
== "Failed to get node from OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||||
|
)
|
||||||
|
assert error_data['block'] == 'opc_repository'
|
||||||
assert error_data['level'] == NotificationLevel.ERROR
|
assert error_data['level'] == NotificationLevel.ERROR
|
||||||
assert error_data['attachment_content'] is not None
|
assert error_data['attachment_content'] is not None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data_invalid_data_type(opc_repository, mock_client):
|
||||||
async def test_write_data_invalid_data_type(opc_repository, mock_client):
|
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
mock_node = AsyncMock()
|
mock_node = MagicMock()
|
||||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||||
|
|
||||||
is_success, error_data = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
is_success, error_data = opc_repository.write_data(
|
||||||
"invalid_type", opc_repository.logger, metadata['metadata'])
|
'ns=2;s=TestNode', 42.0, 'invalid_type', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
opc_repository.validate_connection.assert_called_once()
|
opc_repository.validate_connection.assert_called_once()
|
||||||
mock_client.get_node.assert_called_once_with("ns=2;s=TestNode")
|
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||||
|
|
||||||
assert is_success is False
|
assert is_success is False
|
||||||
assert error_data['notification_id'] == f"OPC_WRITE_DATA_TYPE_ERROR_{opc_repository.id}"
|
assert error_data['notification_id'] == f'OPC_WRITE_DATA_TYPE_ERROR_{opc_repository.id}'
|
||||||
assert error_data['message'] == "Unsupported data type: invalid_type | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
assert (
|
||||||
assert error_data['block'] == "opc_repository"
|
error_data['message']
|
||||||
|
== "Unsupported data type: invalid_type | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||||
|
)
|
||||||
|
assert error_data['block'] == 'opc_repository'
|
||||||
assert error_data['level'] == NotificationLevel.ERROR
|
assert error_data['level'] == NotificationLevel.ERROR
|
||||||
assert error_data.get('attachment_content') is None
|
assert error_data.get('attachment_content') is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data(opc_repository, mock_client):
|
||||||
@patch('laborious.utils.repository.opc_repository.metrics')
|
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||||
async def test_write_data(mock_metrics, opc_repository, mock_client):
|
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
mock_node = AsyncMock()
|
mock_node = MagicMock()
|
||||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||||
|
|
||||||
result = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
result = opc_repository.write_data('ns=2;s=TestNode', 42.0, 'float', metadata['metadata'])
|
||||||
"float", opc_repository.logger, metadata['metadata'])
|
|
||||||
|
|
||||||
mock_client.get_node.assert_called_once_with("ns=2;s=TestNode")
|
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||||
mock_node.write_value.assert_called_once()
|
mock_node.write_value.assert_called_once()
|
||||||
assert result == (True, {})
|
assert result == (True, {'response_time': ANY})
|
||||||
|
|
||||||
mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.assert_called_once_with(
|
|
||||||
pod_id=opc_repository.pod_id,
|
|
||||||
model_name=metadata['metadata']['model_name'],
|
|
||||||
pipeline_name=metadata['metadata']['workflow_name'],
|
|
||||||
opc_server_id=opc_repository.id
|
|
||||||
)
|
|
||||||
mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.return_value.inc.assert_called_once_with()
|
|
||||||
|
|
||||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.assert_called_once_with(
|
|
||||||
pod_id=opc_repository.pod_id,
|
|
||||||
model_name=metadata['metadata']['model_name'],
|
|
||||||
pipeline_name=metadata['metadata']['workflow_name'],
|
|
||||||
opc_server_id=opc_repository.id
|
|
||||||
)
|
|
||||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.return_value.observe.assert_called_once_with(
|
|
||||||
ANY)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_write_data_write_value_failed(opc_repository, mock_client):
|
||||||
async def test_write_data_write_value_failed(opc_repository, mock_client):
|
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
|
||||||
opc_repository.client = mock_client
|
opc_repository.client = mock_client
|
||||||
mock_node = AsyncMock()
|
mock_node = MagicMock()
|
||||||
opc_repository.error_count = 0
|
|
||||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||||
mock_node.write_value.side_effect = Exception("Test error")
|
mock_node.write_value.side_effect = Exception('Test error')
|
||||||
|
|
||||||
is_success, error_data = await opc_repository.write_data("ns=2;s=TestNode", 42.0,
|
is_success, error_data = opc_repository.write_data(
|
||||||
"float", opc_repository.logger, metadata['metadata'])
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
opc_repository.validate_connection.assert_called_once()
|
opc_repository.validate_connection.assert_called_once()
|
||||||
mock_client.get_node.assert_called_once_with("ns=2;s=TestNode")
|
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||||
mock_node.write_value.assert_called_once()
|
mock_node.write_value.assert_called_once()
|
||||||
assert is_success is False
|
assert is_success is False
|
||||||
assert error_data['notification_id'] == f"OPC_WRITE_DATA_ERROR_{opc_repository.id}"
|
assert error_data['notification_id'] == f'OPC_WRITE_DATA_ERROR_{opc_repository.id}'
|
||||||
assert error_data['message'] == "Failed to write data to OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
assert (
|
||||||
assert error_data['block'] == "opc_repository"
|
error_data['message']
|
||||||
|
== "Failed to write data to OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||||
|
)
|
||||||
|
assert error_data['block'] == 'opc_repository'
|
||||||
assert error_data['level'] == NotificationLevel.ERROR
|
assert error_data['level'] == NotificationLevel.ERROR
|
||||||
assert error_data['attachment_content'] is not None
|
assert error_data['attachment_content'] is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_is_reconnectable_opcua_bad():
|
||||||
|
assert is_reconnectable_opcua_bad(BadSessionIdInvalid()) is True
|
||||||
|
assert is_reconnectable_opcua_bad(BadNodeIdUnknown()) is False
|
||||||
|
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, {}))
|
||||||
|
opc_repository.client = mock_client
|
||||||
|
opc_repository._start_reconnect = MagicMock()
|
||||||
|
mock_node = MagicMock()
|
||||||
|
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||||
|
mock_node.write_value.side_effect = BadSessionIdInvalid()
|
||||||
|
|
||||||
|
is_success, error_data = opc_repository.write_data(
|
||||||
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_node.write_value.assert_called_once()
|
||||||
|
opc_repository._start_reconnect.assert_called_once()
|
||||||
|
assert is_success is False
|
||||||
|
assert error_data['opc_error_kind'] == 'session_bad'
|
||||||
|
assert error_data['opc_status'] == 'BadSessionIdInvalid'
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
is_success, error_data = opc_repository.write_data(
|
||||||
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
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):
|
||||||
|
opc_repository.last_reconnection_time = datetime.now()
|
||||||
|
opc_repository.reconnection_interval = 3600
|
||||||
|
|
||||||
|
opc_repository._start_reconnect('BadSessionIdInvalid', 'tok')
|
||||||
|
|
||||||
|
assert opc_repository._reconnect_thread is None
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
is_success, error_data = opc_repository.write_data(
|
||||||
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
opc_repository._start_reconnect.assert_called_once()
|
||||||
|
assert opc_repository._start_reconnect.call_args.args[0] == 'ProtocolClosed'
|
||||||
|
assert is_success is False
|
||||||
|
assert error_data['opc_error_kind'] == 'connection_lost'
|
||||||
|
assert error_data['opc_status'] == 'ProtocolClosed'
|
||||||
|
|
||||||
|
|
||||||
|
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.last_reconnection_time = datetime.now()
|
||||||
|
opc_repository.reconnection_interval = 3600
|
||||||
|
|
||||||
|
is_success, error_data = opc_repository.write_data(
|
||||||
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
assert opc_repository._reconnect_thread 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):
|
||||||
|
opc_repository._session_ready.clear()
|
||||||
|
opc_repository.reconnection_interval = 0
|
||||||
|
opc_repository.last_reconnection_time = None
|
||||||
|
opc_repository._reconnect_locked = MagicMock(
|
||||||
|
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)
|
||||||
|
assert opc_repository._reconnect_locked.call_count == 1
|
||||||
|
assert not opc_repository._reconnect_thread_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)
|
||||||
|
assert opc_repository._reconnect_locked.call_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
is_success, error_data = opc_repository.write_data(
|
||||||
|
'ns=2;s=TestNode', 42.0, 'float', metadata['metadata']
|
||||||
|
)
|
||||||
|
|
||||||
|
assert opc_repository._reconnect_thread 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, {}))
|
||||||
|
opc_repository.client = mock_client
|
||||||
|
opc_repository.reconnection_interval = 0
|
||||||
|
opc_repository.last_reconnection_time = None
|
||||||
|
mock_node = MagicMock()
|
||||||
|
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]
|
||||||
|
|
||||||
|
assert 1 <= opc_repository._start_reconnect.call_count <= 2
|
||||||
|
assert 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)
|
||||||
|
|
||||||
|
|
||||||
|
@patch('laborious.utils.repository.opc_repository.datetime')
|
||||||
|
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, {}))
|
||||||
|
|
||||||
|
result = opc_repository._reconnect_locked()
|
||||||
|
|
||||||
|
opc_repository._disconnect_locked.assert_called_once()
|
||||||
|
opc_repository._connect_locked.assert_called_once()
|
||||||
|
assert result == (True, {})
|
||||||
|
assert opc_repository.last_reconnection_time == datetime(2025, 1, 1, 12, 0, 0)
|
||||||
|
|||||||
@@ -1,14 +1,16 @@
|
|||||||
from os import environ
|
from os import environ
|
||||||
from laborious.utils.connectors_config import (build_mlflow_config,
|
|
||||||
|
from laborious.utils.connectors_config import (
|
||||||
|
build_minio_config,
|
||||||
|
build_mlflow_config,
|
||||||
build_opc_config,
|
build_opc_config,
|
||||||
build_postgres_config,
|
build_plugin_store_config,
|
||||||
build_mongodb_config)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_build_mlflow_config_with_env_vars():
|
def test_build_mlflow_config_with_env_vars():
|
||||||
# Arrange
|
# Arrange
|
||||||
environ['MLFLOW_HOST'] = 'http://test-host'
|
environ['MLFLOW_URL'] = 'http://test-host:8080'
|
||||||
environ['MLFLOW_PORT'] = '8080'
|
|
||||||
environ['MLFLOW_USERNAME'] = 'test-user'
|
environ['MLFLOW_USERNAME'] = 'test-user'
|
||||||
environ['MLFLOW_PASSWORD'] = 'test-pass'
|
environ['MLFLOW_PASSWORD'] = 'test-pass'
|
||||||
|
|
||||||
@@ -16,17 +18,25 @@ def test_build_mlflow_config_with_env_vars():
|
|||||||
config = build_mlflow_config()
|
config = build_mlflow_config()
|
||||||
|
|
||||||
# Assert
|
# Assert
|
||||||
assert config['host'] == 'http://test-host'
|
assert config['url'] == 'http://test-host:8080'
|
||||||
assert config['port'] == 8080
|
|
||||||
assert config['username'] == 'test-user'
|
assert config['username'] == 'test-user'
|
||||||
assert config['password'] == 'test-pass'
|
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():
|
def test_build_mlflow_config_with_defaults():
|
||||||
# Arrange
|
# Arrange
|
||||||
# Clear any existing env vars
|
# Clear any existing env vars
|
||||||
environ.pop('MLFLOW_HOST', None)
|
environ.pop('MLFLOW_URL', None)
|
||||||
environ.pop('MLFLOW_PORT', None)
|
|
||||||
environ.pop('MLFLOW_USERNAME', None)
|
environ.pop('MLFLOW_USERNAME', None)
|
||||||
environ.pop('MLFLOW_PASSWORD', None)
|
environ.pop('MLFLOW_PASSWORD', None)
|
||||||
|
|
||||||
@@ -34,12 +44,30 @@ def test_build_mlflow_config_with_defaults():
|
|||||||
config = build_mlflow_config()
|
config = build_mlflow_config()
|
||||||
|
|
||||||
# Assert
|
# Assert
|
||||||
assert config['host'] == 'http://localhost'
|
assert config['url'] == 'http://localhost:5080'
|
||||||
assert config['port'] == 5080
|
|
||||||
assert config['username'] == 'aignosi'
|
assert config['username'] == 'aignosi'
|
||||||
assert config['password'] == '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():
|
def test_build_opc_config_with_env_vars():
|
||||||
# Arrange
|
# Arrange
|
||||||
environ['OPC_CONFIG'] = '{"opc": {"name": "test-opc", "url": "opc.tcp://test:4840"}}'
|
environ['OPC_CONFIG'] = '{"opc": {"name": "test-opc", "url": "opc.tcp://test:4840"}}'
|
||||||
@@ -88,74 +116,37 @@ def test_build_opc_config_with_defaults():
|
|||||||
assert config['1']['reconnection_interval'] == 120
|
assert config['1']['reconnection_interval'] == 120
|
||||||
|
|
||||||
|
|
||||||
def test_build_postgres_config_with_env_vars():
|
def test_build_minio_config_with_env_vars():
|
||||||
# Arrange
|
environ['MINIO_ENDPOINT_URL'] = 'http://test-host'
|
||||||
environ['POSTGRES_HOST'] = 'test-host'
|
environ['MINIO_ACCESS_KEY'] = 'test-key'
|
||||||
environ['POSTGRES_PORT'] = '5433'
|
environ['MINIO_SECRET_KEY'] = 'test-secret'
|
||||||
environ['POSTGRES_USER'] = 'test-user'
|
environ['MINIO_REGION_NAME'] = 'test-region'
|
||||||
environ['POSTGRES_PASSWORD'] = 'test-pass'
|
environ['MINIO_DEFAULT_BUCKET'] = 'test-bucket'
|
||||||
environ['POSTGRES_DBNAME'] = 'test-db'
|
# Isolate from IDE/CI env (e.g. VS Code may export MINIO_SECURE=true).
|
||||||
environ['POSTGRES_MIN_CONNECTIONS'] = '10'
|
environ['MINIO_SECURE'] = 'false'
|
||||||
environ['POSTGRES_MAX_CONNECTIONS'] = '30'
|
assert build_minio_config() == {
|
||||||
|
'endpoint_url': 'http://test-host',
|
||||||
# Act
|
'access_key': 'test-key',
|
||||||
config = build_postgres_config()
|
'secret_key': 'test-secret',
|
||||||
|
'default_bucket': 'test-bucket',
|
||||||
# Assert
|
'retention_hours': 24,
|
||||||
assert config['host'] == 'test-host'
|
'secure': False,
|
||||||
assert config['port'] == 5433
|
|
||||||
assert config['user'] == 'test-user'
|
|
||||||
assert config['password'] == 'test-pass'
|
|
||||||
assert config['dbname'] == 'test-db'
|
|
||||||
assert config['min_connections'] == 10
|
|
||||||
assert config['max_connections'] == 30
|
|
||||||
|
|
||||||
|
|
||||||
def test_build_postgres_config_with_defaults():
|
|
||||||
# Arrange
|
|
||||||
environ.pop('POSTGRES_HOST', None)
|
|
||||||
environ.pop('POSTGRES_PORT', None)
|
|
||||||
environ.pop('POSTGRES_USER', None)
|
|
||||||
environ.pop('POSTGRES_PASSWORD', None)
|
|
||||||
environ.pop('POSTGRES_DBNAME', None)
|
|
||||||
environ.pop('POSTGRES_MIN_CONNECTIONS', None)
|
|
||||||
environ.pop('POSTGRES_MAX_CONNECTIONS', None)
|
|
||||||
|
|
||||||
# Act
|
|
||||||
config = build_postgres_config()
|
|
||||||
|
|
||||||
# Assert
|
|
||||||
assert config['host'] == 'localhost'
|
|
||||||
assert config['port'] == 5432
|
|
||||||
assert config['user'] == 'sientia'
|
|
||||||
assert config['password'] == 'sientia'
|
|
||||||
assert config['dbname'] == 'sientia'
|
|
||||||
assert config['min_connections'] == 5
|
|
||||||
assert config['max_connections'] == 20
|
|
||||||
|
|
||||||
|
|
||||||
def test_build_mongo_db_config_with_env_vars():
|
|
||||||
environ['MONGODB_USERNAME'] = 'sientia1'
|
|
||||||
environ['MONGODB_PASSWORD'] = 'sientia1'
|
|
||||||
environ['MONGODB_URL'] = 'localhost:27018'
|
|
||||||
environ['MONGODB_DATABASE_NAME'] = 'test_db'
|
|
||||||
environ['MONGODB_TTL_INDEX_HOURS'] = '1'
|
|
||||||
|
|
||||||
assert build_mongodb_config() == {
|
|
||||||
'connection_string': 'mongodb://sientia1:sientia1@localhost:27018',
|
|
||||||
'database_name': 'test_db',
|
|
||||||
'ttl_index_seconds': 3600
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def test_build_mongo_db_config_with_defaults():
|
def test_build_minio_config_with_defaults():
|
||||||
environ.pop('MONGODB_USERNAME', None)
|
environ.pop('MINIO_ENDPOINT_URL', None)
|
||||||
environ.pop('MONGODB_PASSWORD', None)
|
environ.pop('MINIO_ACCESS_KEY', None)
|
||||||
environ.pop('MONGODB_DATABASE_NAME', None)
|
environ.pop('MINIO_SECRET_KEY', None)
|
||||||
environ.pop('MONGODB_URL', None)
|
environ.pop('MINIO_REGION_NAME', None)
|
||||||
environ.pop('MONGODB_TTL_INDEX_HOURS', None)
|
environ.pop('MINIO_DEFAULT_BUCKET', None)
|
||||||
assert build_mongodb_config() == {
|
environ.pop('MINIO_SECURE', None)
|
||||||
'connection_string': 'mongodb://root:wKZDbMNU1c@localhost:27018',
|
environ.pop('MINIO_RETENTION_HOURS', None)
|
||||||
'database_name': 'sientia',
|
assert build_minio_config() == {
|
||||||
'ttl_index_seconds': 3600
|
'endpoint_url': 'http://localhost:9000',
|
||||||
|
'access_key': 'minioadmin',
|
||||||
|
'secret_key': 'minioadmin',
|
||||||
|
'default_bucket': 'laborious',
|
||||||
|
'retention_hours': 24,
|
||||||
|
'secure': False,
|
||||||
}
|
}
|
||||||
|
|||||||
12
tests/laborious/utils/test_dataframe_debug.py
Normal file
12
tests/laborious/utils/test_dataframe_debug.py
Normal file
@@ -0,0 +1,12 @@
|
|||||||
|
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
|
||||||
243
tests/laborious/worker/test_worker.py
Normal file
243
tests/laborious/worker/test_worker.py
Normal file
@@ -0,0 +1,243 @@
|
|||||||
|
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,9 +1,10 @@
|
|||||||
from unittest.mock import call, patch, AsyncMock, ANY
|
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||||
from pytest import mark, fixture
|
|
||||||
|
from pytest import fixture, mark
|
||||||
|
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||||
|
|
||||||
from laborious.activities.activities import Activities
|
from laborious.activities.activities import Activities
|
||||||
from laborious.workflows.sub_workflows.format_and_export_prediction import FormatAndExportPrediction
|
from laborious.workflows.sub_workflows.format_and_export_prediction import FormatAndExportPrediction
|
||||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
|
||||||
|
|
||||||
|
|
||||||
@fixture
|
@fixture
|
||||||
@@ -12,149 +13,667 @@ def format_and_export_prediction():
|
|||||||
|
|
||||||
|
|
||||||
metadata = {
|
metadata = {
|
||||||
"metadata": {
|
'metadata': {
|
||||||
"model_id": "test_model",
|
'model_id': 'test_model',
|
||||||
"model_name": "test_model",
|
'model_name': 'test_model',
|
||||||
"workflow_name": "test_workflow",
|
'workflow_name': 'test_workflow',
|
||||||
"schema_name": "test_schedule",
|
'schema_name': 'test_schedule',
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@mark.asyncio
|
||||||
@patch("laborious.workflows.sub_workflows.format_and_export_prediction.workflow", new_callable=AsyncMock)
|
@patch(
|
||||||
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
|
async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
|
||||||
|
|
||||||
input_data = {
|
input_data = {
|
||||||
'metadata': metadata,
|
'metadata': metadata,
|
||||||
"path_flag": None,
|
'path_flag': None,
|
||||||
"data": {"test": "data"},
|
'data': {'test': 'data'},
|
||||||
"timestamp": "2021-01-01",
|
'timestamp': '2021-01-01',
|
||||||
"model_id": 1,
|
'model_id': 1,
|
||||||
"prediction_confidence": 0,
|
'model_name': metadata['metadata']['model_name'],
|
||||||
"schema": "test_schema",
|
'prediction_confidence': 0,
|
||||||
"table_name": "test_table",
|
'schema': 'test_schema',
|
||||||
"opc_servers": ["test_server"],
|
'table_name': 'test_table',
|
||||||
"opc_output_config": {"test": "config"},
|
'opc_servers': ['test_server'],
|
||||||
"prediction_store_policy": "erl:1"
|
'opc_output_config': {'test': 'config'},
|
||||||
|
'prediction_store_policy': 'erl:1',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
prediction_data = MagicMock()
|
||||||
|
opc_metrics = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
(prediction_data, opc_metrics),
|
||||||
|
MagicMock(),
|
||||||
|
MagicMock(),
|
||||||
|
]
|
||||||
|
|
||||||
await format_and_export_prediction.run(input_data)
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
workflow_mock.execute_local_activity_method.assert_has_calls([
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.format_prediction,
|
Activities.format_prediction,
|
||||||
{
|
{
|
||||||
|
**metadata,
|
||||||
'data': input_data['data'],
|
'data': input_data['data'],
|
||||||
'timestamp': input_data['timestamp'],
|
'timestamp': input_data['timestamp'],
|
||||||
'model_id': input_data['model_id'],
|
'model_id': input_data['model_id'],
|
||||||
'prediction_confidence': input_data['prediction_confidence'],
|
'prediction_confidence': input_data['prediction_confidence'],
|
||||||
'prediction_store_policy': input_data['prediction_store_policy'],
|
'prediction_store_policy': input_data['prediction_store_policy'],
|
||||||
**metadata
|
'model_name': input_data['model_name'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
)])
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.write_opc_data,
|
Activities.write_opc_data,
|
||||||
{
|
{
|
||||||
'opc_output_config': input_data['opc_output_config'],
|
'opc_output_config': input_data['opc_output_config'],
|
||||||
'data': workflow_mock.execute_local_activity_method.return_value,
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
**metadata
|
**metadata,
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.export_data_to_postgres,
|
Activities.export_data_to_postgres,
|
||||||
{
|
{
|
||||||
|
**metadata,
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
'data': workflow_mock.execute_activity_method.return_value,
|
'data': prediction_data,
|
||||||
**metadata,
|
|
||||||
'timestamp_conversion': {
|
'timestamp_conversion': {
|
||||||
'column': 'timestamp',
|
'column': 'timestamp',
|
||||||
'format': DATETIME_FORMAT_WITH_TZ
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
}
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
)])
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'prediction': prediction_data,
|
||||||
|
'opc_metrics': opc_metrics,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
assert workflow_mock.execute_activity_method.call_count == 3
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
assert workflow_mock.execute_local_activity_method.call_count == 1
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@mark.asyncio
|
||||||
@patch("laborious.workflows.sub_workflows.format_and_export_prediction.workflow", new_callable=AsyncMock)
|
@patch(
|
||||||
async def test_run_default_path_flag(workflow_mock, format_and_export_prediction):
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
|
async def test_run_none_path_flag_with_transformed_data(
|
||||||
|
workflow_mock, format_and_export_prediction
|
||||||
|
):
|
||||||
|
# Arrange
|
||||||
input_data = {
|
input_data = {
|
||||||
'metadata': metadata,
|
'metadata': metadata,
|
||||||
"path_flag": "default",
|
'path_flag': None,
|
||||||
"data": {"test": "data"},
|
'data': {'test': 'data'},
|
||||||
"timestamp": "2021-01-01",
|
'transformed_data': {'transformed': 'data'},
|
||||||
"model_id": 1,
|
'timestamp': '2021-01-01',
|
||||||
"prediction_confidence": 0,
|
'model_id': 1,
|
||||||
"schema": "test_schema",
|
'model_name': metadata['metadata']['model_name'],
|
||||||
"table_name": "test_table",
|
'prediction_confidence': 0.9,
|
||||||
"opc_servers": ["test_server"],
|
'schema': 'test_schema',
|
||||||
"opc_output_config": {"test": "config"},
|
'table_name': 'test_table',
|
||||||
"comment": "test_comment"
|
'transform_table_name': 'test_transform_table',
|
||||||
|
'opc_servers': ['test_server'],
|
||||||
|
'opc_output_config': {'test': 'config'},
|
||||||
|
'prediction_store_policy': 'lts:1',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
prediction_data = MagicMock()
|
||||||
|
opc_metrics = MagicMock()
|
||||||
|
transformed_data = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_local_activity_method.side_effect = [
|
||||||
|
prediction_data, # format_prediction
|
||||||
|
transformed_data, # format_transformed_data
|
||||||
|
]
|
||||||
|
|
||||||
|
write_transformed_handler = AsyncMock()
|
||||||
|
workflow_mock.start_activity_method.return_value = write_transformed_handler
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
(prediction_data, opc_metrics), # write_opc_data
|
||||||
|
MagicMock(), # export_data_to_postgres (prediction)
|
||||||
|
MagicMock(), # write_metrics
|
||||||
|
]
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
|
# Assert - format_prediction call
|
||||||
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.format_prediction,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': input_data['data'],
|
||||||
|
'timestamp': input_data['timestamp'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'prediction_confidence': input_data['prediction_confidence'],
|
||||||
|
'prediction_store_policy': input_data['prediction_store_policy'],
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
Activities.format_transformed_data,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': input_data['transformed_data'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - start_activity_method for transformed data export
|
||||||
|
workflow_mock.start_activity_method.assert_called_once_with(
|
||||||
|
Activities.export_payload_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['transform_table_name'],
|
||||||
|
'data': transformed_data,
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - write_opc_data call
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_opc_data,
|
||||||
|
{
|
||||||
|
'opc_output_config': input_data['opc_output_config'],
|
||||||
|
'data': prediction_data,
|
||||||
|
**metadata,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - export_data_to_postgres for prediction call
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['table_name'],
|
||||||
|
'data': prediction_data,
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - write_metrics call
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'prediction': prediction_data,
|
||||||
|
'opc_metrics': opc_metrics,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - verify counts
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
|
assert workflow_mock.execute_local_activity_method.call_count == 2
|
||||||
|
assert workflow_mock.start_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch(
|
||||||
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
|
async def test_run_default_path_flag(workflow_mock, format_and_export_prediction):
|
||||||
|
input_data = {
|
||||||
|
'metadata': metadata,
|
||||||
|
'path_flag': 'default',
|
||||||
|
'data': {'test': 'data'},
|
||||||
|
'timestamp': '2021-01-01',
|
||||||
|
'model_id': 1,
|
||||||
|
'model_name': metadata['metadata']['model_name'],
|
||||||
|
'prediction_confidence': 0,
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'opc_servers': ['test_server'],
|
||||||
|
'opc_output_config': {'test': 'config'},
|
||||||
|
'comment': 'test_comment',
|
||||||
|
}
|
||||||
|
|
||||||
|
prediction_data = MagicMock()
|
||||||
|
opc_metrics = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
(prediction_data, opc_metrics),
|
||||||
|
MagicMock(),
|
||||||
|
MagicMock(),
|
||||||
|
]
|
||||||
|
|
||||||
await format_and_export_prediction.run(input_data)
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
workflow_mock.execute_local_activity_method.assert_has_calls([
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.format_default_prediction,
|
Activities.format_default_prediction,
|
||||||
{
|
{
|
||||||
|
**metadata,
|
||||||
'timestamp': input_data['timestamp'],
|
'timestamp': input_data['timestamp'],
|
||||||
'model_id': input_data['model_id'],
|
'model_id': input_data['model_id'],
|
||||||
'prediction_confidence': input_data['prediction_confidence'],
|
'prediction_confidence': input_data['prediction_confidence'],
|
||||||
'comment': input_data['comment'],
|
'comment': input_data['comment'],
|
||||||
**metadata
|
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.write_opc_data,
|
Activities.write_opc_data,
|
||||||
{
|
{
|
||||||
'opc_output_config': input_data['opc_output_config'],
|
'opc_output_config': input_data['opc_output_config'],
|
||||||
'data': workflow_mock.execute_local_activity_method.return_value,
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
**metadata
|
**metadata,
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.export_data_to_postgres,
|
Activities.export_data_to_postgres,
|
||||||
{
|
{
|
||||||
|
**metadata,
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
'data': workflow_mock.execute_activity_method.return_value,
|
'data': prediction_data,
|
||||||
**metadata,
|
|
||||||
'timestamp_conversion': {
|
'timestamp_conversion': {
|
||||||
'column': 'timestamp',
|
'column': 'timestamp',
|
||||||
'format': DATETIME_FORMAT_WITH_TZ
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
}
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'prediction': prediction_data,
|
||||||
|
'opc_metrics': opc_metrics,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch(
|
||||||
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
|
async def test_run_none_path_flag_with_pi_web_api(workflow_mock, format_and_export_prediction):
|
||||||
|
input_data = {
|
||||||
|
'metadata': metadata,
|
||||||
|
'path_flag': None,
|
||||||
|
'data': {'test': 'data'},
|
||||||
|
'timestamp': '2021-01-01',
|
||||||
|
'model_id': 1,
|
||||||
|
'model_name': metadata['metadata']['model_name'],
|
||||||
|
'prediction_confidence': 0,
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'pi_web_api_output_config': {
|
||||||
|
'endpoint': 'https://test-pi-server.com',
|
||||||
|
'prediction_tags': {},
|
||||||
|
'confidence_tags': {},
|
||||||
|
},
|
||||||
|
'prediction_store_policy': 'erl:1',
|
||||||
|
}
|
||||||
|
|
||||||
|
pi_web_api_data = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
pi_web_api_data, # write_pi_web_api_data
|
||||||
|
MagicMock(), # export_data_to_postgres
|
||||||
|
MagicMock(), # write_metrics
|
||||||
|
]
|
||||||
|
|
||||||
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.format_prediction,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': input_data['data'],
|
||||||
|
'timestamp': input_data['timestamp'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'prediction_confidence': input_data['prediction_confidence'],
|
||||||
|
'prediction_store_policy': input_data['prediction_store_policy'],
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_pi_web_api_data,
|
||||||
|
{
|
||||||
|
'pi_web_api_output_config': input_data['pi_web_api_output_config'],
|
||||||
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
|
**metadata,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['table_name'],
|
||||||
|
'data': pi_web_api_data,
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'prediction': pi_web_api_data,
|
||||||
|
'opc_metrics': {},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch(
|
||||||
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
|
async def test_run_none_path_flag_with_pi_web_api_and_opc(
|
||||||
|
workflow_mock, format_and_export_prediction
|
||||||
|
):
|
||||||
|
input_data = {
|
||||||
|
'metadata': metadata,
|
||||||
|
'path_flag': None,
|
||||||
|
'data': {'test': 'data'},
|
||||||
|
'timestamp': '2021-01-01',
|
||||||
|
'model_id': 1,
|
||||||
|
'model_name': metadata['metadata']['model_name'],
|
||||||
|
'prediction_confidence': 0,
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'opc_output_config': {'test': 'config'},
|
||||||
|
'pi_web_api_output_config': {
|
||||||
|
'endpoint': 'https://test-pi-server.com',
|
||||||
|
'prediction_tags': {},
|
||||||
|
'confidence_tags': {},
|
||||||
|
},
|
||||||
|
'prediction_store_policy': 'erl:1',
|
||||||
|
}
|
||||||
|
|
||||||
|
prediction_data = MagicMock()
|
||||||
|
pi_web_api_data = MagicMock()
|
||||||
|
opc_metrics = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
pi_web_api_data, # write_pi_web_api_data
|
||||||
|
(prediction_data, opc_metrics), # write_opc_data
|
||||||
|
MagicMock(), # export_data_to_postgres
|
||||||
|
MagicMock(), # write_metrics
|
||||||
|
]
|
||||||
|
|
||||||
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_pi_web_api_data,
|
||||||
|
{
|
||||||
|
'pi_web_api_output_config': input_data['pi_web_api_output_config'],
|
||||||
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
|
**metadata,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
Activities.write_opc_data,
|
||||||
|
{
|
||||||
|
'opc_output_config': input_data['opc_output_config'],
|
||||||
|
'data': pi_web_api_data,
|
||||||
|
**metadata,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['table_name'],
|
||||||
|
'data': prediction_data,
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'prediction': prediction_data,
|
||||||
|
'opc_metrics': opc_metrics,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 4
|
||||||
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch(
|
||||||
|
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
)
|
||||||
|
async def test_run_default_path_flag_with_pi_web_api(workflow_mock, format_and_export_prediction):
|
||||||
|
input_data = {
|
||||||
|
'metadata': metadata,
|
||||||
|
'path_flag': 'default',
|
||||||
|
'data': {'test': 'data'},
|
||||||
|
'timestamp': '2021-01-01',
|
||||||
|
'model_id': 1,
|
||||||
|
'model_name': metadata['metadata']['model_name'],
|
||||||
|
'prediction_confidence': 0,
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'pi_web_api_output_config': {
|
||||||
|
'endpoint': 'https://test-pi-server.com',
|
||||||
|
'prediction_tags': {},
|
||||||
|
'confidence_tags': {},
|
||||||
|
},
|
||||||
|
'comment': 'test_comment',
|
||||||
|
}
|
||||||
|
|
||||||
|
pi_web_api_data = MagicMock()
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.side_effect = [
|
||||||
|
pi_web_api_data, # write_pi_web_api_data
|
||||||
|
MagicMock(), # export_data_to_postgres
|
||||||
|
MagicMock(), # write_metrics
|
||||||
|
]
|
||||||
|
|
||||||
|
await format_and_export_prediction.run(input_data)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.write_pi_web_api_data,
|
||||||
|
{
|
||||||
|
'pi_web_api_output_config': input_data['pi_web_api_output_config'],
|
||||||
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
|
**metadata,
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['table_name'],
|
||||||
|
'data': pi_web_api_data,
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
'on_conflict': 'error',
|
||||||
|
'unique_columns': ['model_id', 'timestamp'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
assert workflow_mock.execute_activity_method.call_count == 3
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
assert workflow_mock.execute_local_activity_method.call_count == 1
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
248
tests/laborious/workflows/test_drift.py
Normal file
248
tests/laborious/workflows/test_drift.py
Normal file
@@ -0,0 +1,248 @@
|
|||||||
|
from unittest.mock import ANY, AsyncMock, call, patch
|
||||||
|
|
||||||
|
from pytest import fixture, mark
|
||||||
|
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||||
|
|
||||||
|
from laborious.activities.activities import Activities
|
||||||
|
from laborious.workflows.drift import Drift
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def drift() -> Drift:
|
||||||
|
return Drift()
|
||||||
|
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
'metadata': {
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'drift',
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run(workflow_mock: AsyncMock, drift: Drift):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'source_table_name': 'test_source_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'interval': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'drift_metrics': ['psi', 'ks'],
|
||||||
|
'chunk_period': 'hour',
|
||||||
|
}
|
||||||
|
|
||||||
|
target_name = input_data['model_config']['target']
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
reference_data = {'data': 'test_reference_data'}
|
||||||
|
drift_data = {'drift': 'test_drift_data'}
|
||||||
|
|
||||||
|
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await drift.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Check start_local_activity_method calls
|
||||||
|
# Query format matches psycopg2.sql output (identifiers with double quotes, literals with single quotes)
|
||||||
|
expected_gathering_query = f"""
|
||||||
|
SELECT *
|
||||||
|
FROM "{input_data['schema']}"."{input_data['source_table_name']}"
|
||||||
|
WHERE
|
||||||
|
model_id = '{input_data['model_id']}' AND
|
||||||
|
timestamp > NOW() - INTERVAL '{input_data['interval']} minutes'
|
||||||
|
ORDER BY timestamp ASC
|
||||||
|
"""
|
||||||
|
|
||||||
|
workflow_mock.start_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.load_custom_query,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'query': expected_gathering_query,
|
||||||
|
'datetime_columns': ['timestamp', 'created_at'],
|
||||||
|
'orient': 'records',
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
Activities.get_reference_data,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - Check calculate_drift call
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||||
|
Activities.calculate_drift,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'target_data': target_data,
|
||||||
|
'reference_data': reference_data,
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'target_name': target_name,
|
||||||
|
'drift_metrics': input_data['drift_metrics'],
|
||||||
|
'chunk_period': input_data['chunk_period'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Assert - Check export_data_to_postgres call
|
||||||
|
workflow_mock.execute_activity_method.assert_called_once_with(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': drift_data,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['target_table_name'],
|
||||||
|
'timestamp_conversion': {'column': 'timestamp', 'format': DATETIME_FORMAT_WITH_TZ},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_empty_target_data(workflow_mock: AsyncMock, drift: Drift):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'source_table_name': 'test_source_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'interval': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'drift_metrics': ['psi', 'ks'],
|
||||||
|
}
|
||||||
|
|
||||||
|
target_data = None
|
||||||
|
reference_data = {'data': 'test_reference_data'}
|
||||||
|
|
||||||
|
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock()
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock()
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await drift.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Should not call calculate_drift or export
|
||||||
|
workflow_mock.execute_activity_method.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_empty_drift_data(workflow_mock: AsyncMock, drift: Drift):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'source_table_name': 'test_source_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'interval': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'drift_metrics': ['psi', 'ks'],
|
||||||
|
}
|
||||||
|
|
||||||
|
target_name = input_data['model_config']['target']
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
reference_data = {'data': 'test_reference_data'}
|
||||||
|
drift_data = None
|
||||||
|
|
||||||
|
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock()
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await drift.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Should call calculate_drift but not export
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||||
|
Activities.calculate_drift,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'target_data': target_data,
|
||||||
|
'reference_data': reference_data,
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'target_name': target_name,
|
||||||
|
'drift_metrics': input_data['drift_metrics'],
|
||||||
|
'chunk_period': input_data.get('chunk_period', 'min'),
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_default_chunk_period(workflow_mock: AsyncMock, drift: Drift):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'source_table_name': 'test_source_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'interval': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'drift_metrics': ['psi', 'ks'],
|
||||||
|
# chunk_period not provided, should default to 'min'
|
||||||
|
}
|
||||||
|
|
||||||
|
target_name = input_data['model_config']['target']
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
reference_data = {'data': 'test_reference_data'}
|
||||||
|
drift_data = {'drift': 'test_drift_data'}
|
||||||
|
|
||||||
|
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await drift.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Check calculate_drift call with default chunk_period
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||||
|
Activities.calculate_drift,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'target_data': target_data,
|
||||||
|
'reference_data': reference_data,
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'target_name': target_name,
|
||||||
|
'drift_metrics': input_data['drift_metrics'],
|
||||||
|
'chunk_period': 'min', # Default value
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
@@ -1,5 +1,7 @@
|
|||||||
from unittest.mock import AsyncMock, MagicMock, call, patch, ANY
|
from unittest.mock import ANY, AsyncMock, call, patch
|
||||||
|
|
||||||
from pytest import fixture, mark
|
from pytest import fixture, mark
|
||||||
|
|
||||||
from laborious.activities.activities import Activities
|
from laborious.activities.activities import Activities
|
||||||
from laborious.workflows.minimal_retrain import MinimalRetrain
|
from laborious.workflows.minimal_retrain import MinimalRetrain
|
||||||
|
|
||||||
@@ -10,11 +12,11 @@ def minimal_retrain() -> MinimalRetrain:
|
|||||||
|
|
||||||
|
|
||||||
metadata = {
|
metadata = {
|
||||||
"metadata": {
|
'metadata': {
|
||||||
"model_id": "test_model_id",
|
'model_id': 'test_model_id',
|
||||||
"model_name": "test_model",
|
'model_name': 'test_model',
|
||||||
"workflow_name": "minimal_retrain",
|
'workflow_name': 'minimal_retrain',
|
||||||
"schedule_name": "test_schedule",
|
'schedule_name': 'test_schedule',
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -23,76 +25,300 @@ metadata = {
|
|||||||
@patch('laborious.workflows.minimal_retrain.workflow', new_callable=AsyncMock)
|
@patch('laborious.workflows.minimal_retrain.workflow', new_callable=AsyncMock)
|
||||||
async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
|
async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
|
||||||
input_data = {
|
input_data = {
|
||||||
"model_id": "test_model_id",
|
'model_id': 'test_model_id',
|
||||||
"model_name": "test_model",
|
'model_name': 'test_model',
|
||||||
"workflow_name": "minimal_retrain",
|
'workflow_name': 'minimal_retrain',
|
||||||
"schedule_name": "test_schedule",
|
'schedule_name': 'test_schedule',
|
||||||
"query": "test_query",
|
'query': 'test_query',
|
||||||
"schema": "test_schema",
|
'schema': 'test_schema',
|
||||||
"table_name": "test_table",
|
'table_name': 'test_table',
|
||||||
|
'model_config': {
|
||||||
|
'target': 'test_target',
|
||||||
|
'retention_minutes': 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
storage_result = {
|
||||||
|
'last_timestamp': '2024-01-01 00:00:00+0000',
|
||||||
|
'status': {'success': True},
|
||||||
|
'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}},
|
||||||
|
'bucket': None,
|
||||||
|
'object_key': None,
|
||||||
|
'object_prefix': None,
|
||||||
|
'uri': None,
|
||||||
}
|
}
|
||||||
|
|
||||||
workflow_mock.execute_activity_method = AsyncMock(
|
workflow_mock.execute_activity_method = AsyncMock(
|
||||||
return_value={
|
side_effect=[
|
||||||
"data1": "1",
|
storage_result,
|
||||||
"data2": "2",
|
{'success': True, 'experiment': 'test_experiment'},
|
||||||
}
|
{
|
||||||
|
'success': True,
|
||||||
|
'version': 'test_version',
|
||||||
|
'mlflow_run_id': 'test_mlflow_run_id',
|
||||||
|
'mlflow_experiment_id': 'test_mlflow_experiment_id',
|
||||||
|
},
|
||||||
|
{'report': 'test_report'},
|
||||||
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
await minimal_retrain.run(input_data)
|
await minimal_retrain.run(input_data)
|
||||||
|
|
||||||
workflow_mock.execute_local_activity_method.assert_has_calls(
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
[
|
[
|
||||||
call(
|
call(
|
||||||
Activities.load_custom_query,
|
Activities.load_query_with_minio_offload,
|
||||||
{
|
{
|
||||||
**metadata,
|
**metadata,
|
||||||
"query": input_data["query"],
|
'query': input_data['query'],
|
||||||
'datetime_columns': input_data.get('datetime_columns', [])
|
'datetime_columns': input_data.get('datetime_columns', []),
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.retrain_model,
|
Activities.retrain_model,
|
||||||
{
|
{
|
||||||
**metadata,
|
**metadata,
|
||||||
'data': workflow_mock.execute_local_activity_method.return_value,
|
'data': storage_result,
|
||||||
'model_name': input_data['model_name'],
|
'model_name': input_data['model_name'],
|
||||||
|
'model_config': input_data['model_config'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.update_production_model,
|
Activities.update_production_model,
|
||||||
{
|
{
|
||||||
**metadata,
|
**metadata,
|
||||||
'model_name': input_data['model_name'],
|
'model_name': input_data['model_name'],
|
||||||
'model_id': input_data['model_id'],
|
'success': True,
|
||||||
**workflow_mock.execute_activity_method.return_value,
|
'experiment': 'test_experiment',
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
|
|
||||||
workflow_mock.execute_activity_method.assert_has_calls([
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.format_retrain_report,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'experiment_response': {'success': True, 'experiment': 'test_experiment'},
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'update_report': {
|
||||||
|
'success': True,
|
||||||
|
'version': 'test_version',
|
||||||
|
'mlflow_run_id': 'test_mlflow_run_id',
|
||||||
|
'mlflow_experiment_id': 'test_mlflow_experiment_id',
|
||||||
|
},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.export_data_to_postgres,
|
Activities.export_data_to_postgres,
|
||||||
{
|
{
|
||||||
**metadata,
|
**metadata,
|
||||||
'data': workflow_mock.execute_activity_method.return_value,
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
)
|
)
|
||||||
])
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.minimal_retrain.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
|
||||||
|
input_data = {
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'minimal_retrain',
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'query': 'test_query',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'model_config': {
|
||||||
|
'target': 'test_target',
|
||||||
|
'retention_minutes': 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
storage_result = {
|
||||||
|
'last_timestamp': '2024-01-01 00:00:00+0000',
|
||||||
|
'status': {'success': True},
|
||||||
|
'data': {},
|
||||||
|
'bucket': None,
|
||||||
|
'object_key': None,
|
||||||
|
'object_prefix': None,
|
||||||
|
'uri': None,
|
||||||
|
}
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
storage_result,
|
||||||
|
{'success': True, 'experiment': 'test_experiment'},
|
||||||
|
{
|
||||||
|
'success': True,
|
||||||
|
'version': 'test_version',
|
||||||
|
'mlflow_run_id': 'test_mlflow_run_id',
|
||||||
|
'mlflow_experiment_id': 'test_mlflow_experiment_id',
|
||||||
|
},
|
||||||
|
{'report': 'test_report'},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
from pytest import raises
|
||||||
|
|
||||||
|
with raises(ValueError, match='No data returned from query'):
|
||||||
|
await minimal_retrain.run(input_data)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_called_once_with(
|
||||||
|
Activities.load_query_with_minio_offload,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'query': input_data['query'],
|
||||||
|
'datetime_columns': input_data.get('datetime_columns', []),
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_local_activity_method.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.minimal_retrain.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
|
||||||
|
input_data = {
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'minimal_retrain',
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'query': 'test_query',
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'table_name': 'test_table',
|
||||||
|
'model_config': {
|
||||||
|
'target': 'test_target',
|
||||||
|
'retention_minutes': 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
storage_result = {
|
||||||
|
'last_timestamp': '2024-01-01 00:00:00+0000',
|
||||||
|
'status': {'success': True},
|
||||||
|
'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}},
|
||||||
|
'bucket': None,
|
||||||
|
'object_key': None,
|
||||||
|
'object_prefix': None,
|
||||||
|
'uri': None,
|
||||||
|
}
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
storage_result,
|
||||||
|
{'success': False, 'experiment': 'test_experiment'},
|
||||||
|
{
|
||||||
|
'success': True,
|
||||||
|
'version': 'test_version',
|
||||||
|
'mlflow_run_id': 'test_mlflow_run_id',
|
||||||
|
'mlflow_experiment_id': 'test_mlflow_experiment_id',
|
||||||
|
},
|
||||||
|
{'report': 'test_report'},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await minimal_retrain.run(input_data)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.load_query_with_minio_offload,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'query': input_data['query'],
|
||||||
|
'datetime_columns': input_data.get('datetime_columns', []),
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.retrain_model,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': storage_result,
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_config': input_data['model_config'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.format_retrain_report,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'experiment_response': {'success': False, 'experiment': 'test_experiment'},
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'update_report': {},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['table_name'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 3
|
||||||
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
from unittest.mock import AsyncMock, call, patch, ANY
|
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||||
|
|
||||||
from pytest import fixture, mark
|
from pytest import fixture, mark
|
||||||
|
|
||||||
from laborious.activities.activities import Activities
|
from laborious.activities.activities import Activities
|
||||||
from laborious.workflows.predictions_batch import PredictionsBatch
|
from laborious.workflows.predictions_batch import PredictionsBatch
|
||||||
|
|
||||||
@@ -10,21 +12,19 @@ def predictions_batch() -> PredictionsBatch:
|
|||||||
|
|
||||||
|
|
||||||
metadata = {
|
metadata = {
|
||||||
"metadata": {
|
'model_id': 'test_model_id',
|
||||||
"model_id": "test_model_id",
|
'model_name': 'test_model',
|
||||||
"model_name": "test_model",
|
'workflow_name': 'predictions_batch',
|
||||||
"workflow_name": "predictions_batch",
|
'schedule_name': 'test_schedule',
|
||||||
"schedule_name": "test_schedule",
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@mark.asyncio
|
||||||
@patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock)
|
@patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock)
|
||||||
async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch):
|
async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch):
|
||||||
workflow_mock.execute_local_activity_method.return_value = {
|
activity_return = MagicMock()
|
||||||
'data': 'test_data'
|
workflow_mock.execute_activity_method.return_value = activity_return
|
||||||
}
|
|
||||||
input_data = {
|
input_data = {
|
||||||
'schedule_name': 'test_schedule',
|
'schedule_name': 'test_schedule',
|
||||||
'model_name': 'test_model',
|
'model_name': 'test_model',
|
||||||
@@ -32,57 +32,75 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
|
|||||||
'query': 'SELECT * FROM test',
|
'query': 'SELECT * FROM test',
|
||||||
'schema': 'test_schema',
|
'schema': 'test_schema',
|
||||||
'table_name': 'test_table',
|
'table_name': 'test_table',
|
||||||
|
'transform_table_name': 'test_transform_table',
|
||||||
'opc_output_config': 'test_opc_output_config',
|
'opc_output_config': 'test_opc_output_config',
|
||||||
|
'pi_web_api_output_config': 'test_pi_web_api_output_config',
|
||||||
'datetime_columns': ['timestamp', 'created_at'],
|
'datetime_columns': ['timestamp', 'created_at'],
|
||||||
'prediction_store_policy': 'erl:1',
|
'prediction_store_policy': 'erl:1',
|
||||||
'model_config': {
|
'model_config': {'retention': '30'},
|
||||||
'retention': '30'
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
await predictions_batch.run(input_data)
|
await predictions_batch.run(input_data)
|
||||||
|
|
||||||
workflow_mock.execute_local_activity_method.assert_has_calls([
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
call(
|
call(
|
||||||
Activities.load_custom_query,
|
Activities.load_query_with_minio_offload,
|
||||||
{
|
{
|
||||||
**metadata,
|
'metadata': metadata,
|
||||||
'query': input_data['query'],
|
'query': input_data['query'],
|
||||||
'datetime_columns': input_data.get('datetime_columns', [])
|
'datetime_columns': input_data.get('datetime_columns', []),
|
||||||
|
'model_name': input_data['model_name'],
|
||||||
},
|
},
|
||||||
retry_policy=ANY,
|
retry_policy=ANY,
|
||||||
start_to_close_timeout=ANY
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
])
|
|
||||||
prediction_input = {
|
prediction_input = {
|
||||||
'metadata': metadata,
|
'metadata': {'metadata': metadata},
|
||||||
'data': {'data': 'test_data'},
|
'data': activity_return,
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
|
'transform_table_name': input_data['transform_table_name'],
|
||||||
'model_id': input_data['model_id'],
|
'model_id': input_data['model_id'],
|
||||||
'model_name': input_data['model_name'],
|
'model_name': input_data['model_name'],
|
||||||
'input_filters': input_data.get('input_filters', {
|
'input_filters': input_data.get(
|
||||||
|
'input_filters',
|
||||||
|
{
|
||||||
'EMPTY_DATA': {
|
'EMPTY_DATA': {
|
||||||
'POLICY': 'STOP'
|
'POLICY': 'STOP',
|
||||||
|
'CONFIG': {},
|
||||||
}
|
}
|
||||||
}),
|
},
|
||||||
'mlflow_transform_filters': input_data.get('mlflow_transform_filters', {
|
),
|
||||||
|
'mlflow_transform_filters': input_data.get(
|
||||||
|
'mlflow_transform_filters',
|
||||||
|
{
|
||||||
'API_ERROR': {
|
'API_ERROR': {
|
||||||
'POLICY': 'STOP'
|
'POLICY': 'STOP',
|
||||||
|
'CONFIG': {},
|
||||||
}
|
}
|
||||||
}),
|
},
|
||||||
'mlflow_predict_filters': input_data.get('mlflow_predict_filters', {
|
),
|
||||||
|
'mlflow_predict_filters': input_data.get(
|
||||||
|
'mlflow_predict_filters',
|
||||||
|
{
|
||||||
'API_ERROR': {
|
'API_ERROR': {
|
||||||
'POLICY': 'STOP'
|
'POLICY': 'STOP',
|
||||||
|
'CONFIG': {},
|
||||||
}
|
}
|
||||||
}),
|
},
|
||||||
|
),
|
||||||
'model_config': input_data.get('model_config', {}),
|
'model_config': input_data.get('model_config', {}),
|
||||||
'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']),
|
'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']),
|
||||||
'opc_output_config': input_data.get('opc_output_config', {}),
|
'opc_output_config': input_data.get('opc_output_config', {}),
|
||||||
'prediction_store_policy': input_data.get('prediction_store_policy', 'erl:1')
|
'on_conflict': input_data.get('on_conflict', 'error'),
|
||||||
|
'pi_web_api_output_config': input_data.get('pi_web_api_output_config', {}),
|
||||||
|
'prediction_store_policy': input_data.get('prediction_store_policy', 'lts:1'),
|
||||||
|
'save_transform': input_data.get('save_transform', True),
|
||||||
}
|
}
|
||||||
|
|
||||||
workflow_mock.execute_child_workflow.assert_has_calls([
|
workflow_mock.execute_child_workflow.assert_has_calls(
|
||||||
call(
|
[call('subworkflow.prediction_process', prediction_input)]
|
||||||
'prediction_process', prediction_input)
|
)
|
||||||
])
|
|
||||||
|
|||||||
215
tests/laborious/workflows/test_simple_metrics.py
Normal file
215
tests/laborious/workflows/test_simple_metrics.py
Normal file
@@ -0,0 +1,215 @@
|
|||||||
|
from unittest.mock import ANY, AsyncMock, call, patch
|
||||||
|
|
||||||
|
from pytest import fixture, mark
|
||||||
|
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||||
|
|
||||||
|
from laborious.activities.activities import Activities
|
||||||
|
from laborious.workflows.simple_metrics import SimpleMetrics
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def simple_metrics() -> SimpleMetrics:
|
||||||
|
return SimpleMetrics()
|
||||||
|
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
'metadata': {
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'workflow_name': 'simple_metrics',
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.simple_metrics.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'interval_minutes': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'predictions_table_name': 'test_predictions_table',
|
||||||
|
'data_table_name': 'test_data_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'metrics': ['rmse', 'mse', 'mae', 'r2'],
|
||||||
|
}
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
simple_metrics_data = {'metrics': 'test_simple_metrics_data'}
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None])
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await simple_metrics.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Check load_custom_query call
|
||||||
|
# Query format matches psycopg2.sql output (identifiers with double quotes, literals with single quotes)
|
||||||
|
expected_query = f"""
|
||||||
|
select p."timestamp", p.prediction, ld.value as "target"
|
||||||
|
from "{input_data['schema']}"."{input_data['predictions_table_name']}" p
|
||||||
|
inner join "{input_data['schema']}"."{input_data['data_table_name']}" ld
|
||||||
|
on p."timestamp" = ld."timestamp"
|
||||||
|
where
|
||||||
|
p.model_id = '{input_data['model_id']}' and
|
||||||
|
p.prediction is not null and
|
||||||
|
ld.variable = '{input_data['model_config']['target']}' and
|
||||||
|
ld.value is not null and
|
||||||
|
p."timestamp" >= NOW() - INTERVAL '{input_data['interval_minutes']} minutes'
|
||||||
|
order by
|
||||||
|
p."timestamp" desc;
|
||||||
|
"""
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method.assert_has_calls(
|
||||||
|
[
|
||||||
|
call(
|
||||||
|
Activities.load_custom_query,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'query': expected_query,
|
||||||
|
'datetime_columns': ['timestamp'],
|
||||||
|
'orient': 'records',
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
call(
|
||||||
|
Activities.export_data_to_postgres,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'data': simple_metrics_data,
|
||||||
|
'schema': input_data['schema'],
|
||||||
|
'table_name': input_data['target_table_name'],
|
||||||
|
'timestamp_conversion': {
|
||||||
|
'column': 'timestamp',
|
||||||
|
'format': DATETIME_FORMAT_WITH_TZ,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||||
|
Activities.calculate_simple_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'target_data': target_data,
|
||||||
|
'metrics': input_data['metrics'],
|
||||||
|
'interval_minutes': input_data['interval_minutes'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.simple_metrics.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_empty_target_data(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'interval_minutes': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'predictions_table_name': 'test_predictions_table',
|
||||||
|
'data_table_name': 'test_data_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'metrics': ['rmse', 'mse'],
|
||||||
|
}
|
||||||
|
|
||||||
|
target_data = None
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(return_value=target_data)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await simple_metrics.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Should not call calculate_simple_metrics or export
|
||||||
|
assert workflow_mock.execute_activity_method.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.simple_metrics.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_empty_simple_metrics(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'interval_minutes': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'predictions_table_name': 'test_predictions_table',
|
||||||
|
'data_table_name': 'test_data_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
'metrics': ['rmse', 'mse'],
|
||||||
|
}
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
simple_metrics_data = None
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(return_value=target_data)
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await simple_metrics.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Should call calculate_simple_metrics but not export
|
||||||
|
workflow_mock.execute_activity_method.assert_called_once()
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('laborious.workflows.simple_metrics.workflow', new_callable=AsyncMock)
|
||||||
|
async def test_run_default_metrics(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
|
||||||
|
# Arrange
|
||||||
|
input_data = {
|
||||||
|
'schedule_name': 'test_schedule',
|
||||||
|
'model_name': 'test_model',
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'interval_minutes': 60,
|
||||||
|
'model_config': {'target': 'test_target'},
|
||||||
|
'schema': 'test_schema',
|
||||||
|
'predictions_table_name': 'test_predictions_table',
|
||||||
|
'data_table_name': 'test_data_table',
|
||||||
|
'target_table_name': 'test_target_table',
|
||||||
|
# metrics not provided, should default to ['rmse', 'mse', 'mae', 'r2']
|
||||||
|
}
|
||||||
|
|
||||||
|
target_data = {'data': 'test_target_data'}
|
||||||
|
simple_metrics_data = {'metrics': 'test_simple_metrics_data'}
|
||||||
|
|
||||||
|
workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None])
|
||||||
|
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
|
||||||
|
|
||||||
|
# Act
|
||||||
|
await simple_metrics.run(input_data)
|
||||||
|
|
||||||
|
# Assert - Check calculate_simple_metrics call with default metrics
|
||||||
|
workflow_mock.execute_activity_method.assert_any_call(
|
||||||
|
Activities.load_custom_query,
|
||||||
|
ANY,
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
|
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||||
|
Activities.calculate_simple_metrics,
|
||||||
|
{
|
||||||
|
**metadata,
|
||||||
|
'model_id': input_data['model_id'],
|
||||||
|
'target_data': target_data,
|
||||||
|
'metrics': ['rmse', 'mse', 'mae', 'r2'], # Default value
|
||||||
|
'interval_minutes': input_data['interval_minutes'],
|
||||||
|
},
|
||||||
|
retry_policy=ANY,
|
||||||
|
start_to_close_timeout=ANY,
|
||||||
|
)
|
||||||
376
values.yaml
376
values.yaml
@@ -1,159 +1,76 @@
|
|||||||
# Default values for sientia-module.
|
#
|
||||||
|
# Default values for sientia-laborious-worker using the sientia-module chart (0.6.x).
|
||||||
# This is a YAML-formatted file.
|
# This is a YAML-formatted file.
|
||||||
# Declare variables to be passed into your templates.
|
# Declare variables to be passed into your templates.
|
||||||
|
#
|
||||||
|
|
||||||
# This will set the replicaset count more information can be found here: https://kubernetes.io/docs/concepts/workloads/controllers/replicaset/
|
projectName: &projectName "sientia-laborious-worker"
|
||||||
replicaCount: 1
|
|
||||||
|
|
||||||
# This sets the container image more information can be found here: https://kubernetes.io/docs/concepts/containers/images/
|
# -----------------------------------------------------------------------------
|
||||||
image:
|
# Global configuration shared by all runtimes
|
||||||
repository: aignosi.azurecr.io/sientia-module-courier
|
# -----------------------------------------------------------------------------
|
||||||
# This sets the pull policy for images.
|
global:
|
||||||
|
namespace: sientia
|
||||||
|
|
||||||
|
image:
|
||||||
|
repository: aignosi.azurecr.io/sientia-module
|
||||||
pullPolicy: Always
|
pullPolicy: Always
|
||||||
# Overrides the image tag whose default is the chart appVersion.
|
tag: "1.2.0"
|
||||||
tag: "0.0.2"
|
|
||||||
|
|
||||||
0# This is for the secrets for pulling an image from a private repository more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/pull-image-private-registry/
|
commonLabels: {}
|
||||||
imagePullSecrets:
|
|
||||||
- name: docker-hub-secret
|
|
||||||
# This is to override the chart name.
|
|
||||||
nameOverride: "sientia-laborious-worker"
|
|
||||||
fullnameOverride: "sientia-laborious-worker"
|
|
||||||
namespace: sientia
|
|
||||||
|
|
||||||
# This section builds out the service account more information can be found here: https://kubernetes.io/docs/concepts/security/service-accounts/
|
resources:
|
||||||
serviceAccount:
|
# Resource limits and requests are important for ResourceBasedTuner to work correctly.
|
||||||
# Specifies whether a service account should be created
|
# The tuner monitors system CPU and memory usage, so proper resource limits must be set.
|
||||||
create: true
|
limits:
|
||||||
# Automatically mount a ServiceAccount's API credentials?
|
cpu: 2000m
|
||||||
automount: true
|
memory: 20Gi
|
||||||
# Annotations to add to the service account
|
requests:
|
||||||
annotations: {}
|
cpu: 1000m
|
||||||
# The name of the service account to use.
|
memory: 2Gi
|
||||||
# If not set and create is true, a name is generated using the fullname template
|
|
||||||
name: "sientia-laborious-worker"
|
|
||||||
|
|
||||||
# This is for setting Kubernetes Annotations to a Pod.
|
livenessProbe:
|
||||||
# For more information checkout: https://kubernetes.io/docs/concepts/overview/working-with-objects/annotations/
|
|
||||||
podAnnotations: {}
|
|
||||||
# This is for setting Kubernetes Labels to a Pod.
|
|
||||||
# For more information checkout: https://kubernetes.io/docs/concepts/overview/working-with-objects/labels/
|
|
||||||
podLabels: {}
|
|
||||||
|
|
||||||
podSecurityContext: {}
|
|
||||||
# fsGroup: 2000
|
|
||||||
|
|
||||||
securityContext: {}
|
|
||||||
# capabilities:
|
|
||||||
# drop:
|
|
||||||
# - ALL
|
|
||||||
# readOnlyRootFilesystem: true
|
|
||||||
# runAsNonRoot: true
|
|
||||||
# runAsUser: 1000
|
|
||||||
|
|
||||||
|
|
||||||
resources: {}
|
|
||||||
# We usually recommend not to specify default resources and to leave this as a conscious
|
|
||||||
# choice for the user. This also increases chances charts run on environments with little
|
|
||||||
# resources, such as Minikube. If you do want to specify resources, uncomment the following
|
|
||||||
# lines, adjust them as necessary, and remove the curly braces after 'resources:'.
|
|
||||||
# limits:
|
|
||||||
# cpu: 100m
|
|
||||||
# memory: 128Mi
|
|
||||||
# requests:
|
|
||||||
# cpu: 100m
|
|
||||||
# memory: 128Mi
|
|
||||||
|
|
||||||
# This is to setup the liveness and readiness probes more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/configure-liveness-readiness-startup-probes/
|
|
||||||
livenessProbe:
|
|
||||||
exec:
|
exec:
|
||||||
command:
|
command:
|
||||||
- sh
|
- sh
|
||||||
- -c
|
- -c
|
||||||
- pgrep -f "laborious.worker.worker"
|
- |
|
||||||
initialDelaySeconds: 20
|
curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1'
|
||||||
periodSeconds: 30
|
initialDelaySeconds: 1260
|
||||||
|
|
||||||
readinessProbe:
|
|
||||||
exec:
|
|
||||||
command:
|
|
||||||
- sh
|
|
||||||
- -c
|
|
||||||
- pgrep -f "laborious.worker.worker"
|
|
||||||
initialDelaySeconds: 10
|
|
||||||
periodSeconds: 15
|
periodSeconds: 15
|
||||||
|
timeoutSeconds: 5
|
||||||
|
failureThreshold: 3
|
||||||
|
|
||||||
|
readinessProbe:
|
||||||
|
exec:
|
||||||
|
command:
|
||||||
|
- sh
|
||||||
|
- -c
|
||||||
|
- |
|
||||||
|
curl -sf http://localhost:9090/metrics | grep -q '^app_up{.*} 1'
|
||||||
|
initialDelaySeconds: 1200
|
||||||
|
periodSeconds: 10
|
||||||
|
timeoutSeconds: 3
|
||||||
|
failureThreshold: 2
|
||||||
|
|
||||||
# This section is for setting up autoscaling more information can be found here: https://kubernetes.io/docs/concepts/workloads/autoscaling/
|
autoscaling:
|
||||||
autoscaling:
|
|
||||||
enabled: false
|
enabled: false
|
||||||
minReplicas: 1
|
minReplicas: 1
|
||||||
maxReplicas: 100
|
maxReplicas: 100
|
||||||
targetCPUUtilizationPercentage: 80
|
targetCPUUtilizationPercentage: 80
|
||||||
# targetMemoryUtilizationPercentage: 80
|
# targetMemoryUtilizationPercentage: 80
|
||||||
|
|
||||||
# Additional volumes on the output Deployment definition.
|
# Environment variables shared by all runtimes.
|
||||||
volumes: []
|
env:
|
||||||
# - name: foo
|
|
||||||
# secret:
|
|
||||||
# secretName: mysecret
|
|
||||||
# optional: false
|
|
||||||
|
|
||||||
# Additional volumeMounts on the output Deployment definition.
|
|
||||||
volumeMounts: []
|
|
||||||
# - name: foo
|
|
||||||
# mountPath: "/etc/foo"
|
|
||||||
# readOnly: true
|
|
||||||
|
|
||||||
nodeSelector: {}
|
|
||||||
|
|
||||||
tolerations: []
|
|
||||||
|
|
||||||
affinity: {}
|
|
||||||
|
|
||||||
services:
|
|
||||||
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:
|
|
||||||
# Se true, um recurso ServiceMonitor será criado.
|
|
||||||
enabled: true
|
|
||||||
# O intervalo no qual as métricas devem ser coletadas (ex: 30s, 1m).
|
|
||||||
endpoints:
|
|
||||||
- port: metrics
|
|
||||||
path: /metrics
|
|
||||||
interval: 30s
|
|
||||||
relabelings: []
|
|
||||||
- port: sdk-metrics
|
|
||||||
path: /metrics
|
|
||||||
interval: 30s
|
|
||||||
relabelings: []
|
|
||||||
|
|
||||||
additionalLabels:
|
|
||||||
release: kube-prometheus-stack
|
|
||||||
|
|
||||||
|
|
||||||
env:
|
|
||||||
# Entrypoint variables
|
# Entrypoint variables
|
||||||
- name: GITHUB_REPO_URL
|
- name: GITHUB_REPO_URL
|
||||||
value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git"
|
value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git"
|
||||||
- name: GITHUB_BRANCH
|
- name: GITHUB_BRANCH
|
||||||
value: SIENTIAPDE-1222-ajustar-a-library-para-fazer-o-download-do-courier
|
value: "release/SIENTIAPDE-1646"
|
||||||
- name: PYTHON_APP
|
- name: PYTHON_APP
|
||||||
value: "laborious.worker.worker"
|
value: "laborious.worker.worker"
|
||||||
|
- name: PYPI_SERVER
|
||||||
|
value: "http://library-distribution-server.library.svc.cluster.local:5000"
|
||||||
|
|
||||||
# Application variables
|
# Application variables
|
||||||
- name: POSTGRES_HOST
|
- name: POSTGRES_HOST
|
||||||
@@ -161,33 +78,52 @@ env:
|
|||||||
- name: POSTGRES_PORT
|
- name: POSTGRES_PORT
|
||||||
value: "5432"
|
value: "5432"
|
||||||
- name: POSTGRES_USER
|
- name: POSTGRES_USER
|
||||||
value: "sientia"
|
value: "postgres"
|
||||||
- name: POSTGRES_PASSWORD
|
- name: POSTGRES_PASSWORD
|
||||||
value: "sientia"
|
value: "nFqc81y6kwmr2zuAIx43DhiOosFCVPpeEfTtTWZflkNjB2j1KtEeIANkhFR9mAX3"
|
||||||
- name: POSTGRES_DBNAME
|
- name: POSTGRES_DBNAME
|
||||||
value: "sientia"
|
value: "sientia"
|
||||||
- name: POSTGRES_MIN_CONNECTIONS
|
- name: POSTGRES_MIN_CONNECTIONS
|
||||||
value: "10"
|
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
|
- name: POSTGRES_MAX_CONNECTIONS
|
||||||
value: "30"
|
value: "100"
|
||||||
|
|
||||||
- name: MLFLOW_HOST
|
- name: MLFLOW_URL
|
||||||
value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local"
|
value: "http://sientia-tracker-mlflow-tracking.sientia-tracker.svc.cluster.local:80"
|
||||||
- name: MLFLOW_PORT
|
|
||||||
value: "80"
|
|
||||||
- name: MLFLOW_USERNAME
|
- name: MLFLOW_USERNAME
|
||||||
value: "aignosi"
|
value: "aignosi"
|
||||||
- name: MLFLOW_PASSWORD
|
- name: MLFLOW_PASSWORD
|
||||||
value: "1L0FP50j3ncp123"
|
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
|
- name: OPC_ID
|
||||||
value: "1"
|
value: "1"
|
||||||
|
- name: OPC_SERVER_NAME
|
||||||
|
value: "default_server"
|
||||||
- name: OPC_URL
|
- name: OPC_URL
|
||||||
value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
|
value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
|
||||||
|
|
||||||
- name: KAFKA_BOOTSTRAP_SERVERS
|
|
||||||
value: "kafka.kafka.svc.cluster.local:9092"
|
|
||||||
|
|
||||||
- name: LOG_LEVEL
|
- name: LOG_LEVEL
|
||||||
value: "DEBUG"
|
value: "DEBUG"
|
||||||
- name: HTTP_METRICS_PORT
|
- name: HTTP_METRICS_PORT
|
||||||
@@ -213,17 +149,171 @@ env:
|
|||||||
- name: MONGODB_TTL_INDEX_HOURS
|
- name: MONGODB_TTL_INDEX_HOURS
|
||||||
value: "1"
|
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:
|
ssh:
|
||||||
enabled: true
|
enabled: true
|
||||||
secretName: git-ssh-key-sientia-laborious-worker
|
secretName: git-ssh-key-sientia-laborious-worker
|
||||||
sshPath: /mnt/.ssh
|
sshPath: /mnt/.ssh
|
||||||
knownHostsPath: /mnt/known_hosts
|
knownHostsPath: /mnt/known_hosts
|
||||||
|
|
||||||
# kubectl create secret docker-registry docker-hub-secret --namespace sientia --docker-server=http://aignosi.azurecr.io --docker-username=aignosi --docker-password=5I5zpQ6sRaHqX1hD3dr+2mo647yO3FRc359/wu6gsP+ACRDRz5mp
|
# 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 sientia/sientia-module -n sientia --create-namespace -f ./values.yaml --version 0.5.0
|
# 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 \
|
# kubectl create secret generic git-ssh-key-sientia-laborious-worker \
|
||||||
# --namespace sientia \
|
# --namespace sientia \
|
||||||
# --from-file=ssh-privatekey=git_key \
|
# --from-file=ssh-privatekey=git_key \
|
||||||
# --type=kubernetes.io/ssh-auth
|
# --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