Refactor ModelMetrics to utilize DriftAnalysis for drift detection - Replaced ModelAnalysis with DriftAnalysis in the ModelMetrics class to enhance drift detection capabilities. - Updated method signatures and documentation to reflect the changes in target_name and return values. - Adjusted data handling to ensure compatibility with the new analysis methods and improved clarity in the drift metrics dataframe preparation.
269 lines
9.7 KiB
Python
269 lines
9.7 KiB
Python
"""
|
|
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)
|
|
|
|
|
|
def insert_sample_data(postgres_engine: Engine, model_id: int, values: list[Any]) -> None:
|
|
"""
|
|
Replace laborious_data rows for a model_id with one row per value (sensor_1..n).
|
|
|
|
Args:
|
|
postgres_engine: SQLAlchemy engine.
|
|
model_id: Model id column value.
|
|
values: Per-sensor values; use string 'NULL' for SQL NULL.
|
|
"""
|
|
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}, '2024-01-01 12:00:00+00:00', '2024-01-01 12:00:00+00:00')
|
|
""")
|
|
insert_sql = f"""
|
|
INSERT INTO sientia_data.laborious_data (model_id, variable, value, timestamp, created_at)
|
|
VALUES
|
|
{', '.join(values_sql)}
|
|
"""
|
|
conn.execute(text(insert_sql))
|
|
|
|
|
|
def assert_prediction(
|
|
postgres_engine: Engine,
|
|
model_id: int,
|
|
prediction: float = 0.5,
|
|
prediction_confidence: int | Decimal = 0,
|
|
prediction_status: str = 'Good',
|
|
comments: str = '',
|
|
) -> 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 comments string.
|
|
"""
|
|
import pytest
|
|
|
|
with postgres_engine.connect() as conn:
|
|
result_query = conn.execute(
|
|
text(
|
|
f'SELECT model_id, prediction, prediction_confidence, prediction_status, comments '
|
|
f'FROM sientia_data.predictions WHERE model_id = {model_id} '
|
|
f'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]}"
|
|
assert row[4] == comments, f"Expected comments='{comments}', got {row[4]}"
|
|
|
|
|
|
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)
|
|
)
|
|
)
|
|
|
|
|