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

@@ -16,7 +16,6 @@ with workflow.unsafe.imports_passed_through():
from laborious.activities.storage import Storage from laborious.activities.storage import Storage
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
""" """
Main activities orchestrator for the Laborious system. Main activities orchestrator for the Laborious system.

View File

@@ -13,17 +13,15 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController 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 sientia_do.utils.formatters import create_sample_dict
from laborious import metrics from laborious import metrics
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from laborious.utils.filters.conditional_filters import ( from laborious.utils.filters.conditional_filters import (
filter_empty_data, filter_empty_data,
filter_specific_variables_null_values, filter_specific_variables_null_values,
) )
from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
# Strongly-typed filter function signatures # Strongly-typed filter function signatures
InputFilterFunc = Callable[[DataFrame, dict[str, Any]], bool] InputFilterFunc = Callable[[DataFrame, dict[str, Any]], bool]
@@ -106,7 +104,9 @@ class Gates(MinioManager):
Raises: Raises:
Exception: If BaseActivity initialization fails 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: def close(self) -> None:
""" """
@@ -479,7 +479,7 @@ class Gates(MinioManager):
minio_repo=self.minio_repository, minio_repo=self.minio_repository,
model_name=input_data['model_name'], model_name=input_data['model_name'],
operation='transform', operation='transform',
workflow_metadata=metadata workflow_metadata=metadata,
) )
@activity.defn(name='format_prediction') @activity.defn(name='format_prediction')
@@ -601,7 +601,6 @@ class Gates(MinioManager):
self.info(f'Default prediction formatted: {data.size} rows', metadata) self.info(f'Default prediction formatted: {data.size} rows', metadata)
return data.to_dict() return data.to_dict()
@activity.defn(name='format_retrain_report') @activity.defn(name='format_retrain_report')
async def format_retrain_report(self, input_data: dict[str, Any]) -> dict: async def format_retrain_report(self, input_data: dict[str, Any]) -> dict:
""" """
@@ -672,7 +671,6 @@ class Gates(MinioManager):
return report.to_dict() return report.to_dict()
@activity.defn(name='write_metrics') @activity.defn(name='write_metrics')
async def write_metrics(self, input_data: dict[str, Any]): 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 temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager from laborious.utils.repository.minio_manager import MinioManager
@@ -8,13 +7,12 @@ with workflow.unsafe.imports_passed_through():
from typing import Any from typing import Any
import numpy as np import numpy as np
from io import BytesIO from pandas import to_datetime
from pandas import DataFrame, read_parquet, to_datetime
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController 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 ( from sientia_do.temporal.constants import (
DATETIME_FORMAT, DATETIME_FORMAT,
DATETIME_FORMAT_MS_WITH_TZ, 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 sientia_do.utils.formatters import create_sample_dict
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload 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 from laborious.utils.repository.model_repository import MLFlowRepository
@@ -72,7 +69,9 @@ class MLFlow(MinioManager):
Raises: Raises:
Exception: If MLFlowRepository initialization fails 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_host = mlflow_host
self.mlflow_port = mlflow_port self.mlflow_port = mlflow_port
self.mlflow_username = mlflow_username self.mlflow_username = mlflow_username
@@ -251,7 +250,6 @@ class MLFlow(MinioManager):
self.info('Data predicted successfully', metadata) self.info('Data predicted successfully', metadata)
if not response_data.get('success', False): if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe( return await MinioDataFramePayload.from_dataframe(
dataframe=None, dataframe=None,
@@ -270,10 +268,9 @@ class MLFlow(MinioManager):
workflow_metadata=metadata, workflow_metadata=metadata,
status={ status={
'success': True, 'success': True,
} },
) )
@activity.defn(name='retrain_model') @activity.defn(name='retrain_model')
async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]: async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
""" """
@@ -313,22 +310,10 @@ class MLFlow(MinioManager):
metadata = input_data['metadata'] metadata = input_data['metadata']
try: try:
if 'data' in input_data: # Payload-based retrain input (inline dict or MinIO offloaded).
# New path: payload-based retrain input (inline or MinIO offloaded). payload: MinioDataFramePayload = input_data['data']
data = await MinioDataFramePayload.dataframe_from_wire( data = await payload.retrieve(self.minio_repository, metadata)
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))
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
await self.send_notification_async( await self.send_notification_async(

View File

@@ -1,15 +1,14 @@
import json import json
from temporalio import activity, workflow from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
# Extend the Temporal Postgres activities for convenient query -> MinIO export # Extend the Temporal Postgres activities for convenient query -> MinIO export
import pickle
import traceback import traceback
from datetime import timedelta from datetime import timedelta
from io import BytesIO from io import BytesIO
from os import getenv
from typing import Any from typing import Any
import pandas as pd import pandas as pd
@@ -17,11 +16,11 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController 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.activities.postgres import Postgres
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload 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'}) _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, 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') @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. Run the custom SQL load, then return a MinIO-aware dataframe wire dict.
@@ -92,7 +95,9 @@ class Storage(Postgres, MinioManager):
input_data, input_data,
) )
if not rows: 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 dataframe = None
else: else:
dataframe = pd.DataFrame(rows) dataframe = pd.DataFrame(rows)
@@ -201,8 +206,6 @@ class Storage(Postgres, MinioManager):
return report return report
@activity.defn(name='query_to_minio') @activity.defn(name='query_to_minio')
async def query_to_minio(self, input_data: dict[str, Any]) -> dict[str, Any]: 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 pickle
import re import re
from dataclasses import dataclass, field from collections.abc import Hashable
from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from io import BytesIO from io import BytesIO
from os import getenv from os import getenv
from typing import Any, Hashable, Literal from typing import Any, Literal
from pandas import DataFrame, read_parquet 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.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). # Keys that are part of the serialized wire format (not arbitrary metadata).
_SERIALIZED_FIELD_KEYS = frozenset({'data', 'bucket', 'object_key', 'object_prefix', 'uri'}) _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$' 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. # Relative prefix used for storing offloaded training datasets in MinIO.
# It is also the root directory for retention cleanup listing. # It is also the root directory for retention cleanup listing.
@@ -79,7 +81,6 @@ class MinioDataFramePayload:
object_prefix: str | None = None object_prefix: str | None = None
uri: str | None = None uri: str | None = None
@staticmethod @staticmethod
def estimate_size_bytes(df: DataFrame) -> int: def estimate_size_bytes(df: DataFrame) -> int:
""" """
@@ -117,21 +118,6 @@ class MinioDataFramePayload:
except ValueError: except ValueError:
return None 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 @staticmethod
def cleanup_prefix(self) -> str | None: def cleanup_prefix(self) -> str | None:
""" """
@@ -141,6 +127,12 @@ class MinioDataFramePayload:
return self.object_prefix return self.object_prefix
return None 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 @classmethod
async def from_dataframe( async def from_dataframe(
cls, cls,
@@ -172,7 +164,9 @@ class MinioDataFramePayload:
""" """
if not dataframe or dataframe.empty: 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()) last_timestamp = max(dataframe['timestamp'].values.tolist())
@@ -207,7 +201,9 @@ class MinioDataFramePayload:
last_timestamp=last_timestamp, 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. 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: if self.data is not None:
return DataFrame(self.data) return DataFrame(self.data)
if self.data is None and self.object_key is None: if not self.has_data():
return DataFrame() return DataFrame()
file_bytes = await minio_repo.download_file( 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)) df = read_parquet(BytesIO(file_bytes))
return df 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.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController 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.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository import MinioRepository
class MinioManager(SientiaMonitoring): class MinioManager(SientiaMonitoring):
minio_repository: MinioRepository | None = None 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: if self.minio_repository is None:
self.minio_repository = minio_repository self.minio_repository = minio_repository
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)

View File

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

View File

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

View File

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

View File

@@ -2,7 +2,6 @@ from temporalio import workflow
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
from datetime import timedelta from datetime import timedelta
from collections.abc import Callable
from typing import Any from typing import Any
from sientia_do.temporal.policies import retry_policy from sientia_do.temporal.policies import retry_policy
@@ -120,9 +119,8 @@ class PredictionProcess:
model_id: Any, model_id: Any,
model_name: str, model_name: str,
model_config: dict[str, Any], model_config: dict[str, Any],
save_transform: bool save_transform: bool,
) -> None: ) -> None:
last_timestamp = data.last_timestamp last_timestamp = data.last_timestamp
# Apply input data quality gates # Apply input data quality gates
@@ -149,12 +147,7 @@ class PredictionProcess:
# Request MLFlow model transformation # Request MLFlow model transformation
response_data = await workflow.execute_local_activity_method( response_data = await workflow.execute_local_activity_method(
Activities.request_transform, 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, retry_policy=retry_policy,
start_to_close_timeout=timedelta(minutes=5), 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. 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. imports are executed.
""" """
import sys
from unittest.mock import MagicMock
# Mock sientia module # Mock sientia module
sientia_mock = MagicMock() sientia_mock = MagicMock()
sientia_mock.ModelAnalysis = 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.Gates.__init__')
@patch('laborious.activities.activities.ModelMetrics.__init__') @patch('laborious.activities.activities.ModelMetrics.__init__')
@patch('laborious.activities.activities.API.__init__') @patch('laborious.activities.activities.API.__init__')
@patch('laborious.activities.activities.MinioRepository')
@patch('laborious.activities.activities.MetricsController') @patch('laborious.activities.activities.MetricsController')
def test___init__( def test___init__(
mock_metrics_controller, mock_metrics_controller,
mock_minio_repository,
mock_api_init, mock_api_init,
mock_model_metrics_init, mock_model_metrics_init,
mock_gates_init, mock_gates_init,
@@ -43,6 +45,7 @@ def test___init__(
'secret_key': 'minio123', 'secret_key': 'minio123',
'region_name': 'us-east-1', 'region_name': 'us-east-1',
'default_bucket': 'test', 'default_bucket': 'test',
'retention_hours': 24,
} }
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'} mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
@@ -89,7 +92,8 @@ def test___init__(
dbname=postgres_config['dbname'], dbname=postgres_config['dbname'],
min_connections=postgres_config['min_connections'], min_connections=postgres_config['min_connections'],
max_connections=postgres_config['max_connections'], max_connections=postgres_config['max_connections'],
minio_config=minio_config, retention_hours=minio_config['retention_hours'],
minio_repository=mock_minio_repository.return_value,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value, metrics_controller=mock_metrics_controller.return_value,
@@ -101,7 +105,7 @@ def test___init__(
mlflow_port=mlflow_config['port'], mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'], mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'], mlflow_password=mlflow_config['password'],
minio_config=minio_config, minio_repository=mock_minio_repository.return_value,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value, metrics_controller=mock_metrics_controller.return_value,
@@ -117,6 +121,7 @@ def test___init__(
mock_gates_init.assert_called_once_with( mock_gates_init.assert_called_once_with(
ANY, ANY,
minio_repository=mock_minio_repository.return_value,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value, metrics_controller=mock_metrics_controller.return_value,
@@ -139,6 +144,16 @@ def test___init__(
metrics_controller=mock_metrics_controller.return_value, 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 @mark.asyncio
@patch('laborious.activities.activities.Storage') @patch('laborious.activities.activities.Storage')
@@ -147,7 +162,9 @@ def test___init__(
@patch('laborious.activities.activities.Gates') @patch('laborious.activities.activities.Gates')
@patch('laborious.activities.activities.ModelMetrics') @patch('laborious.activities.activities.ModelMetrics')
@patch('laborious.activities.activities.API') @patch('laborious.activities.activities.API')
@patch('laborious.activities.activities.MinioRepository')
async def test_shutdown( async def test_shutdown(
_mock_minio_repository,
mock_api_init, mock_api_init,
mock_model_metrics_init, mock_model_metrics_init,
mock_gates_init, mock_gates_init,
@@ -172,6 +189,7 @@ async def test_shutdown(
'secret_key': 'minio123', 'secret_key': 'minio123',
'region_name': 'us-east-1', 'region_name': 'us-east-1',
'default_bucket': 'test', 'default_bucket': 'test',
'retention_hours': 24,
} }
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'} 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__(): def test__init__():
api = API( api = API(
base_url='https://test-pi-server.com', base_url='https://test-pi-server.com',

View File

@@ -1,11 +1,29 @@
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pandas import DataFrame
from pytest import fixture, mark from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from laborious.activities.gates import Gates 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 @fixture
def gates_activity(): def gates_activity():
gates = Gates( gates = Gates(
@@ -40,7 +58,7 @@ async def test_input_gate_invalid_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}}, 'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'value': [1, 2, 3]}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -65,7 +83,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
input_data = { input_data = {
**metadata, **metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': []}, 'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -90,7 +108,7 @@ async def test_input_gate_no_filters(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {}, 'filters': {},
'data': {'value': [1, 2, 3]}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'], 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
} }
@@ -108,7 +126,7 @@ async def test_input_gate_with_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': []}, 'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -126,7 +144,7 @@ async def test_input_gate_with_filter_not_caught(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}}, 'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': {'value': [1, 2, 3]}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -144,7 +162,10 @@ async def test_mlflow_response_gate_invalid_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}}, 'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}}, 'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -169,7 +190,10 @@ async def test_mlflow_response_gate_filter_exception(
input_data = { input_data = {
**metadata, **metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}}, 'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}}, 'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -195,7 +219,10 @@ async def test_mlflow_response_gate_no_filters(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {}, 'filters': {},
'data': {'content': {'message': 'success'}}, 'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
),
'type': 'test', 'type': 'test',
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'], 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
} }
@@ -214,10 +241,10 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}}, 'filters': {'API_ERROR': {'policy': 'STOP'}},
'data': { 'data': _minio_payload(
'success': False, {'content': {'message': 'API error occurred', 'traceback': 'error trace'}},
'content': {'message': 'API error occurred', 'traceback': 'error trace'}, status={'success': False, 'message': 'API error occurred'},
}, ),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -237,10 +264,10 @@ async def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}}, 'filters': {'API_ERROR': {'policy': 'STOP'}},
'data': { 'data': _minio_payload(
'success': True, {'content': {'message': 'success'}},
'content': {'message': 'success'}, status={'success': True},
}, ),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -259,10 +286,7 @@ async def test_mlflow_content_gate_invalid_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}}, 'filters': {'INVALID_FILTER': {'POLICY': 'STOP'}},
'data': { 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'success': True,
'content': {'message': 'success'},
},
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -287,10 +311,7 @@ async def test_mlflow_content_gate_filter_exception(
input_data = { input_data = {
**metadata, **metadata,
'filters': {'API_ERROR': {'POLICY': 'STOP'}}, 'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': { 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'success': False,
'content': {'message': 'API error occurred', 'traceback': 'error trace'},
},
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -317,7 +338,7 @@ async def test_mlflow_content_gate_no_filters(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {}, 'filters': {},
'data': {'value': [1, 2, 3]}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test', 'type': 'test',
'path_priority': ['CONTINUE', 'STOP', 'REPEAT'], 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'],
} }
@@ -336,7 +357,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}}, 'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}},
'data': {'value': [None, None, None]}, 'data': _minio_payload(DataFrame({'value': [None, None, None]})),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
} }
@@ -356,7 +377,7 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
input_data = { input_data = {
**metadata, **metadata,
'filters': {'API_ERROR': {'POLICY': 'STOP'}}, 'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': {'content': {'message': 'success'}}, 'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test', 'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'], '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() 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): def test_get_prediction_store_policy_invalid_policy(gates_activity):
# Arrange # Arrange
prediction_store_policy = 'INVALID_POLICY' prediction_store_policy = 'INVALID_POLICY'
@@ -430,10 +467,14 @@ async def test_format_prediction_no_timestamp(gates_activity):
# Arrange # Arrange
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'prediction': {'2023-05-26 11:12:27': 1}, 'prediction': {'2023-05-26 11:12:27': 1},
'response_time': {'2023-05-26 11:12:27': 0.1}, 'response_time': {'2023-05-26 11:12:27': 0.1},
}, }
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'prediction_confidence': 0.9, 'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:1', 'prediction_store_policy': 'lts:1',
@@ -457,7 +498,9 @@ async def test_format_prediction_with_timestamp_erl(gates_activity):
# Arrange # Arrange
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'prediction': { 'prediction': {
'2023-05-26 11:12:27': 1, '2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2, '2023-05-26 11:12:28': 2,
@@ -468,7 +511,9 @@ async def test_format_prediction_with_timestamp_erl(gates_activity):
'2023-05-26 11:12:28': 0.2, '2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3, '2023-05-26 11:12:29': 0.3,
}, },
}, }
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'prediction_confidence': 0.9, 'prediction_confidence': 0.9,
'prediction_store_policy': 'erl:2', 'prediction_store_policy': 'erl:2',
@@ -492,7 +537,9 @@ async def test_format_prediction_with_timestamp_lts(gates_activity):
# Arrange # Arrange
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'prediction': { 'prediction': {
'2023-05-26 11:12:27': 1, '2023-05-26 11:12:27': 1,
'2023-05-26 11:12:28': 2, '2023-05-26 11:12:28': 2,
@@ -503,7 +550,9 @@ async def test_format_prediction_with_timestamp_lts(gates_activity):
'2023-05-26 11:12:28': 0.2, '2023-05-26 11:12:28': 0.2,
'2023-05-26 11:12:29': 0.3, '2023-05-26 11:12:29': 0.3,
}, },
}, }
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'prediction_confidence': 0.9, 'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:2', 'prediction_store_policy': 'lts:2',
@@ -527,11 +576,19 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
# Arrange # Arrange
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'prediction': [1, 2, 3], 'prediction': [1, 2, 3],
'response_time': [0.1, 0.2, 0.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'], 'timestamp': [
}, '2023-05-26 11:12:27',
'2023-05-26 11:12:28',
'2023-05-26 11:12:29',
],
}
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'prediction_confidence': 0.9, 'prediction_confidence': 0.9,
'prediction_store_policy': 'lts:2', 'prediction_store_policy': 'lts:2',
@@ -547,34 +604,51 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
@mark.asyncio @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 # Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'var1': {'2023-05-26 11:12:27': 1.0}, 'var1': {'2023-05-26 11:12:27': 1.0},
'var2': {'2023-05-26 11:12:27': 2.0}, 'var2': {'2023-05-26 11:12:27': 2.0},
}, }
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'model_name': 'test_model',
} }
# Act # Act
result = await gates_activity.format_transformed_data(input_data) result = await gates_activity.format_transformed_data(input_data)
# Assert # Assert
assert result['timestamp'] == {0: '2023-05-26 11:12:27', 1: '2023-05-26 11:12:27'} assert result is payload_result
assert result['variable'] == {0: 'var1', 1: 'var2'} mock_from_dataframe.assert_called_once()
assert result['value'] == {0: 1.0, 1: 2.0} kwargs = mock_from_dataframe.call_args.kwargs
assert result['model_id'] == {0: 'test_model', 1: 'test_model'} 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() gates_activity.info.assert_called()
@mark.asyncio @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 # Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': _minio_payload(
DataFrame(
{
'var1': { 'var1': {
'2023-05-26 11:12:27': 1.0, '2023-05-26 11:12:27': 1.0,
'2023-05-26 11:12:28': 2.0, '2023-05-26 11:12:28': 2.0,
@@ -583,40 +657,53 @@ async def test_format_transformed_data_multiple_rows(gates_activity):
'2023-05-26 11:12:27': 3.0, '2023-05-26 11:12:27': 3.0,
'2023-05-26 11:12:28': 4.0, '2023-05-26 11:12:28': 4.0,
}, },
}, }
)
),
'model_id': 'test_model', 'model_id': 'test_model',
'model_name': 'test_model',
} }
# Act # Act
result = await gates_activity.format_transformed_data(input_data) result = await gates_activity.format_transformed_data(input_data)
# Assert # Assert
assert len(result['timestamp']) == 4 assert result is payload_result
assert len(result['variable']) == 4 mock_from_dataframe.assert_called_once()
assert len(result['value']) == 4 kwargs = mock_from_dataframe.call_args.kwargs
assert len(result['model_id']) == 4 assert kwargs['model_name'] == 'test_model'
assert all(v == 'test_model' for v in result['model_id'].values()) assert kwargs['operation'] == 'transform'
assert set(result['variable'].values()) == {'var1', 'var2'} assert kwargs['workflow_metadata'] == metadata['metadata']
assert kwargs['minio_repo'] is gates_activity.minio_repository
assert 'dataframe' in kwargs
gates_activity.info.assert_called() gates_activity.info.assert_called()
@mark.asyncio @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 # Arrange
payload_result = MagicMock()
mock_from_dataframe.return_value = payload_result
input_data = { input_data = {
**metadata, **metadata,
'data': {}, 'data': _minio_payload(DataFrame()),
'model_id': 'test_model', 'model_id': 'test_model',
'model_name': 'test_model',
} }
# Act # Act
result = await gates_activity.format_transformed_data(input_data) result = await gates_activity.format_transformed_data(input_data)
# Assert # Assert
assert result['timestamp'] == {} assert result is payload_result
assert result['variable'] == {} mock_from_dataframe.assert_called_once()
assert result['value'] == {} kwargs = mock_from_dataframe.call_args.kwargs
assert result['model_id'] == {} 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() gates_activity.info.assert_called()
@@ -711,31 +798,6 @@ async def test_format_retrain_report_failure(gates_activity):
gates_activity.debug.assert_called() 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 @mark.asyncio
@patch('laborious.activities.gates.metrics') @patch('laborious.activities.gates.metrics')
async def test_write_metrics(mock_metrics, gates_activity): 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.MLFlowRepository')
@patch('laborious.activities.mlflow.MinioRepository') @patch('laborious.activities.mlflow.MinioRepository')
def test___init__(mock_minio_repository, mock_mlflow_repository): 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 = MLFlow(
mlflow_host='http://localhost', mlflow_host='http://localhost',
mlflow_port=5000, mlflow_port=5000,
mlflow_username='admin', mlflow_username='admin',
mlflow_password='admin', mlflow_password='admin',
minio_config={ minio_repository=minio_repo,
'endpoint_url': 'http://localhost:9000', logger=logger,
'access_key': 'minio', notification_handler=notification_handler,
'secret_key': 'minio123', metrics_controller=metrics_controller,
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
assert mlflow.mlflow_host == 'http://localhost' 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.MLFlowRepository')
@patch('laborious.activities.mlflow.MinioRepository') @patch('laborious.activities.mlflow.MinioRepository')
def mlflow(mock_minio_repository, mock_mlflow_repository): 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 = MLFlow(
mlflow_host='http://localhost:5000', mlflow_host='http://localhost:5000',
mlflow_port=5000, mlflow_port=5000,
mlflow_username='admin', mlflow_username='admin',
mlflow_password='admin', mlflow_password='admin',
minio_config={ minio_repository=minio_repo,
'endpoint_url': 'http://localhost:9000', logger=logger,
'access_key': 'minio', notification_handler=notification_handler,
'secret_key': 'minio123', metrics_controller=metrics_controller,
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
mlflow.model_monitoring_repository = AsyncMock() mlflow.model_monitoring_repository = AsyncMock()
@@ -96,165 +108,161 @@ metadata = {
@mark.asyncio @mark.asyncio
@patch( @patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire', 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock, new_callable=AsyncMock,
) )
@patch('laborious.activities.mlflow.max') async def test_request_transform_success(mock_from_dataframe, mlflow):
async def test_request_transform_success(mock_max, mock_dataframe_from_wire, mlflow):
mock_max.return_value = '2024-01-02'
data_mock = MagicMock() data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock payload = AsyncMock()
# Mock input data payload.retrieve = AsyncMock(return_value=data_mock)
input_data = { input_data = {
**metadata, **metadata,
'data': [ 'data': payload,
{
'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',
},
],
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': {}, 'model_config': {},
} }
# Mock the transform response transform_response = {'success': True, 'content': MagicMock()}
expected_response = {'prediction': [0.5, 0.6], 'timestamp': ['2024-01-01', '2024-01-02']} mlflow.model_monitoring_repository.transform.return_value = transform_response
mlflow.model_monitoring_repository.transform.return_value = expected_response
data_mock.sort_values.return_value = data_mock data_mock.sort_values.return_value = data_mock
data_mock.drop_duplicates.return_value = data_mock data_mock.drop_duplicates.return_value = data_mock
data_mock.pivot.return_value = data_mock data_mock.pivot.return_value = data_mock
# Call the method
response_data = await mlflow.request_transform(input_data) 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( mlflow.model_monitoring_repository.transform.assert_called_once_with(
'test_model', data_mock, {}, metadata['metadata'] 'test_model', data_mock, {}, metadata['metadata']
) )
mock_from_dataframe.assert_called_once()
assert response_data == mock_from_dataframe.return_value
@mark.asyncio @mark.asyncio
@patch( @patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire', 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock, new_callable=AsyncMock,
) )
@patch('laborious.activities.mlflow.to_datetime') async def test_request_transform_failure(mock_from_dataframe, mlflow):
@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'
data_mock = MagicMock() data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock payload = AsyncMock()
# Mock input data payload.retrieve = AsyncMock(return_value=data_mock)
input_data = { input_data = {
**metadata, **metadata,
'data': { 'data': payload,
'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},
},
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': {}, 'model_config': {},
} }
# Mock the predict response transform_response = {'success': False, 'message': 'Transform failed'}
expected_response = {'prediction': [0.5, 0.6]} mlflow.model_monitoring_repository.transform.return_value = transform_response
mlflow.model_monitoring_repository.predict.return_value = expected_response
# Call the method data_mock.sort_values.return_value = data_mock
response_data = await mlflow.request_predict(input_data) 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) response_data = await mlflow.request_transform(input_data)
data_mock.__setitem__.assert_any_call(
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value mock_from_dataframe.assert_called_once_with(
) dataframe=None,
data_mock.__setitem__.assert_any_call( minio_repo=mlflow.minio_repository,
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value model_name='test_model',
) operation='transform',
status=transform_response,
mock_to_datetime.assert_called_once_with( workflow_metadata=metadata['metadata'],
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']
) )
assert response_data == mock_from_dataframe.return_value
@mark.asyncio @mark.asyncio
@patch('laborious.activities.mlflow.read_parquet') @patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime') @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 = { mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': True, 'success': True,
'experiment': 'test_experiment', 'experiment': 'test_experiment',
'message': 'Model retrained successfully.', 'message': 'Model retrained successfully.',
} }
mlflow.minio_repository.download_file.return_value = b'parquet-bytes' raw_data = MagicMock(columns=['variable', 'timestamp', 'value'])
mock_read_parquet.return_value = MagicMock() payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data)
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
'object_key': 'test_object_key', 'data': payload,
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': { 'model_config': {
'target': 'target', '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 timestamp = raw_data.__getitem__.return_value.max.return_value
raw_data.sort_values.assert_not_called() 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 @mark.asyncio
@patch('laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire', new_callable=AsyncMock)
@patch('laborious.activities.mlflow.to_datetime') @patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_with_payload_data( async def test_retrain_model_success_with_payload_data(mock_to_datetime, mlflow):
mock_to_datetime, mock_dataframe_from_wire, mlflow
):
mlflow.model_monitoring_repository.retrain_model.return_value = { mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': True, 'success': True,
'experiment': 'test_experiment', 'experiment': 'test_experiment',
'message': 'Model retrained successfully.', '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( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
'data': {'data': {'a': [1]}}, 'data': payload,
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': { 'model_config': {
'target': 'target', 'target': 'target',
@@ -347,24 +353,22 @@ async def test_retrain_model_success_with_payload_data(
@mark.asyncio @mark.asyncio
@patch('laborious.activities.mlflow.read_parquet')
@patch('laborious.activities.mlflow.to_datetime') @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 = { mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': False, 'success': False,
'traceback': 'test_traceback', 'traceback': 'test_traceback',
'message': 'Model retrained failed.', 'message': 'Model retrained failed.',
} }
mlflow.minio_repository.download_file.return_value = b'parquet-bytes' raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
mock_read_parquet.return_value = MagicMock( payload = AsyncMock()
columns=['variable', 'timestamp', 'value', 'created_at'] payload.retrieve = AsyncMock(return_value=raw_data)
)
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
'object_key': 'test_object_key', 'data': payload,
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': { 'model_config': {
'target': 'target', '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 timestamp = raw_data.__getitem__.return_value.max.return_value
raw_data.sort_values.assert_called_once_with('created_at', ascending=False) 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 @mark.asyncio
async def test_retrain_model_data_error(mlflow): 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( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
'object_key': 'test_object_key',
'model_name': 'test_model', 'model_name': 'test_model',
'model_config': { 'model_config': {
'target': 'target', 'target': 'target',
@@ -458,7 +455,7 @@ async def test_retrain_model_data_error(mlflow):
assert response == { assert response == {
'success': False, 'success': False,
'message': 'Error loading retrain data: Error loading retrain data', 'message': "Error loading retrain data: 'data'",
'traceback': ANY, 'traceback': ANY,
'timestamp': ANY, 'timestamp': ANY,
} }

View File

@@ -743,6 +743,101 @@ async def test_get_drift_metrics_univariate_error(
raise AssertionError('Expected Exception') 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 @mark.asyncio
async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity): async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
# Arrange # 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( model_metrics_activity.info.assert_called_once_with(
"Calculating simple metrics for model test_model_id: ['rmse', 'mae']", metadata['metadata'] "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', dbname='postgres',
min_connections=1, min_connections=1,
max_connections=10, max_connections=10,
minio_config={ retention_hours=24,
'endpoint_url': 'localhost:9000', minio_repository=mock_minio_repository.return_value,
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(), metrics_controller=AsyncMock(),
@@ -47,6 +42,7 @@ def test___init___not_hasattr(mock_minio_repository):
logger = MagicMock() logger = MagicMock()
notification_handler = MagicMock() notification_handler = MagicMock()
metrics_controller = AsyncMock() metrics_controller = AsyncMock()
minio_repo = mock_minio_repository.return_value
storage = Storage( storage = Storage(
host='localhost', host='localhost',
port=5432, port=5432,
@@ -55,28 +51,16 @@ def test___init___not_hasattr(mock_minio_repository):
dbname='postgres', dbname='postgres',
min_connections=1, min_connections=1,
max_connections=10, max_connections=10,
minio_config={ retention_hours=24,
'endpoint_url': 'localhost:9000', minio_repository=minio_repo,
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
) )
assert isinstance(storage, Postgres) assert isinstance(storage, Postgres)
mock_minio_repository.assert_called_once_with( assert storage.minio_repository is minio_repo
endpoint='localhost:9000', mock_minio_repository.assert_not_called()
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
@patch('laborious.activities.storage.MinioRepository') @patch('laborious.activities.storage.MinioRepository')
@@ -93,27 +77,15 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
dbname='postgres', dbname='postgres',
min_connections=1, min_connections=1,
max_connections=10, max_connections=10,
minio_config={ retention_hours=24,
'endpoint_url': 'localhost:9000', minio_repository=None,
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
) )
mock_minio_repository.assert_called_once_with( assert storage.minio_repository is None
endpoint='localhost:9000', mock_minio_repository.assert_not_called()
access_key='minio',
secret_key='minio123',
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
bucket='test',
)
@patch('laborious.activities.storage.MinioRepository') @patch('laborious.activities.storage.MinioRepository')
@@ -126,13 +98,8 @@ def test___init___done_repository(mock_minio_repository, storage):
dbname='postgres', dbname='postgres',
min_connections=1, min_connections=1,
max_connections=10, max_connections=10,
minio_config={ retention_hours=24,
'endpoint_url': 'localhost:9000', minio_repository=mock_minio_repository.return_value,
'access_key': 'minio',
'secret_key': 'minio123',
'region_name': 'us-east-1',
'default_bucket': 'test',
},
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(), metrics_controller=AsyncMock(),
@@ -231,55 +198,58 @@ def test___del__(storage):
storage.close.assert_called_once() storage.close.assert_called_once()
def test_estimate_payload_size_bytes(storage):
assert storage._estimate_payload_size_bytes({'x': 1}) > 0
@mark.asyncio @mark.asyncio
async def test_load_query_with_minio_offload_no_rows(storage): async def test_load_query_with_minio_offload_no_rows(storage):
storage.load_custom_query = AsyncMock(return_value=None) storage.load_custom_query = AsyncMock(return_value=None)
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( result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'} {**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
) )
assert result['success'] is False assert result == storage_result
mock_from_dataframe.assert_awaited_once()
@mark.asyncio @mark.asyncio
async def test_load_query_with_minio_offload_inline(storage): async def test_load_query_with_minio_offload_inline(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}]) storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None}
with patch(
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
return_value=storage_result,
) as mock_from_dataframe:
result = await storage.load_query_with_minio_offload( result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'my-model', 'key_prefix': 'predictions/s'} {
**metadata,
'query': 'SELECT 1',
'model_name': 'my-model',
'key_prefix': 'predictions/s',
}
) )
assert result.get('success') is True assert result == storage_result
assert 'data' in result mock_from_dataframe.assert_awaited_once()
assert result.get('object_key') is None
@mark.asyncio @mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.MinioDataFramePayload.estimate_size_bytes') async def test_load_query_with_minio_offload_minio(storage):
async def test_load_query_with_minio_offload_minio(mock_estimate, storage):
mock_estimate.return_value = 10**9
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}]) storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage.minio_repository.upload_file = AsyncMock( storage_result = {'success': True, 'data': None, 'object_key': 'object-key'}
return_value={ with patch(
'minio_object_name': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-15_12-30-45.parquet' 'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
} new_callable=AsyncMock,
) return_value=storage_result,
storage.minio_repository.bucket = 'test' ) as mock_from_dataframe:
fixed = datetime.datetime(2024, 1, 15, 12, 30, 45)
with patch('laborious.utils.models.minio_dataframe_payload.now', return_value=fixed):
result = await storage.load_query_with_minio_offload( result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'} {**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
) )
assert result.get('success') is True assert result == storage_result
assert result.get('data') is None mock_from_dataframe.assert_awaited_once()
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()
@mark.asyncio @mark.asyncio
@@ -297,12 +267,117 @@ async def test_cleanup_minio_objects_expired(mock_now, storage):
storage.send_notification_async = AsyncMock() storage.send_notification_async = AsyncMock()
result = await storage.cleanup_minio_objects_expired( 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['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( 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', object_name='sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
metadata=metadata['metadata'], 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 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(): 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 assert MinioDataFramePayload.parse_object_timestamp('bad.parquet') is None
def test_is_offloaded_dict_true_false(): def test_estimate_size_bytes_returns_positive_for_nonempty_frame():
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': None}) is True df = DataFrame({'a': [1, 2]})
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': {}}) is False size = MinioDataFramePayload.estimate_size_bytes(df)
assert MinioDataFramePayload.is_offloaded_dict({'data': {}}) is False assert isinstance(size, int)
assert size > 0
def test_cleanup_prefix_from_payload_dict(): def test_cleanup_prefix_when_offloaded_returns_object_prefix():
p = { payload = MinioDataFramePayload(
'object_key': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-01_00-00-00.parquet', last_timestamp='t',
'bucket': 'b', data=None,
'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_from_payload_dict(p) == 'training_datasets/m' )
assert MinioDataFramePayload.cleanup_prefix(payload) == 'training_datasets/m'
def test_cleanup_prefix_from_explicit_object_prefix(): def test_cleanup_prefix_when_inline_returns_none():
p = {'object_key': 'x.parquet', 'object_prefix': 'my/prefix', 'data': None} payload = MinioDataFramePayload(last_timestamp='t', data={'x': [1]}, object_key=None)
assert MinioDataFramePayload.cleanup_prefix_from_payload_dict(p) == 'my/prefix' assert MinioDataFramePayload.cleanup_prefix(payload) is None
@mark.asyncio def test_has_data_true_when_object_key_set():
async def test_resolve_dict_if_offloaded_noop(): payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key='k')
d = {'success': True, 'data': {'a': [1]}} assert payload.has_data() is True
out = await MinioDataFramePayload.resolve_dict_if_offloaded(d, None, {})
assert out is d
@mark.asyncio @pytest.mark.asyncio
async def test_dataframe_from_wire_list(): async def test_retrieve_inline_dict_as_dataframe():
df = await MinioDataFramePayload.dataframe_from_wire([{'a': 1}], None, {}) payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]})
assert list(df.columns) == ['a'] 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 from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
import mlflow as mlflow_lib import mlflow as mlflow_lib
import numpy as np
import pytest import pytest
from pandas import DataFrame, Timestamp 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'] mlflow_repository.get_cached_operation.return_value, metadata['metadata']
) )
assert output == { assert output['success'] is True
'success': True, assert output['content'] is mlflow_repository.detect_and_parse_datetime_index.return_value
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value,
}
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -1355,7 +1354,7 @@ async def test_predict_success_array(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with( mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name, model_name=model_name,
data=data, data=ANY,
operation='predict', operation='predict',
retention=60, retention=60,
flavor='pyfunc', flavor='pyfunc',
@@ -1363,10 +1362,12 @@ async def test_predict_success_array(mlflow_repository):
) )
assert output['success'] is True assert output['success'] is True
assert output['content'] == { content = output['content']
'prediction': {'index_1': 2, 'index_2': 3}, assert isinstance(content, DataFrame)
'response_time': {'index_1': ANY, 'index_2': ANY}, 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 @pytest.mark.asyncio
@@ -1383,7 +1384,7 @@ async def test_predict_success_df(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with( mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name, model_name=model_name,
data=data, data=ANY,
operation='predict', operation='predict',
retention=60, retention=60,
flavor='pyfunc', flavor='pyfunc',
@@ -1391,10 +1392,12 @@ async def test_predict_success_df(mlflow_repository):
) )
assert output['success'] is True assert output['success'] is True
assert output['content'] == { content = output['content']
'prediction': {'index_1': 2, 'index_2': 3}, assert isinstance(content, DataFrame)
'response_time': {'index_1': ANY, 'index_2': ANY}, 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 @pytest.mark.asyncio
@@ -1409,7 +1412,7 @@ async def test_predict_error(mlflow_repository):
mlflow_repository.get_cached_operation.assert_called_once_with( mlflow_repository.get_cached_operation.assert_called_once_with(
model_name=model_name, model_name=model_name,
data=data, data=ANY,
operation='predict', operation='predict',
retention=60, retention=60,
flavor='pyfunc', flavor='pyfunc',
@@ -1559,3 +1562,119 @@ def test_get_prediction_data_pyfunc(mlflow_repository):
assert 'target' in result.columns assert 'target' in result.columns
assert 'timestamp' in result.columns assert 'timestamp' in result.columns
assert result.index.tolist() == [0, 1] 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', 'secret_key': 'test-secret',
'region_name': 'test-region', 'region_name': 'test-region',
'default_bucket': 'test-bucket', 'default_bucket': 'test-bucket',
'retention_hours': 24,
} }
@@ -117,4 +118,5 @@ def test_build_minio_config_with_defaults():
'secret_key': 'minioadmin', 'secret_key': 'minioadmin',
'region_name': 'us-east-1', 'region_name': 'us-east-1',
'default_bucket': 'laborious', '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'}, 'data': {'test': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0, 'prediction_confidence': 0,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
@@ -58,12 +59,13 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
call( call(
Activities.format_prediction, Activities.format_prediction,
{ {
**metadata,
'data': input_data['data'], 'data': input_data['data'],
'timestamp': input_data['timestamp'], 'timestamp': input_data['timestamp'],
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'prediction_confidence': input_data['prediction_confidence'], 'prediction_confidence': input_data['prediction_confidence'],
'prediction_store_policy': input_data['prediction_store_policy'], 'prediction_store_policy': input_data['prediction_store_policy'],
**metadata, 'model_name': input_data['model_name'],
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -141,6 +143,7 @@ async def test_run_none_path_flag_with_transformed_data(
'transformed_data': {'transformed': 'data'}, 'transformed_data': {'transformed': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0.9, 'prediction_confidence': 0.9,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
@@ -176,12 +179,13 @@ async def test_run_none_path_flag_with_transformed_data(
call( call(
Activities.format_prediction, Activities.format_prediction,
{ {
**metadata,
'data': input_data['data'], 'data': input_data['data'],
'timestamp': input_data['timestamp'], 'timestamp': input_data['timestamp'],
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'prediction_confidence': input_data['prediction_confidence'], 'prediction_confidence': input_data['prediction_confidence'],
'prediction_store_policy': input_data['prediction_store_policy'], 'prediction_store_policy': input_data['prediction_store_policy'],
**metadata, 'model_name': input_data['model_name'],
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -189,9 +193,10 @@ async def test_run_none_path_flag_with_transformed_data(
call( call(
Activities.format_transformed_data, Activities.format_transformed_data,
{ {
**metadata,
'data': input_data['transformed_data'], 'data': input_data['transformed_data'],
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
**metadata, 'model_name': input_data['model_name'],
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=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 # Assert - start_activity_method for transformed data export
workflow_mock.start_activity_method.assert_called_once_with( workflow_mock.start_activity_method.assert_called_once_with(
Activities.export_data_to_postgres, Activities.export_payload_to_postgres,
{ {
**metadata,
'schema': input_data['schema'], 'schema': input_data['schema'],
'table_name': input_data['transform_table_name'], 'table_name': input_data['transform_table_name'],
'data': transformed_data, 'data': transformed_data,
@@ -210,7 +216,6 @@ async def test_run_none_path_flag_with_transformed_data(
'column': 'timestamp', 'column': 'timestamp',
'format': DATETIME_FORMAT_WITH_TZ, 'format': DATETIME_FORMAT_WITH_TZ,
}, },
**metadata,
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=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'}, 'data': {'test': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0, 'prediction_confidence': 0,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
@@ -311,11 +317,11 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction
call( call(
Activities.format_default_prediction, Activities.format_default_prediction,
{ {
**metadata,
'timestamp': input_data['timestamp'], 'timestamp': input_data['timestamp'],
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'prediction_confidence': input_data['prediction_confidence'], 'prediction_confidence': input_data['prediction_confidence'],
'comment': input_data['comment'], 'comment': input_data['comment'],
**metadata,
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -389,6 +395,7 @@ async def test_run_none_path_flag_with_pi_web_api(workflow_mock, format_and_expo
'data': {'test': 'data'}, 'data': {'test': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0, 'prediction_confidence': 0,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', '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( call(
Activities.format_prediction, Activities.format_prediction,
{ {
**metadata,
'data': input_data['data'], 'data': input_data['data'],
'timestamp': input_data['timestamp'], 'timestamp': input_data['timestamp'],
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'prediction_confidence': input_data['prediction_confidence'], 'prediction_confidence': input_data['prediction_confidence'],
'prediction_store_policy': input_data['prediction_store_policy'], 'prediction_store_policy': input_data['prediction_store_policy'],
**metadata, 'model_name': input_data['model_name'],
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -496,6 +504,7 @@ async def test_run_none_path_flag_with_pi_web_api_and_opc(
'data': {'test': 'data'}, 'data': {'test': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0, 'prediction_confidence': 0,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', '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'}, 'data': {'test': 'data'},
'timestamp': '2021-01-01', 'timestamp': '2021-01-01',
'model_id': 1, 'model_id': 1,
'model_name': metadata['metadata']['model_name'],
'prediction_confidence': 0, 'prediction_confidence': 0,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', '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 from pytest import fixture, mark
@@ -27,9 +27,12 @@ metadata = {
async def test_run(workflow_mock, prediction_process): async def test_run(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(return_value=False) prediction_process.path_flag_handler = AsyncMock(return_value=False)
# Arrange # Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = { input_data = {
'metadata': metadata, 'metadata': metadata,
'data': {'test': 'data'}, 'data': data_payload,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
'transform_table_name': 'test_transform_table', 'transform_table_name': 'test_transform_table',
@@ -47,7 +50,6 @@ async def test_run(workflow_mock, prediction_process):
# Mock the activity responses # Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [ workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate ('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform) # mlflow_response_gate (transform)
@@ -63,21 +65,7 @@ async def test_run(workflow_mock, prediction_process):
await prediction_process.run(input_data) await prediction_process.run(input_data)
# Assert # Assert
assert workflow_mock.execute_local_activity_method.call_count == 7 assert workflow_mock.execute_local_activity_method.call_count == 6
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,
)
]
)
workflow_mock.execute_local_activity_method.assert_has_calls( workflow_mock.execute_local_activity_method.assert_has_calls(
[ [
call( call(
@@ -102,7 +90,6 @@ async def test_run(workflow_mock, prediction_process):
'data': input_data['data'], 'data': input_data['data'],
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], 'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=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): async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(return_value=True) prediction_process.path_flag_handler = AsyncMock(return_value=True)
# Arrange # Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = { input_data = {
'metadata': metadata, 'metadata': metadata,
'data': {'test': 'data'}, 'data': data_payload,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
'transform_table_name': 'test_transform_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 # Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [ workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('stop', 0.95, 'Input data with bad quality'), # input_gate ('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) await prediction_process.run(input_data)
# Assert # 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( 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( call(
Activities.input_gate, 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): async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, True]) prediction_process.path_flag_handler = AsyncMock(side_effect=[False, True])
# Arrange # Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = { input_data = {
'metadata': metadata, 'metadata': metadata,
'data': {'test': 'data'}, 'data': data_payload,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
'transform_table_name': 'test_transform_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 # Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [ workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('repeat', 0.95, 'Input data with bad quality'), # input_gate ('repeat', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
('continue', 0.95, 'Error'), # mlflow_response_gate (transform) ('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) await prediction_process.run(input_data)
# Assert # Assert
assert workflow_mock.execute_local_activity_method.call_count == 4 assert workflow_mock.execute_local_activity_method.call_count == 3
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( workflow_mock.execute_local_activity_method.assert_has_calls(
[ [
call( call(
@@ -325,7 +294,6 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
'data': input_data['data'], 'data': input_data['data'],
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], 'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata, **metadata,
}, },
retry_policy=ANY, 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): async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process):
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, True]) prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, True])
# Arrange # Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = { input_data = {
'metadata': metadata, 'metadata': metadata,
'data': {'test': 'data'}, 'data': data_payload,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
'transform_table_name': 'test_transform_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 # Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [ workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate ('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform) # 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) await prediction_process.run(input_data)
# Assert # 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( workflow_mock.execute_local_activity_method.assert_has_calls(
[ [
call( call(
@@ -426,7 +383,6 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
'data': input_data['data'], 'data': input_data['data'],
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], 'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata, **metadata,
}, },
retry_policy=ANY, 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): 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]) prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, False, True])
# Arrange # Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.last_timestamp = '2024-01-01'
input_data = { input_data = {
'metadata': metadata, 'metadata': metadata,
'data': {'test': 'data'}, 'data': data_payload,
'schema': 'test_schema', 'schema': 'test_schema',
'table_name': 'test_table', 'table_name': 'test_table',
'transform_table_name': 'test_transform_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 # Mock the activity responses
workflow_mock.execute_local_activity_method.side_effect = [ workflow_mock.execute_local_activity_method.side_effect = [
'2024-01-01', # get_last_timestamp
('continue', 0.95, 'Input data with bad quality'), # input_gate ('continue', 0.95, 'Input data with bad quality'), # input_gate
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
# mlflow_response_gate (transform) # 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) await prediction_process.run(input_data)
# Assert # Assert
assert workflow_mock.execute_local_activity_method.call_count == 7 assert workflow_mock.execute_local_activity_method.call_count == 6
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( workflow_mock.execute_local_activity_method.assert_has_calls(
[ [
call( call(
@@ -544,7 +489,6 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
'data': input_data['data'], 'data': input_data['data'],
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], 'model_config': input_data['model_config'],
'key_prefix': 'predictions/test_schedule',
**metadata, **metadata,
}, },
retry_policy=ANY, retry_policy=ANY,
@@ -821,3 +765,49 @@ async def test_path_flag_handler_unknown(workflow_mock, prediction_process):
assert result is False assert result is False
workflow_mock.execute_activity_method.assert_not_called() workflow_mock.execute_activity_method.assert_not_called()
workflow_mock.execute_child_workflow.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 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( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'data': {'a': [1]}, 'success': True}, storage_result,
{'success': True, 'experiment': 'test_experiment'}, {'success': True, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -77,7 +80,7 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
Activities.retrain_model, Activities.retrain_model,
{ {
**metadata, **metadata,
'data': {'data': {'a': [1]}, 'success': True}, 'data': storage_result,
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], '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( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'success': False, 'message': 'No data returned from query'}, storage_result,
{'success': True, 'experiment': 'test_experiment'}, {'success': True, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -174,6 +180,9 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
] ]
) )
from pytest import raises
with raises(ValueError, match='No data returned from query'):
await minimal_retrain.run(input_data) await minimal_retrain.run(input_data)
workflow_mock.execute_activity_method.assert_called_once_with( workflow_mock.execute_activity_method.assert_called_once_with(
@@ -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( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'data': {'a': [1]}, 'success': True}, storage_result,
{'success': False, 'experiment': 'test_experiment'}, {'success': False, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -247,7 +259,7 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
Activities.retrain_model, Activities.retrain_model,
{ {
**metadata, **metadata,
'data': {'data': {'a': [1]}, 'success': True}, 'data': storage_result,
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], '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'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'key_prefix': f"predictions/{input_data['schedule_name']}",
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=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