SIENTIAPDE-1712

Remove code validation script and refactor imports in activities and workflows

- Deleted the `validate.sh` script, which was responsible for running code quality checks.
- Cleaned up import statements in `activities.py`, `gates.py`, `mlflow.py`, and `storage.py` by removing unused imports and organizing them.
- Refactored initialization methods in `MinioManager` and `MLFlow` classes for improved readability.
- Updated various workflows to ensure compatibility with the new structure and removed unnecessary comments.
- Enhanced test cases to accommodate changes in the activities and workflows, ensuring proper mocking of dependencies.
This commit is contained in:
vitor-aignosi
2026-03-20 09:14:16 -03:00
parent 981ac700d4
commit 5d0d049082
25 changed files with 1224 additions and 705 deletions

View File

@@ -14,7 +14,6 @@ with workflow.unsafe.imports_passed_through():
from laborious.activities.model_metrics import ModelMetrics
from laborious.activities.opc import OPC
from laborious.activities.storage import Storage
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):

View File

@@ -13,17 +13,15 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ, now
from sientia_do.utils.formatters import create_sample_dict
from laborious import metrics
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from laborious.utils.filters.conditional_filters import (
filter_empty_data,
filter_specific_variables_null_values,
)
from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
# Strongly-typed filter function signatures
InputFilterFunc = Callable[[DataFrame, dict[str, Any]], bool]
@@ -106,7 +104,9 @@ class Gates(MinioManager):
Raises:
Exception: If BaseActivity initialization fails
"""
MinioManager.__init__(self, minio_repository, logger, notification_handler, metrics_controller)
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
def close(self) -> None:
"""
@@ -238,7 +238,7 @@ class Gates(MinioManager):
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
gate_type = input_data['type']
path_priority = input_data['path_priority']
@@ -326,7 +326,7 @@ class Gates(MinioManager):
self.info('Performing mlflow content gate...', metadata)
filters = input_data['filters']
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
@@ -375,7 +375,7 @@ class Gates(MinioManager):
self.info('Nothing was filtered by the mlflow content gate', metadata)
del data
return None, 0, ''
def get_prediction_store_policy(
@@ -479,7 +479,7 @@ class Gates(MinioManager):
minio_repo=self.minio_repository,
model_name=input_data['model_name'],
operation='transform',
workflow_metadata=metadata
workflow_metadata=metadata,
)
@activity.defn(name='format_prediction')
@@ -601,7 +601,6 @@ class Gates(MinioManager):
self.info(f'Default prediction formatted: {data.size} rows', metadata)
return data.to_dict()
@activity.defn(name='format_retrain_report')
async def format_retrain_report(self, input_data: dict[str, Any]) -> dict:
"""
@@ -672,7 +671,6 @@ class Gates(MinioManager):
return report.to_dict()
@activity.defn(name='write_metrics')
async def write_metrics(self, input_data: dict[str, Any]):
"""

View File

@@ -1,4 +1,3 @@
from re import M
from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
@@ -8,13 +7,12 @@ with workflow.unsafe.imports_passed_through():
from typing import Any
import numpy as np
from io import BytesIO
from pandas import DataFrame, read_parquet, to_datetime
from pandas import to_datetime
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.constants import (
DATETIME_FORMAT,
DATETIME_FORMAT_MS_WITH_TZ,
@@ -24,7 +22,6 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.utils.formatters import create_sample_dict
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from sientia_do.repository.minio_repository import MinioRepository
from laborious.utils.repository.model_repository import MLFlowRepository
@@ -72,7 +69,9 @@ class MLFlow(MinioManager):
Raises:
Exception: If MLFlowRepository initialization fails
"""
MinioManager.__init__(self, minio_repository, logger, notification_handler, metrics_controller)
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
self.mlflow_host = mlflow_host
self.mlflow_port = mlflow_port
self.mlflow_username = mlflow_username
@@ -128,7 +127,7 @@ class MLFlow(MinioManager):
"""
metadata = input_data['metadata']
self.info('Transforming data...', metadata)
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
@@ -251,7 +250,6 @@ class MLFlow(MinioManager):
self.info('Data predicted successfully', metadata)
if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
@@ -270,10 +268,9 @@ class MLFlow(MinioManager):
workflow_metadata=metadata,
status={
'success': True,
}
},
)
@activity.defn(name='retrain_model')
async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
@@ -313,22 +310,10 @@ class MLFlow(MinioManager):
metadata = input_data['metadata']
try:
if 'data' in input_data:
# New path: payload-based retrain input (inline or MinIO offloaded).
data = await MinioDataFramePayload.dataframe_from_wire(
input_data['data'],
self.minio_repository,
metadata,
)
else:
# Backward compatibility: legacy query_to_minio contract.
object_key = input_data['object_key']
self.info(f'Loading retrain data from Key: {object_key}', metadata)
file_bytes = await self.minio_repository.download_file(
object_name=object_key,
metadata=metadata,
)
data = read_parquet(BytesIO(file_bytes))
# Payload-based retrain input (inline dict or MinIO offloaded).
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
except Exception as e:
trace = traceback.format_exc()
await self.send_notification_async(

View File

@@ -1,15 +1,14 @@
import json
from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through():
# Extend the Temporal Postgres activities for convenient query -> MinIO export
import pickle
import traceback
from datetime import timedelta
from io import BytesIO
from os import getenv
from typing import Any
import pandas as pd
@@ -17,11 +16,11 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.activities.postgres import Postgres
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from sientia_do.repository.minio_repository import MinioRepository
_LOAD_QUERY_OFFLOAD_SKIP_KEYS = frozenset({'model_name', 'key_prefix', 'size_threshold_bytes'})
@@ -64,10 +63,14 @@ class Storage(Postgres, MinioManager):
metrics_controller=metrics_controller,
)
MinioManager.__init__(self, minio_repository, logger, notification_handler, metrics_controller)
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
@activity.defn(name='load_query_with_minio_offload')
async def load_query_with_minio_offload(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
async def load_query_with_minio_offload(
self, input_data: dict[str, Any]
) -> MinioDataFramePayload:
"""
Run the custom SQL load, then return a MinIO-aware dataframe wire dict.
@@ -92,7 +95,9 @@ class Storage(Postgres, MinioManager):
input_data,
)
if not rows:
self.error('load_query_with_minio_offload failed: No data returned from query', metadata)
self.error(
'load_query_with_minio_offload failed: No data returned from query', metadata
)
dataframe = None
else:
dataframe = pd.DataFrame(rows)
@@ -200,8 +205,6 @@ class Storage(Postgres, MinioManager):
)
return report
@activity.defn(name='query_to_minio')
async def query_to_minio(self, input_data: dict[str, Any]) -> dict[str, Any]:

View File

@@ -12,16 +12,16 @@ Otherwise, it is inlined as a Temporal-friendly ``dict``.
import pickle
import re
from dataclasses import dataclass, field
from collections.abc import Hashable
from dataclasses import dataclass
from datetime import datetime
from io import BytesIO
from os import getenv
from typing import Any, Hashable, Literal
from typing import Any, Literal
from pandas import DataFrame, read_parquet
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, DATETIME_FORMAT_WITH_TZ, now
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, DATETIME_FORMAT_WITH_TZ, now
# Keys that are part of the serialized wire format (not arbitrary metadata).
_SERIALIZED_FIELD_KEYS = frozenset({'data', 'bucket', 'object_key', 'object_prefix', 'uri'})
@@ -30,7 +30,9 @@ _OBJECT_TIMESTAMP_PATTERN = re.compile(
r'-(?:initial|transform)-(\d{4}-\d{2}-\d{2}_\d{2}-\d{2}-\d{2})\.parquet$'
)
OFFLOAD_THRESHOLD_BYTES = int(getenv('SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES', '1.5')) * 1024 * 1024
OFFLOAD_THRESHOLD_BYTES = int(
float(getenv('SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES', '1.5')) * 1024 * 1024
)
# Relative prefix used for storing offloaded training datasets in MinIO.
# It is also the root directory for retention cleanup listing.
@@ -79,7 +81,6 @@ class MinioDataFramePayload:
object_prefix: str | None = None
uri: str | None = None
@staticmethod
def estimate_size_bytes(df: DataFrame) -> int:
"""
@@ -117,21 +118,6 @@ class MinioDataFramePayload:
except ValueError:
return None
@staticmethod
def is_offloaded_dict(payload: dict[str, Any]) -> bool:
"""
Return True if the dict represents a MinIO-backed payload without inline data.
Args:
payload: Flat dict possibly produced by to_dict() / from_dataframe_to_dict().
Return:
bool: True when object_key is set and inline data is absent.
"""
if not payload.get('object_key'):
return False
return payload.get('data') is None
@staticmethod
def cleanup_prefix(self) -> str | None:
"""
@@ -141,6 +127,12 @@ class MinioDataFramePayload:
return self.object_prefix
return None
def has_data(self) -> bool:
"""
Return True if the payload has some data internally or in MinIO.
"""
return (self.data is not None and not self.data != {}) or self.object_key is not None
@classmethod
async def from_dataframe(
cls,
@@ -172,7 +164,9 @@ class MinioDataFramePayload:
"""
if not dataframe or dataframe.empty:
return cls(data=None, last_timestamp=now().strftime(DATETIME_FORMAT_WITH_TZ), status=status)
return cls(
data=None, last_timestamp=now().strftime(DATETIME_FORMAT_WITH_TZ), status=status
)
last_timestamp = max(dataframe['timestamp'].values.tolist())
@@ -207,7 +201,9 @@ class MinioDataFramePayload:
last_timestamp=last_timestamp,
)
async def retrieve(self, minio_repo: MinioRepository, workflow_metadata: dict[str, Any] | None = None) -> DataFrame:
async def retrieve(
self, minio_repo: MinioRepository, workflow_metadata: dict[str, Any] | None = None
) -> DataFrame:
"""
Load parquet from MinIO when object_key is set and populate inline data.
@@ -221,10 +217,11 @@ class MinioDataFramePayload:
if self.data is not None:
return DataFrame(self.data)
if self.data is None and self.object_key is None:
if not self.has_data():
return DataFrame()
file_bytes = await minio_repo.download_file(
object_name=self.object_key, metadata=workflow_metadata)
object_name=self.object_key, metadata=workflow_metadata
)
df = read_parquet(BytesIO(file_bytes))
return df

View File

@@ -1,14 +1,20 @@
from sientia_do.notifications.handlers import NotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.notifications.handlers import NotificationHandler
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository import MinioRepository
class MinioManager(SientiaMonitoring):
minio_repository: MinioRepository | None = None
def __init__(self, minio_repository: MinioRepository | None = None, logger: Logger | None = None, notification_handler: NotificationHandler | None = None, metrics_controller: MetricsController | None = None):
def __init__(
self,
minio_repository: MinioRepository | None = None,
logger: Logger | None = None,
notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
):
if self.minio_repository is None:
self.minio_repository = minio_repository
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
@@ -23,4 +29,4 @@ class MinioManager(SientiaMonitoring):
finally:
self.minio_repository = None
SientiaMonitoring.shutdown(self)
SientiaMonitoring.shutdown(self)

View File

@@ -44,7 +44,7 @@ class Drift:
model_id = '{input_data['model_id']}' AND
timestamp > NOW() - INTERVAL '{input_data['interval']} minutes'
ORDER BY timestamp ASC
"""
""" # nosec B608 - values come from internal Temporal workflow config, not user input
target_data_handler = workflow.start_local_activity_method(
Activities.load_custom_query,

View File

@@ -84,8 +84,8 @@ class MinimalRetrain:
start_to_close_timeout=timedelta(seconds=600),
)
if isinstance(storage_result, dict) and storage_result.get('success') is False:
return
if not storage_result.has_data():
raise ValueError('No data returned from query')
experiment_response = await workflow.execute_activity_method(
Activities.retrain_model,

View File

@@ -45,7 +45,7 @@ class SimpleMetrics:
p."timestamp" >= NOW() - INTERVAL '{interval_minutes} minutes'
order by
p."timestamp" desc;
"""
""" # nosec B608 - values come from internal Temporal workflow config, not user input
target_data = await workflow.execute_local_activity_method(
Activities.load_custom_query,

View File

@@ -2,7 +2,6 @@ from temporalio import workflow
with workflow.unsafe.imports_passed_through():
from datetime import timedelta
from collections.abc import Callable
from typing import Any
from sientia_do.temporal.policies import retry_policy
@@ -120,9 +119,8 @@ class PredictionProcess:
model_id: Any,
model_name: str,
model_config: dict[str, Any],
save_transform: bool
save_transform: bool,
) -> None:
last_timestamp = data.last_timestamp
# Apply input data quality gates
@@ -149,12 +147,7 @@ class PredictionProcess:
# Request MLFlow model transformation
response_data = await workflow.execute_local_activity_method(
Activities.request_transform,
{
**metadata,
'data': data,
'model_name': model_name,
'model_config': model_config
},
{**metadata, 'data': data, 'model_name': model_name, 'model_config': model_config},
retry_policy=retry_policy,
start_to_close_timeout=timedelta(minutes=5),
)

View File

@@ -1,3 +1,49 @@
import os
import sys
from unittest.mock import MagicMock
# The production code converts SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES 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_BYTES', '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.
@@ -6,9 +52,6 @@ during unit tests. The mock is registered in sys.modules before any test
imports are executed.
"""
import sys
from unittest.mock import MagicMock
# Mock sientia module
sientia_mock = MagicMock()
sientia_mock.ModelAnalysis = MagicMock

View File

@@ -17,9 +17,11 @@ from laborious.activities.storage import Storage
@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__(
mock_metrics_controller,
mock_minio_repository,
mock_api_init,
mock_model_metrics_init,
mock_gates_init,
@@ -43,6 +45,7 @@ def test___init__(
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
'retention_hours': 24,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
@@ -89,7 +92,8 @@ def test___init__(
dbname=postgres_config['dbname'],
min_connections=postgres_config['min_connections'],
max_connections=postgres_config['max_connections'],
minio_config=minio_config,
retention_hours=minio_config['retention_hours'],
minio_repository=mock_minio_repository.return_value,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
@@ -101,7 +105,7 @@ def test___init__(
mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'],
minio_config=minio_config,
minio_repository=mock_minio_repository.return_value,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
@@ -117,6 +121,7 @@ def test___init__(
mock_gates_init.assert_called_once_with(
ANY,
minio_repository=mock_minio_repository.return_value,
logger=logger,
notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
@@ -139,6 +144,16 @@ def test___init__(
metrics_controller=mock_metrics_controller.return_value,
)
mock_minio_repository.assert_called_once_with(
endpoint_url=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,
)
@mark.asyncio
@patch('laborious.activities.activities.Storage')
@@ -147,7 +162,9 @@ def test___init__(
@patch('laborious.activities.activities.Gates')
@patch('laborious.activities.activities.ModelMetrics')
@patch('laborious.activities.activities.API')
@patch('laborious.activities.activities.MinioRepository')
async def test_shutdown(
_mock_minio_repository,
mock_api_init,
mock_model_metrics_init,
mock_gates_init,
@@ -172,6 +189,7 @@ async def test_shutdown(
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
'retention_hours': 24,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}

View File

@@ -61,6 +61,60 @@ def base_input_data():
}
@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=AsyncMock(),
)
with patch.object(
SientiaMonitoring,
'get_core_labels',
return_value={
'pod_id': 'test_pod',
'model_name': 'test_model',
'operation_type': '-',
},
):
labels = api_instance.get_pi_web_api_core_labels(
metadata=metadata['metadata'], operation_type=None
)
assert 'operation_type' not in labels
@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=AsyncMock(),
)
with patch.object(
SientiaMonitoring,
'get_core_labels',
return_value={
'pod_id': 'test_pod',
'model_name': 'test_model',
'operation_type': 'write',
},
):
labels = api_instance.get_pi_web_api_core_labels(
metadata=metadata['metadata'], operation_type='write'
)
assert labels['operation_type'] == 'write'
def test__init__():
api = API(
base_url='https://test-pi-server.com',

View File

@@ -1,11 +1,29 @@
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pandas import DataFrame
from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel
from laborious.activities.gates import Gates
def _minio_payload(retrieve_return, status=None):
"""
Build a MinioDataFramePayload-like test double with async retrieve.
Args:
retrieve_return: Value returned from await retrieve(minio_repo, metadata).
status: Optional status dict for MLflow response gate (payload.status).
Return:
MagicMock: Object with async retrieve and optional status.
"""
p = MagicMock()
p.retrieve = AsyncMock(return_value=retrieve_return)
p.status = status
return p
@fixture
def gates_activity():
gates = Gates(
@@ -40,7 +58,7 @@ async def test_input_gate_invalid_filter(gates_activity):
input_data = {
**metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'value': [1, 2, 3]},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -65,7 +83,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': []},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -90,7 +108,7 @@ async def test_input_gate_no_filters(gates_activity):
input_data = {
**metadata,
'filters': {},
'data': {'value': [1, 2, 3]},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
}
@@ -108,7 +126,7 @@ async def test_input_gate_with_filter(gates_activity):
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': []},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -126,7 +144,7 @@ async def test_input_gate_with_filter_not_caught(gates_activity):
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': [1, 2, 3]},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -144,7 +162,10 @@ async def test_mlflow_response_gate_invalid_filter(gates_activity):
input_data = {
**metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}},
'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -169,7 +190,10 @@ async def test_mlflow_response_gate_filter_exception(
input_data = {
**metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}},
'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -195,7 +219,10 @@ async def test_mlflow_response_gate_no_filters(gates_activity):
input_data = {
**metadata,
'filters': {},
'data': {'content': {'message': 'success'}},
'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test',
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
}
@@ -214,10 +241,10 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
input_data = {
**metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}},
'data': {
'success': False,
'content': {'message': 'API error occurred', 'traceback': 'error trace'},
},
'data': _minio_payload(
{'content': {'message': 'API error occurred', 'traceback': 'error trace'}},
status={'success': False, 'message': 'API error occurred'},
),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -237,10 +264,10 @@ async def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
input_data = {
**metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}},
'data': {
'success': True,
'content': {'message': 'success'},
},
'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -259,10 +286,7 @@ async def test_mlflow_content_gate_invalid_filter(gates_activity):
input_data = {
**metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {
'success': True,
'content': {'message': 'success'},
},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -287,10 +311,7 @@ async def test_mlflow_content_gate_filter_exception(
input_data = {
**metadata,
'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': {
'success': False,
'content': {'message': 'API error occurred', 'traceback': 'error trace'},
},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -317,7 +338,7 @@ async def test_mlflow_content_gate_no_filters(gates_activity):
input_data = {
**metadata,
'filters': {},
'data': {'value': [1, 2, 3]},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
}
@@ -336,7 +357,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
input_data = {
**metadata,
'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}},
'data': {'value': [None, None, None]},
'data': _minio_payload(DataFrame({'value': [None, None, None]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -356,7 +377,7 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
input_data = {
**metadata,
'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -369,6 +390,22 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
gates_activity.debug.assert_called()
@mark.asyncio
async def test_mlflow_content_gate_filter_returns_false(gates_activity):
input_data = {
**metadata,
'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
result = await gates_activity.mlflow_content_gate(input_data)
assert result == (None, 0, '')
gates_activity.debug.assert_called()
def test_get_prediction_store_policy_invalid_policy(gates_activity):
# Arrange
prediction_store_policy = 'INVALID_POLICY'
@@ -430,10 +467,14 @@ async def test_format_prediction_no_timestamp(gates_activity):
# Arrange
input_data = {
**metadata,
'data': {
'prediction': {'2023-05-26 11:12:27': 1},
'response_time': {'2023-05-26 11:12:27': 0.1},
},
'data': _minio_payload(
DataFrame(
{
'prediction': {'2023-05-26 11:12:27': 1},
'response_time': {'2023-05-26 11:12:27': 0.1},
}
)
),
'model_id': 'test_model',
'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:1',
@@ -457,18 +498,22 @@ async def test_format_prediction_with_timestamp_erl(gates_activity):
# Arrange
input_data = {
**metadata,
'data': {
'prediction': {
'2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2,
'2023-05-26 11:12:29': 3,
},
'response_time': {
'2023-05-26 11:12:27': 0.1,
'2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3,
},
},
'data': _minio_payload(
DataFrame(
{
'prediction': {
'2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2,
'2023-05-26 11:12:29': 3,
},
'response_time': {
'2023-05-26 11:12:27': 0.1,
'2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3,
},
}
)
),
'model_id': 'test_model',
'prediction_confidence': 0.9,
'prediction_store_policy': 'erl:2',
@@ -492,18 +537,22 @@ async def test_format_prediction_with_timestamp_lts(gates_activity):
# Arrange
input_data = {
**metadata,
'data': {
'prediction': {
'2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2,
'2023-05-26 11:12:29': 3,
},
'response_time': {
'2023-05-26 11:12:27': 0.1,
'2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3,
},
},
'data': _minio_payload(
DataFrame(
{
'prediction': {
'2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2,
'2023-05-26 11:12:29': 3,
},
'response_time': {
'2023-05-26 11:12:27': 0.1,
'2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3,
},
}
)
),
'model_id': 'test_model',
'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:2',
@@ -527,11 +576,19 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
# Arrange
input_data = {
**metadata,
'data': {
'prediction': [1, 2, 3],
'response_time': [0.1, 0.2, 0.3],
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
},
'data': _minio_payload(
DataFrame(
{
'prediction': [1, 2, 3],
'response_time': [0.1, 0.2, 0.3],
'timestamp': [
'2023-05-26 11:12:27',
'2023-05-26 11:12:28',
'2023-05-26 11:12:29',
],
}
)
),
'model_id': 'test_model',
'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:2',
@@ -547,76 +604,106 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
@mark.asyncio
async def test_format_transformed_data_single_row(gates_activity):
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = {
**metadata,
'data': {
'var1': {'2023-05-26 11:12:27': 1.0},
'var2': {'2023-05-26 11:12:27': 2.0},
},
'data': _minio_payload(
DataFrame(
{
'var1': {'2023-05-26 11:12:27': 1.0},
'var2': {'2023-05-26 11:12:27': 2.0},
}
)
),
'model_id': 'test_model',
'model_name': 'test_model',
}
# Act
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert result['timestamp'] == {0: '2023-05-26 11:12:27', 1: '2023-05-26 11:12:27'}
assert result['variable'] == {0: 'var1', 1: 'var2'}
assert result['value'] == {0: 1.0, 1: 2.0}
assert result['model_id'] == {0: 'test_model', 1: 'test_model'}
assert result is payload_result
mock_from_dataframe.assert_called_once()
kwargs = mock_from_dataframe.call_args.kwargs
assert kwargs['model_name'] == 'test_model'
assert kwargs['operation'] == 'transform'
assert kwargs['workflow_metadata'] == metadata['metadata']
assert kwargs['minio_repo'] is gates_activity.minio_repository
assert 'dataframe' in kwargs
gates_activity.info.assert_called()
@mark.asyncio
async def test_format_transformed_data_multiple_rows(gates_activity):
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = {
**metadata,
'data': {
'var1': {
'2023-05-26 11:12:27': 1.0,
'2023-05-26 11:12:28': 2.0,
},
'var2': {
'2023-05-26 11:12:27': 3.0,
'2023-05-26 11:12:28': 4.0,
},
},
'data': _minio_payload(
DataFrame(
{
'var1': {
'2023-05-26 11:12:27': 1.0,
'2023-05-26 11:12:28': 2.0,
},
'var2': {
'2023-05-26 11:12:27': 3.0,
'2023-05-26 11:12:28': 4.0,
},
}
)
),
'model_id': 'test_model',
'model_name': 'test_model',
}
# Act
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert len(result['timestamp']) == 4
assert len(result['variable']) == 4
assert len(result['value']) == 4
assert len(result['model_id']) == 4
assert all(v == 'test_model' for v in result['model_id'].values())
assert set(result['variable'].values()) == {'var1', 'var2'}
assert result is payload_result
mock_from_dataframe.assert_called_once()
kwargs = mock_from_dataframe.call_args.kwargs
assert kwargs['model_name'] == 'test_model'
assert kwargs['operation'] == 'transform'
assert kwargs['workflow_metadata'] == metadata['metadata']
assert kwargs['minio_repo'] is gates_activity.minio_repository
assert 'dataframe' in kwargs
gates_activity.info.assert_called()
@mark.asyncio
async def test_format_transformed_data_empty_data(gates_activity):
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity):
# Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = {
**metadata,
'data': {},
'data': _minio_payload(DataFrame()),
'model_id': 'test_model',
'model_name': 'test_model',
}
# Act
result = await gates_activity.format_transformed_data(input_data)
# Assert
assert result['timestamp'] == {}
assert result['variable'] == {}
assert result['value'] == {}
assert result['model_id'] == {}
assert result is payload_result
mock_from_dataframe.assert_called_once()
kwargs = mock_from_dataframe.call_args.kwargs
assert kwargs['model_name'] == 'test_model'
assert kwargs['operation'] == 'transform'
assert kwargs['workflow_metadata'] == metadata['metadata']
assert kwargs['minio_repo'] is gates_activity.minio_repository
assert 'dataframe' in kwargs
gates_activity.info.assert_called()
@@ -711,31 +798,6 @@ async def test_format_retrain_report_failure(gates_activity):
gates_activity.debug.assert_called()
@mark.asyncio
async def test_get_last_timestamp_with_data(gates_activity):
# Arrange
input_data = {**metadata, 'data': {'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28']}}
# Act
result = await gates_activity.get_last_timestamp(input_data)
# Assert
assert result == '2023-05-26 11:12:28'
@mark.asyncio
async def test_get_last_timestamp_no_data(gates_activity):
# Arrange
input_data = {'data': {}, **metadata}
# Act
result = await gates_activity.get_last_timestamp(input_data)
# Assert
assert isinstance(result, str) # Should be a timestamp string
assert len(result) > 0
@mark.asyncio
@patch('laborious.activities.gates.metrics')
async def test_write_metrics(mock_metrics, gates_activity):

View File

@@ -11,21 +11,27 @@ from laborious.activities.mlflow import MLFlow
@patch('laborious.activities.mlflow.MLFlowRepository')
@patch('laborious.activities.mlflow.MinioRepository')
def test___init__(mock_minio_repository, mock_mlflow_repository):
logger = MagicMock()
notification_handler = MagicMock()
metrics_controller = AsyncMock()
minio_repo = mock_minio_repository(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
mlflow = MLFlow(
mlflow_host='http://localhost',
mlflow_port=5000,
mlflow_username='admin',
mlflow_password='admin',
minio_config={
'endpoint_url': 'http://localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
minio_repository=minio_repo,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
assert mlflow.mlflow_host == 'http://localhost'
@@ -52,21 +58,27 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
@patch('laborious.activities.mlflow.MLFlowRepository')
@patch('laborious.activities.mlflow.MinioRepository')
def mlflow(mock_minio_repository, mock_mlflow_repository):
logger = MagicMock()
notification_handler = MagicMock()
metrics_controller = AsyncMock()
minio_repo = mock_minio_repository(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
mlflow = MLFlow(
mlflow_host='http://localhost:5000',
mlflow_port=5000,
mlflow_username='admin',
mlflow_password='admin',
minio_config={
'endpoint_url': 'http://localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
minio_repository=minio_repo,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
mlflow.model_monitoring_repository = AsyncMock()
@@ -96,165 +108,161 @@ metadata = {
@mark.asyncio
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire',
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.max')
async def test_request_transform_success(mock_max, mock_dataframe_from_wire, mlflow):
mock_max.return_value = '2024-01-02'
async def test_request_transform_success(mock_from_dataframe, mlflow):
data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock
# Mock input data
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=data_mock)
input_data = {
**metadata,
'data': [
{
'timestamp': '2024-01-01',
'variable': 'var1',
'value': 1.0,
'created_at': '2024-01-01 12:00:00',
},
{
'timestamp': '2024-01-01',
'variable': 'var2',
'value': 2.0,
'created_at': '2024-01-01 12:00:00',
},
{
'timestamp': '2024-01-02',
'variable': 'var1',
'value': 3.0,
'created_at': '2024-01-02 12:00:00',
},
{
'timestamp': '2024-01-02',
'variable': 'var2',
'value': 4.0,
'created_at': '2024-01-02 12:00:00',
},
{
'timestamp': '2024-01-02',
'variable': 'var1',
'value': 1.0,
'created_at': '2024-01-01 12:00:00',
},
{
'timestamp': '2024-01-02',
'variable': 'var2',
'value': 1.0,
'created_at': '2024-01-01 12:00:00',
},
],
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
# Mock the transform response
expected_response = {'prediction': [0.5, 0.6], 'timestamp': ['2024-01-01', '2024-01-02']}
mlflow.model_monitoring_repository.transform.return_value = expected_response
transform_response = {'success': True, 'content': MagicMock()}
mlflow.model_monitoring_repository.transform.return_value = transform_response
data_mock.sort_values.return_value = data_mock
data_mock.drop_duplicates.return_value = data_mock
data_mock.pivot.return_value = data_mock
# Call the method
response_data = await mlflow.request_transform(input_data)
# Verify the data was correctly transformed
data_mock.pivot.assert_called_once_with(
index='timestamp', columns='variable', values='value'
)
data_mock.fillna.assert_called_once_with(np.nan, inplace=True)
# mock_dataframe.reset_index.assert_called_once()
data_mock.columns.name = None
# Verify the response
assert response_data == expected_response
# Verify the repository was called with correct arguments
mlflow.model_monitoring_repository.transform.assert_called_once_with(
'test_model', data_mock, {}, metadata['metadata']
)
mock_from_dataframe.assert_called_once()
assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire',
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime')
@patch('laborious.activities.mlflow.max')
async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe_from_wire, mlflow):
mock_max.return_value = '2024-01-02'
async def test_request_transform_failure(mock_from_dataframe, mlflow):
data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock
# Mock input data
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=data_mock)
input_data = {
**metadata,
'data': {
'variable': {
'2024-01-01': 'var1',
'2024-01-02': 'var2',
'2024-01-03': 'var1',
'2024-01-04': 'var2',
},
'value': {'2024-01-01': 1.0, '2024-01-02': 2.0, '2024-01-03': 3.0, '2024-01-04': 4.0},
},
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
# Mock the predict response
expected_response = {'prediction': [0.5, 0.6]}
mlflow.model_monitoring_repository.predict.return_value = expected_response
transform_response = {'success': False, 'message': 'Transform failed'}
mlflow.model_monitoring_repository.transform.return_value = transform_response
# Call the method
response_data = await mlflow.request_predict(input_data)
data_mock.sort_values.return_value = data_mock
data_mock.drop_duplicates.return_value = data_mock
data_mock.pivot.return_value = data_mock
data_mock.replace.assert_called_once_with(np.nan, None, inplace=True)
data_mock.__setitem__.assert_any_call(
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value
)
data_mock.__setitem__.assert_any_call(
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value
)
mock_to_datetime.assert_called_once_with(
data_mock.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ
)
mock_to_datetime.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT)
mock_to_datetime.assert_called_once_with(
data_mock.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ
)
mock_to_datetime.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT)
# Verify the response
assert response_data == expected_response
# Verify the repository was called with correct arguments
mlflow.model_monitoring_repository.predict.assert_called_once_with(
'test_model', data_mock, {}, metadata['metadata']
response_data = await mlflow.request_transform(input_data)
mock_from_dataframe.assert_called_once_with(
dataframe=None,
minio_repo=mlflow.minio_repository,
model_name='test_model',
operation='transform',
status=transform_response,
workflow_metadata=metadata['metadata'],
)
assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch('laborious.activities.mlflow.read_parquet')
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mock_read_parquet, mlflow):
async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
data_mock = MagicMock()
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=data_mock)
input_data = {
**metadata,
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
predict_response = {'success': True, 'content': MagicMock()}
mlflow.model_monitoring_repository.predict.return_value = predict_response
response_data = await mlflow.request_predict(input_data)
data_mock.replace.assert_called_once_with(np.nan, None, inplace=True)
mock_to_datetime.assert_called_once_with(
data_mock.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ
)
mock_to_datetime.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT)
mlflow.model_monitoring_repository.predict.assert_called_once_with(
'test_model', data_mock, {}, metadata['metadata']
)
mock_from_dataframe.assert_called_once()
assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime')
async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, mlflow):
data_mock = MagicMock()
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=data_mock)
input_data = {
**metadata,
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
predict_response = {'success': False, 'message': 'Predict failed'}
mlflow.model_monitoring_repository.predict.return_value = predict_response
response_data = await mlflow.request_predict(input_data)
mock_from_dataframe.assert_called_once_with(
dataframe=None,
minio_repo=mlflow.minio_repository,
model_name='test_model',
operation='predict',
status=predict_response,
workflow_metadata=metadata['metadata'],
)
assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlflow):
mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': True,
'experiment': 'test_experiment',
'message': 'Model retrained successfully.',
}
mlflow.minio_repository.download_file.return_value = b'parquet-bytes'
mock_read_parquet.return_value = MagicMock()
raw_data = MagicMock(columns=['variable', 'timestamp', 'value'])
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data)
response = await mlflow.retrain_model(
{
**metadata,
'object_key': 'test_object_key',
'data': payload,
'model_name': 'test_model',
'model_config': {
'target': 'target',
@@ -264,8 +272,6 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mock
}
)
raw_data = mock_read_parquet.return_value
timestamp = raw_data.__getitem__.return_value.max.return_value
raw_data.sort_values.assert_not_called()
@@ -317,22 +323,22 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mock
@mark.asyncio
@patch('laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire', new_callable=AsyncMock)
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_with_payload_data(
mock_to_datetime, mock_dataframe_from_wire, mlflow
):
async def test_retrain_model_success_with_payload_data(mock_to_datetime, mlflow):
mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': True,
'experiment': 'test_experiment',
'message': 'Model retrained successfully.',
}
mock_dataframe_from_wire.return_value = MagicMock()
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data)
response = await mlflow.retrain_model(
{
**metadata,
'data': {'data': {'a': [1]}},
'data': payload,
'model_name': 'test_model',
'model_config': {
'target': 'target',
@@ -347,24 +353,22 @@ async def test_retrain_model_success_with_payload_data(
@mark.asyncio
@patch('laborious.activities.mlflow.read_parquet')
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mock_read_parquet, mlflow):
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow):
mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': False,
'traceback': 'test_traceback',
'message': 'Model retrained failed.',
}
mlflow.minio_repository.download_file.return_value = b'parquet-bytes'
mock_read_parquet.return_value = MagicMock(
columns=['variable', 'timestamp', 'value', 'created_at']
)
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data)
response = await mlflow.retrain_model(
{
**metadata,
'object_key': 'test_object_key',
'data': payload,
'model_name': 'test_model',
'model_config': {
'target': 'target',
@@ -374,8 +378,6 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mock_re
}
)
raw_data = mock_read_parquet.return_value
timestamp = raw_data.__getitem__.return_value.max.return_value
raw_data.sort_values.assert_called_once_with('created_at', ascending=False)
@@ -439,14 +441,9 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mock_re
@mark.asyncio
async def test_retrain_model_data_error(mlflow):
mlflow.minio_repository.download_file.side_effect = Exception(
'Error loading retrain data'
)
response = await mlflow.retrain_model(
{
**metadata,
'object_key': 'test_object_key',
'model_name': 'test_model',
'model_config': {
'target': 'target',
@@ -458,7 +455,7 @@ async def test_retrain_model_data_error(mlflow):
assert response == {
'success': False,
'message': 'Error loading retrain data: Error loading retrain data',
'message': "Error loading retrain data: 'data'",
'traceback': ANY,
'timestamp': ANY,
}

View File

@@ -743,6 +743,101 @@ async def test_get_drift_metrics_univariate_error(
raise AssertionError('Expected Exception')
@mark.asyncio
@patch('laborious.activities.model_metrics.to_datetime')
@patch('laborious.activities.model_metrics.time.time')
@patch('laborious.activities.model_metrics.ModelAnalysis')
@patch('laborious.activities.model_metrics.metrics')
async 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:
await 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.assert_called_with(
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
)
else:
raise AssertionError('Expected Exception')
@mark.asyncio
@patch('laborious.activities.model_metrics.to_datetime')
@patch('laborious.activities.model_metrics.time.time')
@patch('laborious.activities.model_metrics.ModelAnalysis')
@patch('laborious.activities.model_metrics.metrics')
async 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:
await 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 getting drift metrics: Dataframe error', metadata['metadata']
)
model_metrics_activity.emit_metric.assert_called_with(
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
)
else:
raise AssertionError('Expected Exception')
@mark.asyncio
async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
# Arrange
@@ -987,3 +1082,27 @@ async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_me
model_metrics_activity.info.assert_called_once_with(
"Calculating simple metrics for model test_model_id: ['rmse', 'mae']", metadata['metadata']
)
@mark.asyncio
async 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(await model_metrics_activity.calculate_simple_metrics(input_data))
assert len(result['metric']) == 1
assert result['metric'].values[0] == 'rmse'

View File

@@ -29,13 +29,8 @@ def storage(mock_minio_repository):
dbname='postgres',
min_connections=1,
max_connections=10,
minio_config={
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
retention_hours=24,
minio_repository=mock_minio_repository.return_value,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
@@ -47,6 +42,7 @@ def test___init___not_hasattr(mock_minio_repository):
logger = MagicMock()
notification_handler = MagicMock()
metrics_controller = AsyncMock()
minio_repo = mock_minio_repository.return_value
storage = Storage(
host='localhost',
port=5432,
@@ -55,28 +51,16 @@ def test___init___not_hasattr(mock_minio_repository):
dbname='postgres',
min_connections=1,
max_connections=10,
minio_config={
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
retention_hours=24,
minio_repository=minio_repo,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
assert isinstance(storage, Postgres)
mock_minio_repository.assert_called_once_with(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
assert storage.minio_repository is minio_repo
mock_minio_repository.assert_not_called()
@patch('laborious.activities.storage.MinioRepository')
@@ -93,27 +77,15 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
dbname='postgres',
min_connections=1,
max_connections=10,
minio_config={
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
retention_hours=24,
minio_repository=None,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
mock_minio_repository.assert_called_once_with(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
assert storage.minio_repository is None
mock_minio_repository.assert_not_called()
@patch('laborious.activities.storage.MinioRepository')
@@ -126,13 +98,8 @@ def test___init___done_repository(mock_minio_repository, storage):
dbname='postgres',
min_connections=1,
max_connections=10,
minio_config={
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
retention_hours=24,
minio_repository=mock_minio_repository.return_value,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
@@ -231,55 +198,58 @@ def test___del__(storage):
storage.close.assert_called_once()
def test_estimate_payload_size_bytes(storage):
assert storage._estimate_payload_size_bytes({'x': 1}) > 0
@mark.asyncio
async def test_load_query_with_minio_offload_no_rows(storage):
storage.load_custom_query = AsyncMock(return_value=None)
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result['success'] is False
storage_result = {'success': False}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result == storage_result
mock_from_dataframe.assert_awaited_once()
@mark.asyncio
async def test_load_query_with_minio_offload_inline(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'my-model', 'key_prefix': 'predictions/s'}
)
assert result.get('success') is True
assert 'data' in result
assert result.get('object_key') is None
storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = await 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_awaited_once()
@mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.MinioDataFramePayload.estimate_size_bytes')
async def test_load_query_with_minio_offload_minio(mock_estimate, storage):
mock_estimate.return_value = 10**9
async def test_load_query_with_minio_offload_minio(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage.minio_repository.upload_file = AsyncMock(
return_value={
'minio_object_name': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-15_12-30-45.parquet'
}
)
storage.minio_repository.bucket = 'test'
fixed = datetime.datetime(2024, 1, 15, 12, 30, 45)
with patch('laborious.utils.models.minio_dataframe_payload.now', return_value=fixed):
storage_result = {'success': True, 'data': None, 'object_key': 'object-key'}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result.get('success') is True
assert result.get('data') is None
assert (
result['object_key']
== 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-15_12-30-45.parquet'
)
storage.minio_repository.upload_file.assert_called_once()
assert result == storage_result
mock_from_dataframe.assert_awaited_once()
@mark.asyncio
@@ -297,12 +267,117 @@ async def test_cleanup_minio_objects_expired(mock_now, storage):
storage.send_notification_async = AsyncMock()
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'prefixes': ['training_datasets/m']}
{**metadata, 'prefix': 'training_datasets/m'}
)
assert result['success'] is True
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'],
)
@mark.asyncio
async def test_load_query_with_minio_offload_minio_not_initialized(storage):
storage.minio_repository = None
with raises(ValueError, match='Minio repository not initialized'):
await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm'}
)
@mark.asyncio
async def test_export_payload_to_postgres(storage):
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=MagicMock())
storage.export_data_to_postgres = AsyncMock(return_value={'success': True})
result = await storage.export_payload_to_postgres(
{**metadata, 'data': payload, 'schema': 'public', 'table': 't'}
)
payload.retrieve.assert_awaited_once_with(storage.minio_repository, metadata['metadata'])
storage.export_data_to_postgres.assert_awaited_once()
assert result == {'success': True}
@mark.asyncio
async def test_cleanup_minio_objects_expired_minio_not_initialized(storage):
storage.minio_repository = None
with raises(ValueError, match='Minio repository not initialized'):
await storage.cleanup_minio_objects_expired({**metadata, 'prefix': 'test'})
@mark.asyncio
@patch('laborious.activities.storage.now')
async def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = AsyncMock(
return_value=['some/random/key-without-timestamp.parquet']
)
storage.minio_repository.delete_file = AsyncMock()
storage.send_notification_async = AsyncMock()
result = await storage.cleanup_minio_objects_expired({**metadata, 'prefix': 'test'})
assert result['deleted_count'] == 0
assert result['failed_count'] == 0
storage.minio_repository.delete_file.assert_not_called()
@mark.asyncio
@patch('laborious.activities.storage.now')
async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
old_key = 'training_datasets/m/m-initial-2024-12-01_00-00-00.parquet'
storage.minio_repository.list_objects = AsyncMock(return_value=[old_key])
storage.minio_repository.delete_file = AsyncMock(side_effect=Exception('delete error'))
storage.send_notification_async = AsyncMock()
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'prefix': 'training_datasets/m'}
)
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'
@mark.asyncio
@patch('laborious.activities.storage.now')
async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = AsyncMock(side_effect=Exception('list error'))
storage.send_notification_async = AsyncMock()
storage.error = MagicMock()
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'prefix': 'training_datasets/m'}
)
assert result['deleted_count'] == 0
assert result['failed_count'] == 0
storage.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
message='Error cleaning up MinIO objects: list error',
block='cleanup_minio_objects_expired',
level=NotificationLevel.ERROR,
attachment_content=ANY,
)
storage.error.assert_called_once()

View File

@@ -1,8 +1,14 @@
from datetime import datetime
from io import BytesIO
from unittest.mock import AsyncMock, MagicMock, patch
from pytest import mark
import pytest
from pandas import DataFrame
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from laborious.utils.models.minio_dataframe_payload import (
MinioDataFramePayload,
_build_object_key,
)
def test_parse_object_timestamp_hyphenated_model():
@@ -21,34 +27,190 @@ def test_parse_object_timestamp_invalid():
assert MinioDataFramePayload.parse_object_timestamp('bad.parquet') is None
def test_is_offloaded_dict_true_false():
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': None}) is True
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': {}}) is False
assert MinioDataFramePayload.is_offloaded_dict({'data': {}}) is False
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_from_payload_dict():
p = {
'object_key': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-01_00-00-00.parquet',
'bucket': 'b',
'data': None,
}
assert MinioDataFramePayload.cleanup_prefix_from_payload_dict(p) == 'training_datasets/m'
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_from_explicit_object_prefix():
p = {'object_key': 'x.parquet', 'object_prefix': 'my/prefix', 'data': None}
assert MinioDataFramePayload.cleanup_prefix_from_payload_dict(p) == 'my/prefix'
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
@mark.asyncio
async def test_resolve_dict_if_offloaded_noop():
d = {'success': True, 'data': {'a': [1]}}
out = await MinioDataFramePayload.resolve_dict_if_offloaded(d, None, {})
assert out is d
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
@mark.asyncio
async def test_dataframe_from_wire_list():
df = await MinioDataFramePayload.dataframe_from_wire([{'a': 1}], None, {})
assert list(df.columns) == ['a']
@pytest.mark.asyncio
async def test_retrieve_inline_dict_as_dataframe():
payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]})
minio = AsyncMock()
out = await payload.retrieve(minio, {'metadata': {}})
assert list(out.columns) == ['a']
minio.download_file.assert_not_called()
@pytest.mark.asyncio
async def test_retrieve_downloads_parquet_when_offloaded():
source = DataFrame({'a': [1, 2]})
buf = BytesIO()
source.to_parquet(buf, engine='pyarrow', index=True)
file_bytes = buf.getvalue()
payload = MinioDataFramePayload(
last_timestamp='t',
data=None,
object_key='training_datasets/m/f.parquet',
object_prefix='training_datasets/m',
)
minio = AsyncMock()
minio.download_file = AsyncMock(return_value=file_bytes)
out = await payload.retrieve(minio, {'metadata': {}})
minio.download_file.assert_awaited_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 == 'training_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet'
assert prefix == 'training_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 == 'training_datasets/my-model'
assert key.startswith('training_datasets/my-model/')
def test_estimate_size_bytes_fallback():
df = DataFrame({'a': [1, 2]})
original_to_dict = df.to_dict
df.to_dict = lambda *a, **kw: (_ for _ in ()).throw(RuntimeError('to_dict failed'))
size = MinioDataFramePayload.estimate_size_bytes(df)
df.to_dict = original_to_dict
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
@pytest.mark.asyncio
async def test_retrieve_empty_when_no_data():
payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key=None)
minio = AsyncMock()
out = await payload.retrieve(minio, {})
assert out.empty
minio.download_file.assert_not_called()
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
async def test_from_dataframe_none(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = AsyncMock()
result = await 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
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
async def test_from_dataframe_empty(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = AsyncMock()
mock_df = MagicMock()
mock_df.__bool__ = MagicMock(return_value=True)
mock_df.empty = True
result = await 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
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
async def test_from_dataframe_inline():
minio = AsyncMock()
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = await 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'
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now')
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0)
async def test_from_dataframe_offloaded(mock_now):
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
minio = AsyncMock()
minio.upload_file = AsyncMock(return_value={'minio_object_name': 'full/key.parquet'})
minio.bucket = 'test-bucket'
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = await 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_awaited_once()

View File

@@ -2,6 +2,7 @@ from datetime import UTC, datetime
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
import mlflow as mlflow_lib
import numpy as np
import pytest
from pandas import DataFrame, Timestamp
@@ -1310,10 +1311,8 @@ async def test_transform_success(mlflow_repository):
mlflow_repository.get_cached_operation.return_value, metadata['metadata']
)
assert output == {
'success': True,
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value,
}
assert output['success'] is True
assert output['content'] is mlflow_repository.detect_and_parse_datetime_index.return_value
@pytest.mark.asyncio
@@ -1355,7 +1354,7 @@ async def test_predict_success_array(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name,
data=data,
data=ANY,
operation='predict',
retention=60,
flavor='pyfunc',
@@ -1363,10 +1362,12 @@ async def test_predict_success_array(mlflow_repository):
)
assert output['success'] is True
assert output['content'] == {
'prediction': {'index_1': 2, 'index_2': 3},
'response_time': {'index_1': ANY, 'index_2': ANY},
}
content = output['content']
assert isinstance(content, DataFrame)
assert 'prediction' in content.columns
assert 'response_time' in content.columns
assert list(content.columns) == ['prediction', 'response_time']
assert content.index.tolist() == data.index.tolist()
@pytest.mark.asyncio
@@ -1383,7 +1384,7 @@ async def test_predict_success_df(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name,
data=data,
data=ANY,
operation='predict',
retention=60,
flavor='pyfunc',
@@ -1391,10 +1392,12 @@ async def test_predict_success_df(mlflow_repository):
)
assert output['success'] is True
assert output['content'] == {
'prediction': {'index_1': 2, 'index_2': 3},
'response_time': {'index_1': ANY, 'index_2': ANY},
}
content = output['content']
assert isinstance(content, DataFrame)
assert 'prediction' in content.columns
assert 'response_time' in content.columns
assert list(content.columns) == ['prediction', 'response_time']
assert content.index.tolist() == data.index.tolist()
@pytest.mark.asyncio
@@ -1409,7 +1412,7 @@ async def test_predict_error(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name,
data=data,
data=ANY,
operation='predict',
retention=60,
flavor='pyfunc',
@@ -1559,3 +1562,119 @@ def test_get_prediction_data_pyfunc(mlflow_repository):
assert 'target' in result.columns
assert 'timestamp' in result.columns
assert result.index.tolist() == [0, 1]
@patch('laborious.utils.repository.model_repository.pd.merge')
@patch('laborious.utils.repository.model_repository.isinstance')
@pytest.mark.asyncio
async def test_fit_models_skip_transform(isinstance_mock, pd_merge, mlflow_repository):
isinstance_mock.return_value = True
data_model = MagicMock(target_variable='feat_2')
prediction_model = MagicMock()
mlflow_repository.download_model = AsyncMock(
side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
)
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
return_value=MagicMock(
drop_duplicates=MagicMock(return_value=MagicMock(columns=['feat_1']))
)
)
mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame())
data = MagicMock()
output = await mlflow_repository.fit_models(
'model_name',
data,
'latest_production_id',
metadata['metadata'],
'sklearn',
True,
'pyfunc',
'feat_1',
)
data_model.fit.assert_not_called()
assert output['data_model'] == {'model': data_model, 'artifact_path': 'artifact_path'}
@patch('laborious.utils.repository.model_repository.force_memory_release')
@patch('laborious.utils.repository.model_repository.path')
@patch('laborious.utils.repository.model_repository.rmtree')
@pytest.mark.asyncio
async def test_create_new_experiment_path_not_exists(
_rmtree, path, force_memory_release, mlflow, mlflow_repository
):
model_name = 'model_name'
data = MagicMock()
prediction_data = MagicMock(spec=DataFrame)
retrain_data = {
'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
'prediction_data': prediction_data,
}
mlflow_repository.get_model_params = MagicMock(
return_value={
'transform_flavor': 'sklearn',
'predict_flavor': 'pyfunc',
'target_name': 'target_name',
}
)
mlflow_repository.get_experiment = MagicMock()
mlflow_repository.get_next_run_name = MagicMock()
mlflow_repository.log_model = AsyncMock()
path.exists.return_value = False
path.join.return_value = './tmp/artifacts/model_name'
await mlflow_repository.create_new_experiment(
model_name,
data,
retrain_data,
'latest_production_id',
metadata['metadata'],
'sklearn',
'pyfunc',
)
_rmtree.assert_not_called()
@pytest.mark.asyncio
async def test_update_production_model_by_run_id_transition_error(mlflow, mlflow_repository):
mlflow_repository.client.get_registered_model.return_value = MagicMock(
latest_versions=[
MagicMock(version='1'),
MagicMock(version='2'),
]
)
mlflow_repository.client.transition_model_version_stage.side_effect = Exception(
'transition error'
)
with pytest.raises(Exception, match='transition error'):
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
mlflow_repository.emit_metric.assert_called_with(
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
)
@pytest.mark.asyncio
async def test_predict_success_ndarray(mlflow_repository):
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
model_name = 'model'
mlflow_repository.get_cached_operation = AsyncMock(return_value=np.array([5.0, 6.0]))
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
assert output['success'] is True
content = output['content']
assert isinstance(content, DataFrame)
assert 'prediction' in content.columns
assert 'response_time' in content.columns
assert content.index.tolist() == data.index.tolist()

View File

@@ -102,6 +102,7 @@ def test_build_minio_config_with_env_vars():
'secret_key': 'test-secret',
'region_name': 'test-region',
'default_bucket': 'test-bucket',
'retention_hours': 24,
}
@@ -117,4 +118,5 @@ def test_build_minio_config_with_defaults():
'secret_key': 'minioadmin',
'region_name': 'us-east-1',
'default_bucket': 'laborious',
'retention_hours': 24,
}

View File

@@ -34,6 +34,7 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
'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',
@@ -58,12 +59,13 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
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'],
**metadata,
'model_name': input_data['model_name'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -141,6 +143,7 @@ async def test_run_none_path_flag_with_transformed_data(
'transformed_data': {'transformed': 'data'},
'timestamp': '2021-01-01',
'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0.9,
'schema': 'test_schema',
'table_name': 'test_table',
@@ -176,12 +179,13 @@ async def test_run_none_path_flag_with_transformed_data(
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'],
**metadata,
'model_name': input_data['model_name'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -189,9 +193,10 @@ async def test_run_none_path_flag_with_transformed_data(
call(
Activities.format_transformed_data,
{
**metadata,
'data': input_data['transformed_data'],
'model_id': input_data['model_id'],
**metadata,
'model_name': input_data['model_name'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -201,8 +206,9 @@ async def test_run_none_path_flag_with_transformed_data(
# Assert - start_activity_method for transformed data export
workflow_mock.start_activity_method.assert_called_once_with(
Activities.export_data_to_postgres,
Activities.export_payload_to_postgres,
{
**metadata,
'schema': input_data['schema'],
'table_name': input_data['transform_table_name'],
'data': transformed_data,
@@ -210,7 +216,6 @@ async def test_run_none_path_flag_with_transformed_data(
'column': 'timestamp',
'format': DATETIME_FORMAT_WITH_TZ,
},
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -287,6 +292,7 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction
'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',
@@ -311,11 +317,11 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction
call(
Activities.format_default_prediction,
{
**metadata,
'timestamp': input_data['timestamp'],
'model_id': input_data['model_id'],
'prediction_confidence': input_data['prediction_confidence'],
'comment': input_data['comment'],
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -389,6 +395,7 @@ async def test_run_none_path_flag_with_pi_web_api(workflow_mock, format_and_expo
'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',
@@ -415,12 +422,13 @@ async def test_run_none_path_flag_with_pi_web_api(workflow_mock, format_and_expo
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'],
**metadata,
'model_name': input_data['model_name'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -496,6 +504,7 @@ async def test_run_none_path_flag_with_pi_web_api_and_opc(
'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',
@@ -597,6 +606,7 @@ async def test_run_default_path_flag_with_pi_web_api(workflow_mock, format_and_e
'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',

View File

@@ -1,4 +1,4 @@
from unittest.mock import ANY, AsyncMock, call, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pytest import fixture, mark
@@ -27,9 +27,12 @@ metadata = {
async def test_run(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(return_value=False)
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': {'test': 'data'},
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
@@ -47,7 +50,6 @@ async def test_run(workflow_mock, prediction_process):
# Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform)
@@ -63,21 +65,7 @@ async def test_run(workflow_mock, prediction_process):
await prediction_process.run(input_data)
# Assert
assert workflow_mock.execute_local_activity_method.call_count == 7
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.get_last_timestamp,
{
**metadata,
'data': input_data['data'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
)
]
)
assert workflow_mock.execute_local_activity_method.call_count == 6
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
@@ -102,7 +90,6 @@ async def test_run(workflow_mock, prediction_process):
'data': input_data['data'],
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -201,9 +188,12 @@ async def test_run(workflow_mock, prediction_process):
async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(return_value=True)
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': {'test': 'data'},
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
@@ -219,7 +209,6 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
# Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('stop', 0.95, 'Input data with bad quality'), # input_gate
]
@@ -227,18 +216,9 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
await prediction_process.run(input_data)
# Assert
assert workflow_mock.execute_local_activity_method.call_count == 2
assert workflow_mock.execute_local_activity_method.call_count == 1
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.get_last_timestamp,
{
'data': input_data['data'],
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
),
call(
Activities.input_gate,
{
@@ -260,9 +240,12 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, True])
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': {'test': 'data'},
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
@@ -278,7 +261,6 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
# Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('repeat', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
('continue', 0.95, 'Error'), # mlflow_response_gate (transform)
@@ -288,20 +270,7 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
await prediction_process.run(input_data)
# Assert
assert workflow_mock.execute_local_activity_method.call_count == 4
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.get_last_timestamp,
{
'data': input_data['data'],
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
)
]
)
assert workflow_mock.execute_local_activity_method.call_count == 3
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
@@ -325,7 +294,6 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
'data': input_data['data'],
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata,
},
retry_policy=ANY,
@@ -357,9 +325,12 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, True])
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': {'test': 'data'},
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
@@ -375,7 +346,6 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
# Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform)
@@ -388,21 +358,8 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
await prediction_process.run(input_data)
# Assert
assert workflow_mock.execute_local_activity_method.call_count == 5
assert workflow_mock.execute_local_activity_method.call_count == 4
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.get_last_timestamp,
{
'data': input_data['data'],
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
)
]
)
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
@@ -426,7 +383,6 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
'data': input_data['data'],
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata,
},
retry_policy=ANY,
@@ -474,9 +430,12 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, False, True])
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': {'test': 'data'},
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
@@ -492,7 +451,6 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
# Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform)
@@ -507,20 +465,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
await prediction_process.run(input_data)
# Assert
assert workflow_mock.execute_local_activity_method.call_count == 7
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.get_last_timestamp,
{
'data': input_data['data'],
**metadata,
},
retry_policy=ANY,
start_to_close_timeout=ANY,
)
]
)
assert workflow_mock.execute_local_activity_method.call_count == 6
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
@@ -544,7 +489,6 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
'data': input_data['data'],
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata,
},
retry_policy=ANY,
@@ -821,3 +765,49 @@ async def test_path_flag_handler_unknown(workflow_mock, prediction_process):
assert result is False
workflow_mock.execute_activity_method.assert_not_called()
workflow_mock.execute_child_workflow.assert_not_called()
@mark.asyncio
@patch('laborious.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(return_value=False)
prediction_process.cleanup_prefixes = {'training_datasets/test'}
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = {
'metadata': metadata,
'data': data_payload,
'schema': 'test_schema',
'table_name': 'test_table',
'transform_table_name': 'test_transform_table',
'model_id': 1,
'input_filters': {'test': 'filter'},
'mlflow_transform_filters': {'test': 'filter'},
'mlflow_predict_filters': {'test': 'filter'},
'model_name': 'test_model_name',
'model_config': {'retention': '30'},
'path_priority': ['continue', 'repeat', 'stop'],
'opc_output_config': {'test': 'config'},
'pi_web_api_output_config': {'test': 'config'},
'prediction_store_policy': 'lts:1',
}
workflow_mock.execute_local_activity_method.side_effect = [
('continue', 0.95, 'ok'),
{'content': 'transformed_data', 'timestamp': '2024-01-01'},
('continue', 0.95, ''),
('continue', 0.95, ''),
{'content': 'predicted_data', 'timestamp': '2024-01-01'},
('continue', 0.95, ''),
]
await prediction_process.run(input_data)
workflow_mock.execute_activity_method.assert_any_call(
Activities.cleanup_minio_objects_expired,
{**metadata, 'prefix': 'training_datasets/test'},
retry_policy=ANY,
start_to_close_timeout=ANY,
)

View File

@@ -1,4 +1,4 @@
from unittest.mock import ANY, AsyncMock, call, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pytest import fixture, mark
@@ -39,9 +39,12 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = True
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
{'data': {'a': [1]}, 'success': True},
storage_result,
{'success': True, 'experiment': 'test_experiment'},
{
'success': True,
@@ -77,7 +80,7 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
Activities.retrain_model,
{
**metadata,
'data': {'data': {'a': [1]}, 'success': True},
'data': storage_result,
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
},
@@ -160,9 +163,12 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = False
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
{'success': False, 'message': 'No data returned from query'},
storage_result,
{'success': True, 'experiment': 'test_experiment'},
{
'success': True,
@@ -174,7 +180,10 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
]
)
await minimal_retrain.run(input_data)
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,
@@ -209,9 +218,12 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = True
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
{'data': {'a': [1]}, 'success': True},
storage_result,
{'success': False, 'experiment': 'test_experiment'},
{
'success': True,
@@ -247,7 +259,7 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
Activities.retrain_model,
{
**metadata,
'data': {'data': {'a': [1]}, 'success': True},
'data': storage_result,
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
},

View File

@@ -54,7 +54,6 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'],
'key_prefix': f"predictions/{input_data['schedule_name']}",
},
retry_policy=ANY,
start_to_close_timeout=ANY,

View File

@@ -1,124 +0,0 @@
#!/bin/bash
# Model Manager Code Validation Script
# This script runs all code quality checks before committing or deploying
set -e # Exit on any error
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
BLUE='\033[0;34m'
NC='\033[0m' # No Color
# Args
FIX_MODE=false
while [[ $# -gt 0 ]]; do
case "$1" in
--fix)
FIX_MODE=true
shift
;;
-h|--help)
echo "Usage: $0 [--fix]"
echo " --fix Apply Ruff auto-fixes (format and lint fixes)."
exit 0
;;
*)
echo -e "${RED}Unknown option: $1${NC}"
echo "Usage: $0 [--fix]"
exit 2
;;
esac
done
echo -e "${BLUE}╔════════════════════════════════════════════════════════╗${NC}"
echo -e "${BLUE}║ Model Manager - Code Validation Suite ║${NC}"
echo -e "${BLUE}╚════════════════════════════════════════════════════════╝${NC}"
echo ""
# Check if virtual environment is activated
if [[ -z "${VIRTUAL_ENV}" ]] && [[ -z "${CONDA_DEFAULT_ENV}" ]]; then
echo -e "${YELLOW}⚠️ Warning: No virtual environment detected${NC}"
echo -e "${YELLOW} Consider activating your venv/conda environment${NC}"
echo ""
fi
# Function to run a validation step
run_step() {
local step_name=$1
local step_command=$2
echo -e "${BLUE}━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━${NC}"
echo -e "${BLUE}${step_name}${NC}"
echo -e "${BLUE}━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━${NC}"
if eval "$step_command"; then
echo -e "${GREEN}${step_name} - PASSED${NC}"
echo ""
return 0
else
echo -e "${RED}${step_name} - FAILED${NC}"
echo ""
return 1
fi
}
# Track failures
FAILED_STEPS=()
# Step 1: Code Formatting Check (Ruff)
# - default: check only
# - --fix: write changes
if ! run_step "1. Code Formatting (Ruff)" "if \$FIX_MODE; then ruff format laborious/ tests/; else ruff format --check laborious/ tests/ e2e/; fi"; then
FAILED_STEPS+=("Code Formatting")
fi
# Step 2: Linting (Ruff)
# - default: check only
# - --fix: apply autofixes
if ! run_step "2. Code Linting (Ruff)" "if \$FIX_MODE; then ruff check --fix laborious/ tests/; else ruff check laborious/ tests/ e2e/; fi"; then
FAILED_STEPS+=("Linting")
fi
# Step 3: Type Checking (mypy)
if ! run_step "3. Type Checking (mypy)" "mypy laborious/"; then
FAILED_STEPS+=("Type Checking")
fi
# Step 4: Security Analysis (Bandit)
if ! run_step "4. Security Analysis (Bandit)" "bandit -c pyproject.toml -r laborious/ -ll -q"; then
FAILED_STEPS+=("Security Analysis")
fi
# Step 5: Unit Tests (pytest)
if ! run_step "5. Unit Tests (pytest)" "pytest tests/ --cov=laborious --cov-report=term-missing --cov-report=xml --cov-report=html --cov-fail-under=80 -q"; then
FAILED_STEPS+=("Unit Tests")
fi
# Summary
echo -e "${BLUE}╔════════════════════════════════════════════════════════╗${NC}"
echo -e "${BLUE}║ Validation Summary ║${NC}"
echo -e "${BLUE}╚════════════════════════════════════════════════════════╝${NC}"
echo ""
if [ ${#FAILED_STEPS[@]} -eq 0 ]; then
echo -e "${GREEN}✅ All validation checks passed!${NC}"
echo -e "${GREEN} Your code is ready for commit/deployment.${NC}"
echo ""
exit 0
else
echo -e "${RED}❌ Validation failed for the following steps:${NC}"
for step in "${FAILED_STEPS[@]}"; do
echo -e "${RED}${step}${NC}"
done
echo ""
echo -e "${YELLOW}💡 Tips:${NC}"
echo -e "${YELLOW} • Run 'ruff format laborious/ tests/' to auto-fix formatting${NC}"
echo -e "${YELLOW} • Run 'ruff check --fix laborious/ tests/' to auto-fix linting issues${NC}"
echo -e "${YELLOW} • Review mypy errors and add type hints where needed${NC}"
echo -e "${YELLOW} • Check bandit warnings for security issues${NC}"
echo -e "${YELLOW} • Fix failing tests or improve test coverage${NC}"
echo ""
exit 1
fi