SIENTIAPDE-1712

Implement MinIO Offload and Retention Features

- Added configuration options for MinIO retention hours and offload threshold in README.
- Introduced MinIO payload offloading for large DataFrame-derived payloads, storing them as parquet files.
- Updated activities to utilize MinIO for data loading and cleanup, including new methods for offloading and retention management.
- Refactored existing activities to integrate MinIO functionality, ensuring compatibility with previous workflows.
- Removed the legacy MinioRepository class, consolidating MinIO operations under a new manager structure.
- Updated requirements to use the latest version of the sientia-dataops-library.
This commit is contained in:
vitor-aignosi
2026-03-19 17:29:43 -03:00
parent 9dc3cb3ba0
commit 981ac700d4
25 changed files with 994 additions and 681 deletions

View File

@@ -716,10 +716,80 @@ The Laborious system exposes comprehensive Prometheus metrics for operational vi
| `HTTP_METRICS_PORT` | Prometheus metrics port | `9090` | No | | `HTTP_METRICS_PORT` | Prometheus metrics port | `9090` | No |
| `HTTP_SDK_METRICS_PORT` | Temporal SDK metrics port | `9091` | No | | `HTTP_SDK_METRICS_PORT` | Temporal SDK metrics port | `9091` | No |
| `POD_ID` | Kubernetes pod identifier | `None` | No | | `POD_ID` | Kubernetes pod identifier | `None` | No |
| `SIENTIA_MINIO_RETENTION_HOURS` | Retention window for offloaded MinIO objects | `168` | No |
| `SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES` | Offload threshold for DataFrame-derived payloads | `int(1.5 * 1024 * 1024)` | No |
### MinIO Payload Offload & Retention
Laborious uses MinIO to prevent Temporal workflow history from carrying very large in-memory payloads (pandas `DataFrame`-derived dicts).
Whenever a payload exceeds a configurable size threshold, it is stored as a parquet file in MinIO and the workflow history only keeps a lightweight reference.
Notes:
- `SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES` supports:
- Integer bytes (e.g. `"1572864"`)
- Float MiB (e.g. `"1.5"`), converted to bytes as `MiB * 1024 * 1024`
- Fallback behavior uses `1.5 MiB` when the env var is missing or invalid.
#### Wire Contract: `MinioDataFramePayload`
The payload is implemented in `laborious/utils/models/minio_dataframe_payload.py`.
The dataclass does **not** store a pandas `DataFrame` field.
Instead, the `DataFrame` is only used at build time by:
- `MinioDataFramePayload.from_dataframe(...)`
- `MinioDataFramePayload.from_dataframe_to_dict(...)`
After evaluation, the payload is serialized for Temporal as a flat dict:
- **Inline path**: `data` contains `df.to_dict()`, and MinIO keys (`object_key`, `bucket`, ...) are absent / `None`.
- **MinIO path**: the dict contains:
- `bucket`
- `object_key` (full MinIO object name returned by `MinioRepository.upload_file`)
- `object_prefix` (directory prefix used for cleanup listing; relative to the repository namespace)
- `uri` (best-effort `s3://<bucket>/<...>` string)
- `data` is omitted / set to `None`.
When an activity needs pandas operations, it resolves references using:
- `MinioDataFramePayload.dataframe_from_wire(...)`
#### MinIO Object Naming (Retention Parsing)
MinIO object basename (required convention):
`{model_name}-{operation}-{timestamp}.parquet`
Where:
- `model_name`: model identifier used by the pipeline
- `operation`: `initial` (SQL/query load before transform) or `transform` (after MLFlow transform)
- `timestamp`: `DATETIME_FORMAT_FILENAME` from `sientia_do.temporal.constants`
The relative object key (under the repository namespace) is always shaped as:
`training_datasets/{model_name}/{basename}`
Retention cleanup parses timestamps from the basename using the `-initial-` / `-transform-` anchors.
`model_name` may contain hyphens; parsing is resilient to it.
#### Workflows / Activities Integration
Predictions batch uses MinIO offload as follows:
1. `predictions_batch` calls `Activities.load_query_with_minio_offload`
- On success, it puts the serialized `MinioDataFramePayload` dict into `prediction_input["data"]`.
2. `sub_workflows/prediction_process`
- Tracks which MinIO prefixes were referenced for offloaded payloads.
- Runs `Activities.cleanup_minio_objects_expired` in a `finally` block (only when MinIO offload happened).
3. `laborious/activities/gates.py` and `laborious/activities/mlflow.py`
- Resolve offloaded payloads transparently before constructing pandas `DataFrame` objects.
#### Legacy: `query_to_minio` (Minimal Retrain)
`Storage.query_to_minio` is intentionally kept with its legacy behavior for `minimal_retrain`.
It always uploads parquet and returns `{success, object_key, uri}`.
It is not used by predictions batch MinIO offload, and its objects are not part of the retention parser described above.
Legacy MinIO object layout (relative key):
`training_datasets/{model_name}/{object_prefix}_{timestamp}.parquet` where `object_prefix` is sanitized
(slashes replaced by underscores) to keep a stable model-level directory.
### OPC Configuration ### OPC Configuration
For multiple OPC servers, use the `OPC_CONFIG` environment variable: For multiple OPC servers, use the `OPC_CONFIG` environment variable:

View File

@@ -3,6 +3,7 @@ Pytest configuration and fixtures for E2E tests.
""" """
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
from io import BytesIO
import pandas as pd import pandas as pd
import pytest import pytest
@@ -216,11 +217,28 @@ def metrics_controller(mock_logger):
def mock_minio_repository(): def mock_minio_repository():
"""Mock MinIO repository for object storage operations.""" """Mock MinIO repository for object storage operations."""
mock_repo = MagicMock() mock_repo = MagicMock()
# Provide at least valid parquet bytes so that MinioDataFramePayload.retrieve()
# can decode the payload if offloading is exercised in an integration scenario.
parquet_df = pd.DataFrame({'a': [1]})
parquet_buffer = BytesIO()
parquet_df.to_parquet(parquet_buffer, engine='pyarrow', index=True)
parquet_bytes = parquet_buffer.getvalue()
# Mock repository methods # sientia_do MinioRepository API
mock_repo.put_parquet_from_dataframe = AsyncMock(return_value='test-object-key') mock_repo.bucket = 'test-bucket'
mock_repo.get_parquet_as_dataframe = AsyncMock(return_value=pd.DataFrame()) mock_repo.upload_file = AsyncMock(
mock_repo.minio_bucket = 'test-bucket' side_effect=lambda file_bytes, relative_key, content_type='application/octet-stream', bucket=None, metadata=None: {
'minio_object_name': f'sientia/streamlit-connectors/{relative_key}',
'original_filename': relative_key.rsplit('/', 1)[-1],
'uploaded_at': '2024-01-01T00:00:00Z',
'sha256_hash': 'deadbeef',
}
)
mock_repo.download_file = AsyncMock(return_value=parquet_bytes)
mock_repo.list_objects = AsyncMock(return_value=[])
mock_repo.delete_file = AsyncMock()
mock_repo.close = MagicMock()
return mock_repo return mock_repo
@@ -261,7 +279,7 @@ def patch_create_engine(postgres_engine):
@pytest_asyncio.fixture @pytest_asyncio.fixture
def patch_minio_repository(mock_minio_repository): def patch_minio_repository(mock_minio_repository):
"""Patch MinioRepository to return mock.""" """Patch MinioRepository to return mock."""
with patch('laborious.utils.repository.minio_repository.MinioRepository', return_value=mock_minio_repository): with patch('sientia_do.repository.minio_repository.MinioRepository', return_value=mock_minio_repository):
yield yield
@pytest_asyncio.fixture @pytest_asyncio.fixture
@@ -433,6 +451,8 @@ async def temporal_worker(temporal_test_env, test_activities):
workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction], workflows=[PredictionsBatch, PredictionProcess, FormatAndExportPrediction],
activities=[ activities=[
test_activities.load_custom_query, test_activities.load_custom_query,
test_activities.load_query_with_minio_offload,
test_activities.cleanup_minio_objects_expired,
test_activities.get_last_timestamp, test_activities.get_last_timestamp,
test_activities.input_gate, test_activities.input_gate,
test_activities.request_transform, test_activities.request_transform,

View File

@@ -6,6 +6,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.repository.minio_repository import MinioRepository
from laborious.activities.api import API from laborious.activities.api import API
from laborious.activities.gates import Gates from laborious.activities.gates import Gates
@@ -13,6 +14,7 @@ with workflow.unsafe.imports_passed_through():
from laborious.activities.model_metrics import ModelMetrics from laborious.activities.model_metrics import ModelMetrics
from laborious.activities.opc import OPC from laborious.activities.opc import OPC
from laborious.activities.storage import Storage from laborious.activities.storage import Storage
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
@@ -73,6 +75,16 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
""" """
metrics_controller = MetricsController(logger=logger) metrics_controller = MetricsController(logger=logger)
minio_repository = MinioRepository(
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=metrics_controller,
)
# Initialize parent classes # Initialize parent classes
Storage.__init__( Storage.__init__(
self, self,
@@ -83,7 +95,8 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
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=minio_repository,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
@@ -95,7 +108,7 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
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=minio_repository,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
@@ -103,6 +116,7 @@ class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):
Gates.__init__( Gates.__init__(
self, self,
minio_repository=minio_repository,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller, metrics_controller=metrics_controller,

View File

@@ -1,5 +1,8 @@
from sientia_do.repository.minio_repository import MinioRepository
from temporalio import activity, workflow from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
import traceback import traceback
from collections.abc import Callable, Mapping from collections.abc import Callable, Mapping
@@ -15,6 +18,7 @@ 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 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,
@@ -63,7 +67,7 @@ mlflow_content_path_confidence: Mapping[str, int] = {
} }
class Gates(SientiaMonitoring): class Gates(MinioManager):
""" """
Data quality gates and filtering activities for the Laborious system. Data quality gates and filtering activities for the Laborious system.
@@ -83,11 +87,14 @@ class Gates(SientiaMonitoring):
mlflow_content_filter_functions (dict): Mapping of MLFlow content filter names to functions mlflow_content_filter_functions (dict): Mapping of MLFlow content filter names to functions
""" """
minio_repository: MinioRepository | None = None
def __init__( def __init__(
self, self,
logger: Logger, minio_repository: MinioRepository | None = None,
notification_handler: NotificationHandler, logger: Logger | None = None,
metrics_controller: MetricsController, notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
): ):
""" """
Initialize data quality gates with logging and notification capabilities. Initialize data quality gates with logging and notification capabilities.
@@ -99,13 +106,14 @@ class Gates(SientiaMonitoring):
Raises: Raises:
Exception: If BaseActivity initialization fails Exception: If BaseActivity initialization fails
""" """
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) MinioManager.__init__(self, minio_repository, logger, notification_handler, metrics_controller)
def close(self) -> None: def close(self) -> None:
""" """
Close the gates activity and clean up resources. Close the gates activity and clean up resources.
""" """
SientiaMonitoring.shutdown(self)
MinioManager.close(self)
def __del__(self): def __del__(self):
self.close() self.close()
@@ -149,7 +157,8 @@ class Gates(SientiaMonitoring):
self.info('Performing input gate...', metadata) self.info('Performing input gate...', metadata)
filters = input_data['filters'] filters = input_data['filters']
data = DataFrame(input_data['data']) payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
path_priority = input_data['path_priority'] path_priority = input_data['path_priority']
filter_output = [] filter_output = []
@@ -183,6 +192,9 @@ class Gates(SientiaMonitoring):
return path_flag, input_path_confidence[path_flag], 'Input data with bad quality' return path_flag, input_path_confidence[path_flag], 'Input data with bad quality'
self.info('Nothing was filtered by the input gate', metadata) self.info('Nothing was filtered by the input gate', metadata)
del data
return None, 0, '' return None, 0, ''
@activity.defn(name='mlflow_response_gate') @activity.defn(name='mlflow_response_gate')
@@ -223,7 +235,10 @@ class Gates(SientiaMonitoring):
self.info('Performing mlflow response gate...', metadata) self.info('Performing mlflow response gate...', metadata)
filters = input_data['filters'] filters = input_data['filters']
data = input_data['data']
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
gate_type = input_data['type'] gate_type = input_data['type']
path_priority = input_data['path_priority'] path_priority = input_data['path_priority']
@@ -233,13 +248,16 @@ class Gates(SientiaMonitoring):
self.debug(f'Filters: {filters}', metadata) self.debug(f'Filters: {filters}', metadata)
comments = [] comments = []
status = payload.status or {}
for fil, config in filters.items(): for fil, config in filters.items():
if fil not in mlflow_response_filter_functions: if fil not in mlflow_response_filter_functions:
continue continue
try: try:
if mlflow_response_filter_functions[fil](data, config): if mlflow_response_filter_functions[fil](status, config):
filter_output.append(config['policy']) filter_output.append(config['policy'])
comments.append(data['content']['message']) comments.append(status['message'])
await self.send_notification_async( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}', notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}',
@@ -265,6 +283,9 @@ class Gates(SientiaMonitoring):
return path_flag, mlflow_response_path_confidence[path_flag], ', '.join(comments) return path_flag, mlflow_response_path_confidence[path_flag], ', '.join(comments)
self.info('Nothing was filtered by the mlflow response gate', metadata) self.info('Nothing was filtered by the mlflow response gate', metadata)
del data
return None, 0, '' return None, 0, ''
@activity.defn(name='mlflow_content_gate') @activity.defn(name='mlflow_content_gate')
@@ -305,7 +326,10 @@ class Gates(SientiaMonitoring):
self.info('Performing mlflow content gate...', metadata) self.info('Performing mlflow content gate...', metadata)
filters = input_data['filters'] filters = input_data['filters']
data = DataFrame(input_data['data'])
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
gate_type = input_data['type'] gate_type = input_data['type']
path_priority = input_data['path_priority'] path_priority = input_data['path_priority']
@@ -349,6 +373,9 @@ class Gates(SientiaMonitoring):
) )
self.info('Nothing was filtered by the mlflow content gate', metadata) self.info('Nothing was filtered by the mlflow content gate', metadata)
del data
return None, 0, '' return None, 0, ''
def get_prediction_store_policy( def get_prediction_store_policy(
@@ -403,7 +430,7 @@ class Gates(SientiaMonitoring):
return policy_type, int(policy_value) return policy_type, int(policy_value)
@activity.defn(name='format_transformed_data') @activity.defn(name='format_transformed_data')
async def format_transformed_data(self, input_data: dict[str, Any]) -> dict: async def format_transformed_data(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
""" """
Format transformed data for storage and export operations. Format transformed data for storage and export operations.
@@ -438,7 +465,8 @@ class Gates(SientiaMonitoring):
self.info('Formatting transformed data...', metadata) self.info('Formatting transformed data...', metadata)
data = DataFrame(input_data['data']) payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
data['timestamp'] = data.index data['timestamp'] = data.index
data = data.reset_index(drop=True) data = data.reset_index(drop=True)
@@ -446,7 +474,13 @@ class Gates(SientiaMonitoring):
data = data.melt(id_vars='timestamp', var_name='variable', value_name='value') data = data.melt(id_vars='timestamp', var_name='variable', value_name='value')
data['model_id'] = model_id data['model_id'] = model_id
return data.to_dict() return await MinioDataFramePayload.from_dataframe(
dataframe=data,
minio_repo=self.minio_repository,
model_name=input_data['model_name'],
operation='transform',
workflow_metadata=metadata
)
@activity.defn(name='format_prediction') @activity.defn(name='format_prediction')
async def format_prediction(self, input_data: dict[str, Any]) -> dict: async def format_prediction(self, input_data: dict[str, Any]) -> dict:
@@ -477,7 +511,8 @@ class Gates(SientiaMonitoring):
prediction_store_policy = input_data['prediction_store_policy'] prediction_store_policy = input_data['prediction_store_policy']
self.info('Formatting prediction...', metadata) self.info('Formatting prediction...', metadata)
data = DataFrame(input_data['data']) payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
# Create timestamp column from index and reset index # Create timestamp column from index and reset index
data['timestamp'] = data.index data['timestamp'] = data.index
@@ -566,6 +601,7 @@ class Gates(SientiaMonitoring):
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:
""" """
@@ -636,45 +672,6 @@ class Gates(SientiaMonitoring):
return report.to_dict() return report.to_dict()
@activity.defn(name='get_last_timestamp')
async def get_last_timestamp(self, input_data: dict[str, Any]) -> str:
"""
Extract the most recent timestamp from prediction data.
This method analyzes prediction data to find the latest timestamp,
enabling incremental processing and data continuity tracking.
It handles empty datasets gracefully by returning the current time
as a fallback timestamp.
The method is essential for:
1. Incremental data processing workflows
2. Data continuity validation
3. Timestamp-based data loading optimization
4. Workflow execution tracking
Args:
input_data (dict): Input data containing:
- data (dict[str, Any]): Prediction data to analyze
Returns:
str: Formatted timestamp string in UTC with timezone
"""
metadata = input_data['metadata']
self.info('Getting last timestamp...', metadata)
data = DataFrame(input_data['data'])
self.debug(f'Input data: {data.head(5).to_string()}', metadata)
if data.empty:
return now().strftime(DATETIME_FORMAT_WITH_TZ)
max_timestamp = max(data['timestamp'].values.tolist())
self.info(f'Last timestamp: {max_timestamp}', metadata)
return max_timestamp
@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,11 +1,15 @@
from re import M
from temporalio import activity, workflow from temporalio import activity, workflow
from laborious.utils.repository.minio_manager import MinioManager
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
import traceback import traceback
from typing import Any from typing import Any
import numpy as np import numpy as np
from pandas import DataFrame, to_datetime from io import BytesIO
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
@@ -19,11 +23,12 @@ 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.repository.minio_repository import MinioRepository 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
class MLFlow(SientiaMonitoring): class MLFlow(MinioManager):
""" """
MLFlow integration activities for model inference operations. MLFlow integration activities for model inference operations.
@@ -47,11 +52,11 @@ class MLFlow(SientiaMonitoring):
mlflow_host: str, mlflow_host: str,
mlflow_port: int, mlflow_port: int,
mlflow_username: str, mlflow_username: str,
minio_config: dict[str, Any],
mlflow_password: str, mlflow_password: str,
logger: Logger, minio_repository: MinioRepository | None = None,
notification_handler: NotificationHandler, logger: Logger | None = None,
metrics_controller: MetricsController, notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
): ):
""" """
Initialize MLFlow activities with server configuration. Initialize MLFlow activities with server configuration.
@@ -67,7 +72,7 @@ class MLFlow(SientiaMonitoring):
Raises: Raises:
Exception: If MLFlowRepository initialization fails Exception: If MLFlowRepository initialization fails
""" """
SientiaMonitoring.__init__(self, 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
@@ -82,32 +87,17 @@ class MLFlow(SientiaMonitoring):
metrics_controller, metrics_controller,
) )
if not hasattr(self, 'minio_repository'):
self.minio_repository: MinioRepository | None = None
if self.minio_repository is None:
self.minio_repository = MinioRepository(
logger=logger,
notification_handler=notification_handler,
minio_endpoint_url=minio_config['endpoint_url'],
minio_access_key=minio_config['access_key'],
minio_secret_key=minio_config['secret_key'],
minio_region_name=minio_config['region_name'],
minio_default_bucket=minio_config['default_bucket'],
metrics_controller=metrics_controller,
)
def close(self) -> None: def close(self) -> None:
""" """
Close the MLFlow activity and clean up resources. Close the MLFlow activity and clean up resources.
""" """
SientiaMonitoring.shutdown(self) MinioManager.close(self)
def __del__(self): def __del__(self):
self.close() self.close()
@activity.defn(name='request_transform') @activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]: async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
""" """
Transform input data using MLFlow models. Transform input data using MLFlow models.
@@ -138,7 +128,10 @@ class MLFlow(SientiaMonitoring):
""" """
metadata = input_data['metadata'] metadata = input_data['metadata']
self.info('Transforming data...', metadata) self.info('Transforming data...', metadata)
data = DataFrame(input_data['data'])
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name'] model_name = input_data['model_name']
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
@@ -176,10 +169,29 @@ class MLFlow(SientiaMonitoring):
self.info('Data transformed successfully', metadata) self.info('Data transformed successfully', metadata)
return response_data if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
operation='transform',
status=response_data,
workflow_metadata=metadata,
)
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
operation='transform',
workflow_metadata=metadata,
status={
'success': True,
},
)
@activity.defn(name='request_predict') @activity.defn(name='request_predict')
async def request_predict(self, input_data: dict[str, Any]) -> dict[str, Any]: async def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
""" """
Execute predictions using MLFlow models. Execute predictions using MLFlow models.
@@ -210,7 +222,10 @@ class MLFlow(SientiaMonitoring):
""" """
metadata = input_data['metadata'] metadata = input_data['metadata']
self.info('Predicting data...', metadata) self.info('Predicting data...', metadata)
data = DataFrame(input_data['data'])
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name'] model_name = input_data['model_name']
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
@@ -236,7 +251,28 @@ class MLFlow(SientiaMonitoring):
self.info('Data predicted successfully', metadata) self.info('Data predicted successfully', metadata)
return response_data
if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
operation='predict',
status=response_data,
workflow_metadata=metadata,
)
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
operation='predict',
workflow_metadata=metadata,
status={
'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]:
@@ -275,14 +311,24 @@ class MLFlow(SientiaMonitoring):
raise ValueError('Minio repository not initialized') raise ValueError('Minio repository not initialized')
metadata = input_data['metadata'] metadata = input_data['metadata']
object_key = input_data['object_key']
self.info(f'Loading retrain data from Key: {object_key}', metadata)
try: try:
data = await self.minio_repository.get_parquet_as_dataframe( if 'data' in input_data:
object_key=object_key, metadata=metadata # New path: payload-based retrain input (inline or MinIO offloaded).
) data = await MinioDataFramePayload.dataframe_from_wire(
input_data['data'],
self.minio_repository,
metadata,
)
else:
# Backward compatibility: legacy query_to_minio contract.
object_key = input_data['object_key']
self.info(f'Loading retrain data from Key: {object_key}', metadata)
file_bytes = await self.minio_repository.download_file(
object_name=object_key,
metadata=metadata,
)
data = read_parquet(BytesIO(file_bytes))
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,8 +1,15 @@
import json
from temporalio import activity, workflow from temporalio import activity, workflow
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 io import BytesIO
from os import getenv
from typing import Any from typing import Any
import pandas as pd import pandas as pd
@@ -13,15 +20,20 @@ with workflow.unsafe.imports_passed_through():
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.repository.minio_repository import MinioRepository 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'})
class Storage(Postgres): class Storage(Postgres, MinioManager):
""" """
Extensions for Postgres activities with a helper to export query results Extensions for Postgres activities with a helper to export query results
directly to MinIO as Parquet and return the object name. directly to MinIO as Parquet and return the object name.
""" """
minio_repository: MinioRepository | None = None
def __init__( def __init__(
self, self,
host: str, host: str,
@@ -31,12 +43,15 @@ class Storage(Postgres):
dbname: str, dbname: str,
min_connections: int, min_connections: int,
max_connections: int, max_connections: int,
minio_config: dict[str, Any], retention_hours: int = 24,
logger: Logger, minio_repository: MinioRepository | None = None,
notification_handler: NotificationHandler, logger: Logger | None = None,
metrics_controller: MetricsController, notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
): ):
super().__init__( self.retention_hours = retention_hours
Postgres.__init__(
self,
host=host, host=host,
port=port, port=port,
user=user, user=user,
@@ -49,20 +64,144 @@ class Storage(Postgres):
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
) )
if not hasattr(self, 'minio_repository'): MinioManager.__init__(self, minio_repository, logger, notification_handler, metrics_controller)
self.minio_repository: MinioRepository | None = None
@activity.defn(name='load_query_with_minio_offload')
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.
Args (input_data):
metadata (dict): Workflow metadata (same as load_custom_query).
query (str): SQL query.
datetime_columns (list[str], optional): Datetime column names.
model_name (str): Model name for object key basename.
key_prefix (str, optional): Directory prefix inside the bucket.
size_threshold_bytes (int, optional): Override env offload threshold.
Returns:
dict[str, Any]: Flat ``MinioDataFramePayload`` dict or ``success: False`` on failure.
"""
if self.minio_repository is None: if self.minio_repository is None:
self.minio_repository = MinioRepository( raise ValueError('Minio repository not initialized')
logger=logger,
notification_handler=notification_handler, metadata: dict = input_data.get('metadata', {})
minio_endpoint_url=minio_config['endpoint_url'], model_name = input_data['model_name']
minio_access_key=minio_config['access_key'],
minio_secret_key=minio_config['secret_key'], rows = await self.load_custom_query(
minio_region_name=minio_config['region_name'], input_data,
minio_default_bucket=minio_config['default_bucket'], )
metrics_controller=metrics_controller, if not rows:
self.error('load_query_with_minio_offload failed: No data returned from query', metadata)
dataframe = None
else:
dataframe = pd.DataFrame(rows)
return await MinioDataFramePayload.from_dataframe(
dataframe,
minio_repo=self.minio_repository,
workflow_metadata=metadata,
model_name=model_name,
operation='initial',
)
@activity.defn(name='export_payload_to_postgres')
async def export_payload_to_postgres(self, input_data: dict[str, Any]) -> dict:
"""
Export a payload to PostgreSQL.
"""
metadata = input_data.get('metadata')
payload: MinioDataFramePayload = input_data['data']
data = await payload.retrieve(self.minio_repository, metadata)
return await self.export_data_to_postgres(
{
**input_data,
'data': data,
}
)
@activity.defn(name='cleanup_minio_objects_expired')
async def cleanup_minio_objects_expired(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Delete objects under the given prefixes that are older than the retention window.
Args (input_data):
metadata (dict): Workflow metadata for logging and metrics.
prefixes (list[str]): Key prefixes to scan (one level or subtree per prefix).
Returns:
dict[str, Any]: ``success``, ``deleted_count``, and optional ``message``.
"""
if self.minio_repository is None:
raise ValueError('Minio repository not initialized')
metadata = input_data.get('metadata', {})
prefix = input_data['prefix']
base = now()
cutoff = (base.replace(tzinfo=None) if base.tzinfo else base) - timedelta(
hours=self.retention_hours
)
report: dict[str, Any] = {
'failed': {},
'deleted': {},
'failed_count': 0,
'deleted_count': 0,
}
try:
keys = await self.minio_repository.list_objects(
prefix=prefix,
recursive=True,
metadata=metadata,
) )
for key in keys:
try:
ts = MinioDataFramePayload.parse_object_timestamp(key)
if ts is None:
continue
if ts >= cutoff:
continue
await self.minio_repository.delete_file(
object_name=key,
metadata=metadata,
)
except Exception as e:
report['failed'][key] = {
'success': False,
'message': str(e),
}
report['failed_count'] += 1
continue
report['deleted'][key] = {
'success': True,
'message': 'Deleted',
}
report['deleted_count'] += 1
except Exception as e:
trace = traceback.format_exc()
await self.send_notification_async(
metadata=metadata,
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
message=f'Error cleaning up MinIO objects: {e}',
block='cleanup_minio_objects_expired',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
self.error(trace, metadata)
else:
await self.send_notification_async(
metadata=metadata,
notification_id='CLEANUP_MINIO_OBJECTS_EXPIRED',
message='MinIO objects cleaned up successfully',
block='cleanup_minio_objects_expired',
level=NotificationLevel.INFO,
attachment_content=json.dumps(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]:
@@ -83,11 +222,18 @@ class Storage(Postgres):
raise ValueError('Minio repository not initialized') raise ValueError('Minio repository not initialized')
metadata = input_data.get('metadata', {}) metadata = input_data.get('metadata', {})
model_name = input_data.get('model_name') or metadata.get('model_name') or 'unknown'
object_prefix = input_data.get('object_prefix', 'datasets/retrain') object_prefix = input_data.get('object_prefix', 'datasets/retrain')
timestamp = now().strftime(DATETIME_FORMAT_FILENAME) timestamp = now().strftime(DATETIME_FORMAT_FILENAME)
object_name = f'{object_prefix}_{timestamp}.parquet' # Keep a stable model-level layout for minimal_retrain:
uri = f's3://{self.minio_repository.minio_bucket}/{object_name}' # training_datasets/<model_name>/<filename>
# Sanitize object_prefix to avoid extra subdirectories in the relative key.
safe_prefix = str(object_prefix).strip().strip('/').replace('/', '_')
filename = f'{safe_prefix}_{timestamp}.parquet'
relative_key = f'training_datasets/{model_name}/{filename}'
bucket = getattr(self.minio_repository, 'bucket', 'streamlit-connectors')
uri = f's3://{bucket}/{relative_key}'
try: try:
data = await self.load_custom_query(input_data) data = await self.load_custom_query(input_data)
@@ -98,12 +244,20 @@ class Storage(Postgres):
# Ensure we have a DataFrame # Ensure we have a DataFrame
data = pd.DataFrame(data) data = pd.DataFrame(data)
# Write parquet to memory and upload via persistent client # Convert DataFrame -> parquet bytes, then upload using the new MinIO interface.
await self.minio_repository.store_dataframe_as_parquet( parquet_buffer = BytesIO()
dataframe=data, uri=uri, object_name=object_name, metadata=metadata data.to_parquet(parquet_buffer, engine='pyarrow', index=True)
file_bytes = parquet_buffer.getvalue()
upload_result = await self.minio_repository.upload_file(
file_bytes=file_bytes,
relative_key=relative_key,
metadata=metadata,
) )
return {'success': True, 'object_key': object_name, 'uri': uri} object_key_full = upload_result.get('minio_object_name', relative_key)
uri = f's3://{bucket}/{object_key_full}'
return {'success': True, 'object_key': object_key_full, 'uri': uri}
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(
@@ -121,18 +275,8 @@ class Storage(Postgres):
def close(self) -> None: def close(self) -> None:
"""Close Storage resources (MinIO client and Postgres engine).""" """Close Storage resources (MinIO client and Postgres engine)."""
try: Postgres.close(self)
if hasattr(self, 'minio_repository') and self.minio_repository is not None: MinioManager.close(self)
try:
self.minio_repository.close()
finally:
self.minio_repository = None
finally:
# Ensure Postgres resources are disposed as well
try:
super().close()
except Exception:
self.logger.error('Error closing Postgres resources')
def __del__(self): def __del__(self):
self.close() self.close()

View File

@@ -88,4 +88,5 @@ def build_minio_config() -> dict[str, Any]:
'secret_key': getenv('MINIO_SECRET_KEY', 'minioadmin'), 'secret_key': getenv('MINIO_SECRET_KEY', 'minioadmin'),
'region_name': getenv('MINIO_REGION_NAME', 'us-east-1'), 'region_name': getenv('MINIO_REGION_NAME', 'us-east-1'),
'default_bucket': getenv('MINIO_DEFAULT_BUCKET', 'laborious'), 'default_bucket': getenv('MINIO_DEFAULT_BUCKET', 'laborious'),
'retention_hours': int(getenv('MINIO_RETENTION_HOURS', '24')),
} }

View File

View File

@@ -0,0 +1,230 @@
"""
MinIO-backed DataFrame payload for Temporal workflows.
Data is never stored as a pandas ``DataFrame`` field on the dataclass.
Instead, the DataFrame is only provided as an input to:
`from_dataframe` / `from_dataframe_to_dict`.
At build time, the DataFrame is evaluated for its serialized size; if it exceeds
the configured threshold, it is serialized to parquet bytes and uploaded to MinIO.
Otherwise, it is inlined as a Temporal-friendly ``dict``.
"""
import pickle
import re
from dataclasses import dataclass, field
from datetime import datetime
from io import BytesIO
from os import getenv
from typing import Any, Hashable, Literal
from pandas import DataFrame, read_parquet
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, DATETIME_FORMAT_WITH_TZ, now
from sientia_do.repository.minio_repository import MinioRepository
# Keys that are part of the serialized wire format (not arbitrary metadata).
_SERIALIZED_FIELD_KEYS = frozenset({'data', 'bucket', 'object_key', 'object_prefix', 'uri'})
_OBJECT_TIMESTAMP_PATTERN = re.compile(
r'-(?:initial|transform)-(\d{4}-\d{2}-\d{2}_\d{2}-\d{2}-\d{2})\.parquet$'
)
OFFLOAD_THRESHOLD_BYTES = int(getenv('SIENTIA_MINIO_OFFLOAD_THRESHOLD_BYTES', '1.5')) * 1024 * 1024
# Relative prefix used for storing offloaded training datasets in MinIO.
# It is also the root directory for retention cleanup listing.
TRAINING_DATASETS_PREFIX = 'training_datasets'
OperationKind = Literal['initial', 'transform', 'predict']
def _build_object_key(
model_name: str, operation: OperationKind, timestamp: str
) -> tuple[str, str | None]:
"""
Build the MinIO object key and the directory prefix used for retention listing.
Args:
model_name: Registered model name used in the pipeline.
operation: Either initial (pre-transform load) or transform (post-MLFlow transform).
timestamp: Filename timestamp segment from DATETIME_FORMAT_FILENAME.
Return:
tuple[str, str | None]: Full object key and normalized prefix (or None if at bucket root).
"""
# Naming convention:
# - Directory is always `training_datasets/<model_name>`
# - Filename follows the retention-parsing pattern
basename = f'{model_name}-{operation}-{timestamp}.parquet'
model_dir = model_name.strip().strip('/')
prefix = f'{TRAINING_DATASETS_PREFIX}/{model_dir}'
return f'{prefix}/{basename}', prefix
@dataclass
class MinioDataFramePayload:
"""
Serializable payload after a DataFrame was evaluated: inline tabular dict and/or MinIO keys.
Build from a live DataFrame only via `from_dataframe` / `from_dataframe_to_dict`.
Rehydrate from Temporal via `from_dict`. The DataFrame is not a field on this class.
"""
last_timestamp: str
status: dict[str, Any] | None = None
data: dict[Hashable, Any] | None = None
bucket: str | None = None
object_key: str | None = None
object_prefix: str | None = None
uri: str | None = None
@staticmethod
def estimate_size_bytes(df: DataFrame) -> int:
"""
Approximate serialized size of the DataFrame as the default-orient dict.
Args:
df: DataFrame whose tabular content size is estimated.
Return:
int: Estimated size in bytes (pickle of dict representation).
"""
try:
return len(pickle.dumps(df.to_dict()))
except Exception:
return len(pickle.dumps(df))
@staticmethod
def parse_object_timestamp(object_key: str) -> datetime | None:
"""
Parse the timestamp embedded in the object key basename (before .parquet).
Args:
object_key: S3/MinIO object key whose basename follows
``{model}-{initial|transform}-{DATETIME_FORMAT_FILENAME}.parquet``.
Return:
datetime | None: Parsed UTC-naive datetime from the key, or None if not matched.
"""
basename = object_key.rsplit('/', 1)[-1]
match = _OBJECT_TIMESTAMP_PATTERN.search(basename)
if not match:
return None
try:
return datetime.strptime(match.group(1), DATETIME_FORMAT_FILENAME)
except ValueError:
return None
@staticmethod
def is_offloaded_dict(payload: dict[str, Any]) -> bool:
"""
Return True if the dict represents a MinIO-backed payload without inline data.
Args:
payload: Flat dict possibly produced by to_dict() / from_dataframe_to_dict().
Return:
bool: True when object_key is set and inline data is absent.
"""
if not payload.get('object_key'):
return False
return payload.get('data') is None
@staticmethod
def cleanup_prefix(self) -> str | None:
"""
Return True if cleanup is enabled for this payload.
"""
if self.object_key is not None and self.data is None:
return self.object_prefix
return None
@classmethod
async def from_dataframe(
cls,
dataframe: DataFrame | None,
minio_repo: MinioRepository,
model_name: str,
operation: OperationKind,
status: dict[str, Any] | None = None,
workflow_metadata: dict | None = None,
) -> 'MinioDataFramePayload':
"""
Evaluate the DataFrame size, then either inline dict or upload parquet to MinIO.
The DataFrame is not stored on the returned instance.
Args:
dataframe: Tabular data to evaluate and persist (inline or MinIO).
metadata: Small metadata dict merged into the payload (e.g. success, message).
minio_repo: sientia_do MinioRepository (or compatible) with `upload_file()`.
workflow_metadata: Metadata passed to MinIO store for logging/metrics.
model_name: Registered model name used in the object basename.
operation: Either ``initial`` (query load) or ``transform`` (post-transform).
key_prefix: Backward-compatible parameter (currently ignored for object naming).
size_threshold_bytes: Byte limit before offload. When None, the module-level
environment-derived default is used.
Return:
MinioDataFramePayload: Instance with data and/or MinIO fields set.
"""
if not dataframe or dataframe.empty:
return cls(data=None, last_timestamp=now().strftime(DATETIME_FORMAT_WITH_TZ), status=status)
last_timestamp = max(dataframe['timestamp'].values.tolist())
if cls.estimate_size_bytes(dataframe) <= OFFLOAD_THRESHOLD_BYTES:
return cls(data=dataframe.to_dict(), last_timestamp=last_timestamp)
timestamp = now().strftime(DATETIME_FORMAT_FILENAME)
object_key, object_prefix = _build_object_key(model_name, operation, timestamp)
# Upload using the relative object key. The upstream repository will
# prefix it internally under its MinIO namespace.
parquet_buffer = BytesIO()
dataframe.to_parquet(parquet_buffer, engine='pyarrow', index=True)
file_bytes = parquet_buffer.getvalue()
upload_result = await minio_repo.upload_file(
file_bytes=file_bytes,
relative_key=object_key,
metadata=workflow_metadata,
)
bucket = minio_repo.bucket
object_key_full = upload_result.get('minio_object_name', object_key)
uri = f's3://{bucket}/{object_key_full}' if bucket else None
return cls(
data=None,
bucket=bucket,
object_key=object_key_full,
object_prefix=object_prefix,
uri=uri,
last_timestamp=last_timestamp,
)
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.
Args:
minio_repo: sientia_do MinioRepository (or compatible) with download_file().
workflow_metadata: Metadata passed to MinIO read for logging/metrics.
Return:
dict[str, Any]: Flat dict with data filled (same keys as to_dict after load).
"""
if self.data is not None:
return DataFrame(self.data)
if self.data is None and self.object_key is None:
return DataFrame()
file_bytes = await minio_repo.download_file(
object_name=self.object_key, metadata=workflow_metadata)
df = read_parquet(BytesIO(file_bytes))
return df

View File

@@ -0,0 +1,26 @@
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.notifications.handlers import NotificationHandler
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
class MinioManager(SientiaMonitoring):
minio_repository: MinioRepository | None = None
def __init__(self, minio_repository: MinioRepository | None = None, logger: Logger | None = None, notification_handler: NotificationHandler | None = None, metrics_controller: MetricsController | None = None):
if self.minio_repository is None:
self.minio_repository = minio_repository
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
def close(self) -> None:
"""
Close the MinioManager and clean up resources.
"""
if self.minio_repository is not None:
try:
self.minio_repository.close()
finally:
self.minio_repository = None
SientiaMonitoring.shutdown(self)

View File

@@ -1,215 +0,0 @@
"""
MinIO repository utilities.
This module provides a lightweight repository around a MinIO/S3-compatible
object storage using boto3. It supports creating buckets on demand and
storing/loading pandas DataFrames in Parquet format.
"""
import time
from io import BytesIO
from typing import Any
import boto3
from botocore.config import Config
from botocore.exceptions import ClientError
from pandas import DataFrame, read_parquet
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from laborious import metrics
class MinioRepository(SientiaMonitoring):
"""
Repository for interacting with a MinIO (S3-compatible) object storage.
This class encapsulates a reusable `boto3` S3 client and convenience
helpers to persist and retrieve pandas DataFrames as Parquet files.
Attributes:
storage_options (dict): Options compatible with pandas s3fs usage.
minio_bucket (str): Default bucket name used for operations.
minio_endpoint_url (str): MinIO endpoint URL.
minio_region_name (str): MinIO region name.
s3_client (Any): Reusable S3 client from `boto3`.
logger (Logger): Observability logger.
notification_handler (NotificationHandler): Notifications handler.
"""
def __init__(
self,
minio_endpoint_url: str,
minio_access_key: str,
minio_secret_key: str,
minio_region_name: str,
minio_default_bucket: str,
logger: Logger,
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
"""Initialize the repository and S3 client.
Args:
minio_endpoint_url (str): MinIO endpoint URL.
minio_access_key (str): Access key (AK).
minio_secret_key (str): Secret key (SK).
minio_region_name (str): Region name for the client.
minio_default_bucket (str): Default bucket name to operate on.
logger (Logger): Logger instance for structured logs.
notification_handler (NotificationHandler): Notification handler.
"""
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
# MinIO settings shared with pandas s3fs
self.storage_options = {
'key': minio_access_key,
'secret': minio_secret_key,
'client_kwargs': {'endpoint_url': minio_endpoint_url},
}
self.minio_bucket = minio_default_bucket
self.minio_endpoint_url = minio_endpoint_url
self.minio_region_name = minio_region_name
logger.info(
f'Connecting to MinIO at {self.minio_endpoint_url}, default bucket: {self.minio_bucket}'
)
# Reusable MinIO client
self.s3_client: Any = boto3.client(
's3',
endpoint_url=self.minio_endpoint_url,
aws_access_key_id=self.storage_options['key'],
aws_secret_access_key=self.storage_options['secret'],
region_name=self.minio_region_name,
config=Config(
signature_version='s3v4',
s3={'addressing_style': 'path'},
retries={'max_attempts': 5, 'mode': 'standard'},
connect_timeout=5,
read_timeout=120,
),
)
def close(self):
"""Close the underlying S3 client."""
self.s3_client.close()
async def create_bucket(self, metadata: dict[str, Any]) -> None:
core_labels = {
**self.get_core_labels(metadata, operation_type='create_bucket'),
'bucket_name': self.minio_bucket,
'object_name': '-',
}
self.info(f"Creating bucket '{self.minio_bucket}'", metadata)
start_time = time.time()
try:
self.s3_client.create_bucket(Bucket=self.minio_bucket)
except Exception as e:
await self.emit_metric(metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MINIO_WRITE_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MINIO_WRITE_COUNT, tags=core_labels)
async def ensure_bucket_exists(self, metadata: dict[str, Any]) -> None:
"""Ensure the default bucket exists; create it if missing.
Args:
metadata (dict[str, Any]): Metadata used for structured logging.
"""
self.info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
core_labels = {
**self.get_core_labels(metadata, operation_type='head_bucket'),
'bucket_name': self.minio_bucket,
'object_name': '-',
}
self.info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
start_time = time.time()
try:
self.s3_client.head_bucket(Bucket=self.minio_bucket)
except ClientError:
await self.create_bucket(metadata)
except Exception as e:
await self.emit_metric(metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=core_labels)
raise e
else:
await self.observe_lag(start_time, metrics.MINIO_READ_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MINIO_READ_COUNT, tags=core_labels)
async def store_dataframe_as_parquet(
self, dataframe: DataFrame, uri: str, object_name: str, metadata: dict[str, Any]
):
"""Persist a DataFrame as a Parquet object in the default bucket.
Args:
dataframe (DataFrame): DataFrame to persist.
uri (str): Human-friendly URI used for logging context.
object_name (str): Object key (path/key within the bucket).
metadata (dict[str, Any]): Metadata used for structured logging.
"""
await self.ensure_bucket_exists(metadata)
self.info(f'Storing dataframe as parquet in {uri}', metadata)
buffer = BytesIO()
dataframe.to_parquet(buffer, engine='pyarrow', index=True)
buffer.seek(0)
core_labels = {
**self.get_core_labels(metadata, operation_type='put_object'),
'bucket_name': self.minio_bucket,
'object_name': object_name,
}
start_time = time.time()
try:
self.s3_client.put_object(
Bucket=self.minio_bucket, Key=object_name, Body=buffer.getvalue()
)
except Exception as e:
await self.emit_metric(metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MINIO_WRITE_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MINIO_WRITE_COUNT, tags=core_labels)
self.info(f'Dataframe stored as parquet in {uri}', metadata)
async def get_parquet_as_dataframe(
self, object_key: str, metadata: dict[str, Any]
) -> DataFrame:
"""Load a Parquet object from the default bucket into a DataFrame.
Args:
object_key (str): Object key to retrieve from the bucket.
metadata (dict[str, Any]): Metadata used for structured logging.
Returns:
DataFrame: Loaded DataFrame.
"""
self.info(f'Getting parquet as dataframe from {object_key}', metadata)
core_labels = {
**self.get_core_labels(metadata, operation_type='get_object'),
'bucket_name': self.minio_bucket,
'object_name': object_key,
}
start_time = time.time()
try:
response = self.s3_client.get_object(Bucket=self.minio_bucket, Key=object_key)
except Exception as e:
await self.emit_metric(metric_object=metrics.MINIO_READ_ERROR_COUNT, tags=core_labels)
raise e
await self.observe_lag(start_time, metrics.MINIO_READ_LAG, core_labels)
await self.emit_metric(metric_object=metrics.MINIO_READ_COUNT, tags=core_labels)
# Read the content into a BytesIO buffer to support seek operations
buffer = BytesIO(response['Body'].read())
return read_parquet(buffer)

View File

@@ -1134,7 +1134,7 @@ class MLFlowRepository(SientiaMonitoring):
async def transform( async def transform(
self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict
): ) -> dict[str, Any]:
""" """
Transform data using a cached transformation model. Transform data using a cached transformation model.
@@ -1196,7 +1196,7 @@ class MLFlowRepository(SientiaMonitoring):
transformed_data = self.detect_and_parse_datetime_index(transformed_data, metadata) transformed_data = self.detect_and_parse_datetime_index(transformed_data, metadata)
return {'success': True, 'content': transformed_data.to_dict()} return {'success': True, 'content': transformed_data}
except Exception as e: except Exception as e:
return { return {
@@ -1290,7 +1290,7 @@ class MLFlowRepository(SientiaMonitoring):
predict_data.index = input_index predict_data.index = input_index
predict_data['response_time'] = (end_time - start_time).total_seconds() predict_data['response_time'] = (end_time - start_time).total_seconds()
return {'success': True, 'content': predict_data.to_dict()} return {'success': True, 'content': predict_data}
except Exception as e: except Exception as e:
return { return {

View File

@@ -153,6 +153,7 @@ async def main():
other_workflows=[], other_workflows=[],
activities=[ activities=[
activities.load_custom_query, activities.load_custom_query,
activities.load_query_with_minio_offload,
activities.query_to_minio, activities.query_to_minio,
activities.retrain_model, activities.retrain_model,
activities.update_production_model, activities.update_production_model,
@@ -202,8 +203,9 @@ async def main():
activities.get_last_timestamp, activities.get_last_timestamp,
# OPC # OPC
activities.write_opc_data, activities.write_opc_data,
# Postgres # Postgres / MinIO offload
activities.load_custom_query, activities.load_query_with_minio_offload,
activities.cleanup_minio_objects_expired,
activities.repeat_last_prediction, activities.repeat_last_prediction,
activities.export_data_to_postgres, activities.export_data_to_postgres,
activities.write_metrics, activities.write_metrics,

View File

@@ -73,26 +73,25 @@ class MinimalRetrain:
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
storage_result = await workflow.execute_activity_method( storage_result = await workflow.execute_activity_method(
Activities.query_to_minio, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'query': input_data['query'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': model_name, 'model_name': model_name,
'object_prefix': f'retrain_datasets/{model_name}/data',
}, },
retry_policy=retry_policy, retry_policy=retry_policy,
start_to_close_timeout=timedelta(seconds=600), start_to_close_timeout=timedelta(seconds=600),
) )
if not storage_result['success']: if isinstance(storage_result, dict) and storage_result.get('success') is False:
return return
experiment_response = await workflow.execute_activity_method( experiment_response = await workflow.execute_activity_method(
Activities.retrain_model, Activities.retrain_model,
{ {
**metadata, **metadata,
'object_key': storage_result['object_key'], 'data': storage_result,
'model_name': model_name, 'model_name': model_name,
'model_config': model_config, 'model_config': model_config,
}, },

View File

@@ -83,13 +83,14 @@ class PredictionsBatch:
} }
} }
# Load data using custom query # Load data using custom query with optional MinIO offload for large frames
data = await workflow.execute_local_activity_method( data = await workflow.execute_activity_method(
Activities.load_custom_query, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'query': input_data['query'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'],
}, },
retry_policy=retry_policy, retry_policy=retry_policy,
start_to_close_timeout=timedelta(seconds=300), start_to_close_timeout=timedelta(seconds=300),

View File

@@ -105,6 +105,7 @@ class FormatAndExportPrediction:
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'prediction_confidence': prediction_confidence, 'prediction_confidence': prediction_confidence,
'prediction_store_policy': input_data['prediction_store_policy'], 'prediction_store_policy': input_data['prediction_store_policy'],
'model_name': input_data['model_name'],
}, },
retry_policy=retry_policy, retry_policy=retry_policy,
start_to_close_timeout=timedelta(seconds=60), start_to_close_timeout=timedelta(seconds=60),
@@ -118,13 +119,14 @@ class FormatAndExportPrediction:
**metadata, **metadata,
'data': transformed_data, 'data': transformed_data,
'model_id': input_data['model_id'], 'model_id': input_data['model_id'],
'model_name': input_data['model_name'],
}, },
retry_policy=retry_policy, retry_policy=retry_policy,
start_to_close_timeout=timedelta(seconds=60), start_to_close_timeout=timedelta(seconds=60),
) )
write_transformed_handler = workflow.start_activity_method( write_transformed_handler = workflow.start_activity_method(
Activities.export_data_to_postgres, Activities.export_payload_to_postgres,
{ {
**metadata, **metadata,
'schema': input_data['schema'], 'schema': input_data['schema'],

View File

@@ -2,11 +2,13 @@ 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
from laborious.activities.activities import Activities from laborious.activities.activities import Activities
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
@workflow.defn(name='subworkflow.prediction_process') @workflow.defn(name='subworkflow.prediction_process')
@@ -37,6 +39,8 @@ class PredictionProcess:
8. Export Delegation: Delegates to FormatAndExportPrediction workflow 8. Export Delegation: Delegates to FormatAndExportPrediction workflow
""" """
cleanup_prefixes: set[str] = set()
@workflow.run @workflow.run
async def run(self, input_data: dict[str, Any]): async def run(self, input_data: dict[str, Any]):
""" """
@@ -87,13 +91,39 @@ class PredictionProcess:
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
save_transform = input_data.get('save_transform', True) save_transform = input_data.get('save_transform', True)
# Get last timestamp for incremental processing prefix = data.cleanup_prefix()
last_timestamp = await workflow.execute_local_activity_method(
Activities.get_last_timestamp, try:
{**metadata, 'data': data}, await self._run_prediction_pipeline(
retry_policy=retry_policy, input_data,
start_to_close_timeout=timedelta(minutes=1), metadata,
) data,
model_id,
model_name,
model_config,
save_transform,
)
finally:
if self.cleanup_prefixes:
await workflow.execute_activity_method(
Activities.cleanup_minio_objects_expired,
{**metadata, 'prefix': prefix},
retry_policy=retry_policy,
start_to_close_timeout=timedelta(minutes=5),
)
async def _run_prediction_pipeline(
self,
input_data: dict[str, Any],
metadata: dict[str, Any],
data: Any,
model_id: Any,
model_name: str,
model_config: dict[str, Any],
save_transform: bool
) -> None:
last_timestamp = data.last_timestamp
# Apply input data quality gates # Apply input data quality gates
gate_input = { gate_input = {
@@ -119,7 +149,12 @@ 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),
) )
@@ -221,7 +256,7 @@ class PredictionProcess:
async def path_flag_handler( async def path_flag_handler(
self, self,
data: dict, data: MinioDataFramePayload,
path_flag: str, path_flag: str,
input_data: dict, input_data: dict,
confidence: int, confidence: int,

View File

@@ -3,7 +3,7 @@ psycopg2-binary
sqlalchemy sqlalchemy
asyncua asyncua
redis redis
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.8.2 git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.10.2
git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.40.7 git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.40.7
prometheus-client prometheus-client
botocore botocore

View File

@@ -38,14 +38,13 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
) )
mock_minio_repository.assert_called_once_with( mock_minio_repository.assert_called_once_with(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=ANY, logger=ANY,
notification_handler=ANY, notification_handler=ANY,
minio_endpoint_url='http://localhost:9000',
minio_access_key='minio',
minio_secret_key='minio123',
minio_region_name='us-east-1',
minio_default_bucket='test',
metrics_controller=ANY, metrics_controller=ANY,
bucket='test',
) )
@@ -96,10 +95,15 @@ metadata = {
@mark.asyncio @mark.asyncio
@patch('laborious.activities.mlflow.DataFrame') @patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.max') @patch('laborious.activities.mlflow.max')
async def test_request_transform_success(mock_max, mock_dataframe, mlflow): async def test_request_transform_success(mock_max, mock_dataframe_from_wire, mlflow):
mock_max.return_value = '2024-01-02' mock_max.return_value = '2024-01-02'
data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock
# Mock input data # Mock input data
input_data = { input_data = {
**metadata, **metadata,
@@ -149,37 +153,41 @@ async def test_request_transform_success(mock_max, mock_dataframe, mlflow):
expected_response = {'prediction': [0.5, 0.6], 'timestamp': ['2024-01-01', '2024-01-02']} expected_response = {'prediction': [0.5, 0.6], 'timestamp': ['2024-01-01', '2024-01-02']}
mlflow.model_monitoring_repository.transform.return_value = expected_response mlflow.model_monitoring_repository.transform.return_value = expected_response
mock_dataframe.return_value.sort_values.return_value = mock_dataframe.return_value data_mock.sort_values.return_value = data_mock
mock_dataframe.return_value.drop_duplicates.return_value = mock_dataframe.return_value data_mock.drop_duplicates.return_value = data_mock
data_mock.pivot.return_value = data_mock
# Call the method # 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 # Verify the data was correctly transformed
mock_dataframe.assert_called_once_with(input_data['data']) data_mock.pivot.assert_called_once_with(
mock_dataframe.return_value.pivot.assert_called_once_with(
index='timestamp', columns='variable', values='value' index='timestamp', columns='variable', values='value'
) )
mock_dataframe = mock_dataframe.return_value.pivot.return_value data_mock.fillna.assert_called_once_with(np.nan, inplace=True)
mock_dataframe.fillna.assert_called_once_with(np.nan, inplace=True)
# mock_dataframe.reset_index.assert_called_once() # mock_dataframe.reset_index.assert_called_once()
mock_dataframe.columns.name = None data_mock.columns.name = None
# Verify the response # Verify the response
assert response_data == expected_response assert response_data == expected_response
# Verify the repository was called with correct arguments # 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', mock_dataframe, {}, metadata['metadata'] 'test_model', data_mock, {}, metadata['metadata']
) )
@mark.asyncio @mark.asyncio
@patch('laborious.activities.mlflow.DataFrame') @patch(
'laborious.activities.mlflow.MinioDataFramePayload.dataframe_from_wire',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime') @patch('laborious.activities.mlflow.to_datetime')
@patch('laborious.activities.mlflow.max') @patch('laborious.activities.mlflow.max')
async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe, mlflow): async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe_from_wire, mlflow):
mock_max.return_value = '2024-01-02' mock_max.return_value = '2024-01-02'
data_mock = MagicMock()
mock_dataframe_from_wire.return_value = data_mock
# Mock input data # Mock input data
input_data = { input_data = {
**metadata, **metadata,
@@ -203,22 +211,21 @@ async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe, mlflo
# Call the method # Call the method
response_data = await mlflow.request_predict(input_data) response_data = await mlflow.request_predict(input_data)
mock_dataframe.assert_called_once_with(input_data['data']) data_mock.replace.assert_called_once_with(np.nan, None, inplace=True)
mock_dataframe.return_value.replace.assert_called_once_with(np.nan, None, inplace=True) data_mock.__setitem__.assert_any_call(
mock_dataframe.return_value.__setitem__.assert_any_call(
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value 'timestamp', mock_to_datetime.return_value.dt.strftime.return_value
) )
mock_dataframe.return_value.__setitem__.assert_any_call( data_mock.__setitem__.assert_any_call(
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value 'timestamp', mock_to_datetime.return_value.dt.strftime.return_value
) )
mock_to_datetime.assert_called_once_with( mock_to_datetime.assert_called_once_with(
mock_dataframe.return_value.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ 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.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT)
mock_to_datetime.assert_called_once_with( mock_to_datetime.assert_called_once_with(
mock_dataframe.return_value.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ 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.return_value.dt.strftime.assert_called_once_with(DATETIME_FORMAT)
@@ -227,20 +234,22 @@ async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe, mlflo
# Verify the repository was called with correct arguments # Verify the repository was called with correct arguments
mlflow.model_monitoring_repository.predict.assert_called_once_with( mlflow.model_monitoring_repository.predict.assert_called_once_with(
'test_model', mock_dataframe.return_value, {}, metadata['metadata'] 'test_model', data_mock, {}, metadata['metadata']
) )
@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_success_retrain(mock_to_datetime, mlflow): async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mock_read_parquet, 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.get_parquet_as_dataframe.return_value = MagicMock() mlflow.minio_repository.download_file.return_value = b'parquet-bytes'
mock_read_parquet.return_value = MagicMock()
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
@@ -255,7 +264,7 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
} }
) )
raw_data = mlflow.minio_repository.get_parquet_as_dataframe.return_value 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
@@ -308,15 +317,47 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
@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_data_fail_retrain(mock_to_datetime, mlflow): async def test_retrain_model_success_with_payload_data(
mock_to_datetime, mock_dataframe_from_wire, mlflow
):
mlflow.model_monitoring_repository.retrain_model.return_value = {
'success': True,
'experiment': 'test_experiment',
'message': 'Model retrained successfully.',
}
mock_dataframe_from_wire.return_value = MagicMock()
response = await mlflow.retrain_model(
{
**metadata,
'data': {'data': {'a': [1]}},
'model_name': 'test_model',
'model_config': {
'target': 'target',
'transform_flavor': 'sklearn',
'predict_flavor': 'pyfunc',
},
}
)
assert response['success'] is True
mlflow.minio_repository.download_file.assert_not_called()
@mark.asyncio
@patch('laborious.activities.mlflow.read_parquet')
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mock_read_parquet, mlflow):
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.get_parquet_as_dataframe.return_value = MagicMock( mlflow.minio_repository.download_file.return_value = b'parquet-bytes'
mock_read_parquet.return_value = MagicMock(
columns=['variable', 'timestamp', 'value', 'created_at'] columns=['variable', 'timestamp', 'value', 'created_at']
) )
@@ -333,7 +374,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
} }
) )
raw_data = mlflow.minio_repository.get_parquet_as_dataframe.return_value 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
@@ -398,7 +439,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
@mark.asyncio @mark.asyncio
async def test_retrain_model_data_error(mlflow): async def test_retrain_model_data_error(mlflow):
mlflow.minio_repository.get_parquet_as_dataframe.side_effect = Exception( mlflow.minio_repository.download_file.side_effect = Exception(
'Error loading retrain data' 'Error loading retrain data'
) )

View File

@@ -1,4 +1,5 @@
import datetime import datetime
import os
from unittest.mock import ANY, AsyncMock, MagicMock, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
from pytest import fixture, mark, raises from pytest import fixture, mark, raises
@@ -68,14 +69,13 @@ def test___init___not_hasattr(mock_minio_repository):
assert isinstance(storage, Postgres) assert isinstance(storage, Postgres)
mock_minio_repository.assert_called_once_with( mock_minio_repository.assert_called_once_with(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
minio_endpoint_url='localhost:9000',
minio_access_key='minio',
minio_secret_key='minio123',
minio_region_name='us-east-1',
minio_default_bucket='test',
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
bucket='test',
) )
@@ -106,14 +106,13 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
) )
mock_minio_repository.assert_called_once_with( mock_minio_repository.assert_called_once_with(
endpoint='localhost:9000',
access_key='minio',
secret_key='minio123',
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
minio_endpoint_url='localhost:9000',
minio_access_key='minio',
minio_secret_key='minio123',
minio_region_name='us-east-1',
minio_default_bucket='test',
metrics_controller=metrics_controller, metrics_controller=metrics_controller,
bucket='test',
) )
@@ -169,30 +168,38 @@ async def test_query_to_minio_success(now, dataframe, storage):
data = [{'a': 1}, {'a': 2}, {'a': 3}] data = [{'a': 1}, {'a': 2}, {'a': 3}]
storage.load_custom_query = AsyncMock(return_value=data) storage.load_custom_query = AsyncMock(return_value=data)
now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0) now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0)
storage.minio_repository.store_dataframe_as_parquet = AsyncMock() storage.minio_repository.upload_file = AsyncMock(
storage.minio_repository.minio_bucket = 'test' return_value={
'minio_object_name': 'sientia/streamlit-connectors/training_datasets/test_model/test_2024-01-01_00-00-00.parquet'
}
)
storage.minio_repository.bucket = 'test'
result = await storage.query_to_minio({'object_prefix': 'test', **metadata}) result = await storage.query_to_minio({'object_prefix': 'test', **metadata})
dataframe.assert_called_once_with(data) dataframe.assert_called_once_with(data)
storage.minio_repository.store_dataframe_as_parquet.assert_called_once_with( storage.minio_repository.upload_file.assert_called_once_with(
dataframe=dataframe.return_value, file_bytes=ANY,
uri='s3://test/test_2024-01-01_00-00-00.parquet', relative_key='training_datasets/test_model/test_2024-01-01_00-00-00.parquet',
object_name='test_2024-01-01_00-00-00.parquet',
metadata=metadata['metadata'], metadata=metadata['metadata'],
) )
assert result['success'] is True assert result['success'] is True
assert result['object_key'] == 'test_2024-01-01_00-00-00.parquet' assert (
assert result['uri'] == 's3://test/test_2024-01-01_00-00-00.parquet' result['object_key']
== 'sientia/streamlit-connectors/training_datasets/test_model/test_2024-01-01_00-00-00.parquet'
)
assert (
result['uri']
== 's3://test/sientia/streamlit-connectors/training_datasets/test_model/test_2024-01-01_00-00-00.parquet'
)
@mark.asyncio @mark.asyncio
async def test_query_to_minio_error(storage): async def test_query_to_minio_error(storage):
storage.send_notification = MagicMock() storage.send_notification = MagicMock()
storage.send_notification_async = AsyncMock() storage.send_notification_async = AsyncMock()
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
storage.load_custom_query = AsyncMock(side_effect=Exception('test')) storage.load_custom_query = AsyncMock(side_effect=Exception('test'))
result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'}) result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'})
@@ -222,3 +229,80 @@ def test___del__(storage):
storage.__del__() storage.__del__()
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
async def test_load_query_with_minio_offload_no_rows(storage):
storage.load_custom_query = AsyncMock(return_value=None)
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result['success'] is False
@mark.asyncio
async def test_load_query_with_minio_offload_inline(storage):
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'my-model', 'key_prefix': 'predictions/s'}
)
assert result.get('success') is True
assert 'data' in result
assert result.get('object_key') is None
@mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.MinioDataFramePayload.estimate_size_bytes')
async def test_load_query_with_minio_offload_minio(mock_estimate, storage):
mock_estimate.return_value = 10**9
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
storage.minio_repository.upload_file = AsyncMock(
return_value={
'minio_object_name': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-15_12-30-45.parquet'
}
)
storage.minio_repository.bucket = 'test'
fixed = datetime.datetime(2024, 1, 15, 12, 30, 45)
with patch('laborious.utils.models.minio_dataframe_payload.now', return_value=fixed):
result = await storage.load_query_with_minio_offload(
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
)
assert result.get('success') is True
assert result.get('data') is None
assert (
result['object_key']
== 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-15_12-30-45.parquet'
)
storage.minio_repository.upload_file.assert_called_once()
@mark.asyncio
@patch.dict(os.environ, {'SIENTIA_MINIO_RETENTION_HOURS': '1'})
@patch('laborious.activities.storage.now')
async def test_cleanup_minio_objects_expired(mock_now, storage):
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
storage.minio_repository.list_objects = AsyncMock(
return_value=[
'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
'sientia/streamlit-connectors/training_datasets/m/m-initial-2025-01-10_12-00-00.parquet',
]
)
storage.minio_repository.delete_file = AsyncMock()
storage.send_notification_async = AsyncMock()
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'prefixes': ['training_datasets/m']}
)
assert result['success'] is True
assert result['deleted_count'] == 1
storage.minio_repository.delete_file.assert_called_once_with(
object_name='sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
metadata=metadata['metadata'],
)

View File

@@ -0,0 +1,54 @@
from datetime import datetime
from pytest import mark
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
def test_parse_object_timestamp_hyphenated_model():
key = 'predictions/sched/my-long-model-initial-2024-06-15_10-30-45.parquet'
ts = MinioDataFramePayload.parse_object_timestamp(key)
assert ts == datetime(2024, 6, 15, 10, 30, 45)
def test_parse_object_timestamp_transform():
key = 'p/m-transform-2024-01-02_03-04-05.parquet'
ts = MinioDataFramePayload.parse_object_timestamp(key)
assert ts == datetime(2024, 1, 2, 3, 4, 5)
def test_parse_object_timestamp_invalid():
assert MinioDataFramePayload.parse_object_timestamp('bad.parquet') is None
def test_is_offloaded_dict_true_false():
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': None}) is True
assert MinioDataFramePayload.is_offloaded_dict({'object_key': 'k', 'data': {}}) is False
assert MinioDataFramePayload.is_offloaded_dict({'data': {}}) is False
def test_cleanup_prefix_from_payload_dict():
p = {
'object_key': 'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-01-01_00-00-00.parquet',
'bucket': 'b',
'data': None,
}
assert MinioDataFramePayload.cleanup_prefix_from_payload_dict(p) == 'training_datasets/m'
def test_cleanup_prefix_from_explicit_object_prefix():
p = {'object_key': 'x.parquet', 'object_prefix': 'my/prefix', 'data': None}
assert MinioDataFramePayload.cleanup_prefix_from_payload_dict(p) == 'my/prefix'
@mark.asyncio
async def test_resolve_dict_if_offloaded_noop():
d = {'success': True, 'data': {'a': [1]}}
out = await MinioDataFramePayload.resolve_dict_if_offloaded(d, None, {})
assert out is d
@mark.asyncio
async def test_dataframe_from_wire_list():
df = await MinioDataFramePayload.dataframe_from_wire([{'a': 1}], None, {})
assert list(df.columns) == ['a']

View File

@@ -1,245 +0,0 @@
from unittest.mock import ANY, AsyncMock, MagicMock, patch
from botocore.utils import ClientError
from pytest import fixture, mark, raises
from laborious import metrics
from laborious.utils.repository.minio_repository import MinioRepository
@patch('laborious.utils.repository.minio_repository.boto3')
@patch('laborious.utils.repository.minio_repository.Config')
def test___init___(mock_config, mock_boto3):
minio_repository = MinioRepository(
minio_endpoint_url='localhost:9000',
minio_access_key='minio',
minio_secret_key='minio123',
minio_region_name='us-east-1',
minio_default_bucket='test',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
)
assert minio_repository.storage_options == {
'key': 'minio',
'secret': 'minio123',
'client_kwargs': {'endpoint_url': 'localhost:9000'},
}
assert minio_repository.minio_bucket == 'test'
assert minio_repository.minio_endpoint_url == 'localhost:9000'
assert minio_repository.minio_region_name == 'us-east-1'
mock_config.assert_called_once_with(
signature_version='s3v4',
s3={'addressing_style': 'path'},
retries={'max_attempts': 5, 'mode': 'standard'},
connect_timeout=5,
read_timeout=120,
)
mock_boto3.client.assert_called_once_with(
's3',
endpoint_url='localhost:9000',
aws_access_key_id='minio',
aws_secret_access_key='minio123',
region_name='us-east-1',
config=mock_config.return_value,
)
@fixture
@patch('laborious.utils.repository.minio_repository.Config')
@patch('laborious.utils.repository.minio_repository.boto3')
def minio_repository(mock_boto3, mock_config):
minio_repository = MinioRepository(
minio_endpoint_url='localhost:9000',
minio_access_key='minio',
minio_secret_key='minio123',
minio_region_name='us-east-1',
minio_default_bucket='test',
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
)
minio_repository.emit_metric = AsyncMock()
minio_repository.observe_lag = AsyncMock()
minio_repository.send_notification = MagicMock()
minio_repository.send_notification_async = AsyncMock()
return minio_repository
def test_close(minio_repository):
minio_repository.close()
minio_repository.s3_client.close.assert_called_once()
@mark.asyncio
async def test_create_bucket_success(minio_repository):
await minio_repository.create_bucket({})
minio_repository.s3_client.create_bucket.assert_called_once_with(Bucket='test')
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_WRITE_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_WRITE_COUNT, tags=ANY
)
@mark.asyncio
async def test_create_bucket_error(minio_repository):
minio_repository.s3_client.create_bucket.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.create_bucket({})
minio_repository.s3_client.create_bucket.assert_called_once_with(Bucket='test')
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=ANY
)
minio_repository.observe_lag.assert_not_called()
@mark.asyncio
async def test_ensure_bucket_exists_bucket_exists(minio_repository):
assert await minio_repository.ensure_bucket_exists({}) is None
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_READ_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_COUNT, tags=ANY
)
@mark.asyncio
async def test_ensure_bucket_exists_bucket_not_exists_create_success(minio_repository):
minio_repository.s3_client.head_bucket.side_effect = ClientError(
error_response={'Error': {'Code': '404'}}, operation_name='head_bucket'
)
minio_repository.create_bucket = AsyncMock()
assert await minio_repository.ensure_bucket_exists({}) is None
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
minio_repository.create_bucket.assert_called_once_with({})
minio_repository.observe_lag.assert_not_called()
minio_repository.emit_metric.assert_not_called()
@mark.asyncio
async def test_ensure_bucket_exists_bucket_not_exists_create_error(minio_repository):
minio_repository.s3_client.head_bucket.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.ensure_bucket_exists({})
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=ANY
)
minio_repository.observe_lag.assert_not_called()
@mark.asyncio
@patch('laborious.utils.repository.minio_repository.BytesIO')
async def test_store_dataframe_as_parquet_success(mock_bytesio, minio_repository):
input_data = MagicMock()
minio_repository.ensure_bucket_exists = AsyncMock()
await minio_repository.store_dataframe_as_parquet(
dataframe=input_data, uri='s3://test/test.parquet', object_name='test.parquet', metadata={}
)
minio_repository.ensure_bucket_exists.assert_called_once_with({})
mock_bytesio.assert_called_once()
input_data.to_parquet.assert_called_once_with(
mock_bytesio.return_value, engine='pyarrow', index=True
)
mock_bytesio.return_value.seek.assert_called_once_with(0)
minio_repository.s3_client.put_object.assert_called_once_with(
Bucket='test', Key='test.parquet', Body=mock_bytesio.return_value.getvalue.return_value
)
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_WRITE_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_WRITE_COUNT, tags=ANY
)
@mark.asyncio
@patch('laborious.utils.repository.minio_repository.BytesIO')
async def test_store_dataframe_as_parquet_error(mock_bytesio, minio_repository):
input_data = MagicMock()
minio_repository.ensure_bucket_exists = AsyncMock()
minio_repository.s3_client.put_object.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.store_dataframe_as_parquet(
dataframe=input_data,
uri='s3://test/test.parquet',
object_name='test.parquet',
metadata={},
)
minio_repository.ensure_bucket_exists.assert_called_once_with({})
mock_bytesio.assert_called_once()
input_data.to_parquet.assert_called_once_with(
mock_bytesio.return_value, engine='pyarrow', index=True
)
mock_bytesio.return_value.seek.assert_called_once_with(0)
minio_repository.s3_client.put_object.assert_called_once_with(
Bucket='test', Key='test.parquet', Body=mock_bytesio.return_value.getvalue.return_value
)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=ANY
)
@mark.asyncio
@patch('laborious.utils.repository.minio_repository.BytesIO')
@patch('laborious.utils.repository.minio_repository.read_parquet')
async def test_get_parquet_as_dataframe_success(mock_read_parquet, mock_bytesio, minio_repository):
input_data = {'Body': MagicMock(read=MagicMock(return_value=b'test'))}
minio_repository.s3_client.get_object.return_value = input_data
output = await minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
minio_repository.s3_client.get_object.assert_called_once_with(Bucket='test', Key='test.parquet')
mock_bytesio.assert_called_once_with(input_data['Body'].read.return_value)
mock_read_parquet.assert_called_once_with(mock_bytesio.return_value)
assert output == mock_read_parquet.return_value
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_READ_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_COUNT, tags=ANY
)
@mark.asyncio
@patch('laborious.utils.repository.minio_repository.BytesIO')
@patch('laborious.utils.repository.minio_repository.read_parquet')
async def test_get_parquet_as_dataframe_error(mock_read_parquet, mock_bytesio, minio_repository):
minio_repository.s3_client.get_object.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
minio_repository.s3_client.get_object.assert_called_once_with(
Bucket='test', Key='test.parquet'
)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_ERROR_COUNT, tags=ANY
)
minio_repository.observe_lag.assert_not_called()

View File

@@ -17,6 +17,7 @@ metadata = {
'model_name': 'test_model', 'model_name': 'test_model',
'workflow_name': 'test_workflow', 'workflow_name': 'test_workflow',
'schema_name': 'test_schedule', 'schema_name': 'test_schedule',
'schedule_name': 'test_schedule',
}, },
} }
@@ -101,6 +102,7 @@ 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,
@@ -323,6 +325,7 @@ 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,
@@ -423,6 +426,7 @@ 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,
@@ -540,6 +544,7 @@ 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,

View File

@@ -41,7 +41,7 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
workflow_mock.execute_activity_method = AsyncMock( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'success': True, 'object_key': 'test_object_key'}, {'data': {'a': [1]}, 'success': True},
{'success': True, 'experiment': 'test_experiment'}, {'success': True, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -58,13 +58,12 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
workflow_mock.execute_activity_method.assert_has_calls( workflow_mock.execute_activity_method.assert_has_calls(
[ [
call( call(
Activities.query_to_minio, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'query': input_data['query'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'object_prefix': f'retrain_datasets/{input_data["model_name"]}/data',
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -78,7 +77,7 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
Activities.retrain_model, Activities.retrain_model,
{ {
**metadata, **metadata,
'object_key': 'test_object_key', 'data': {'data': {'a': [1]}, 'success': True},
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'model_config': input_data['model_config'], 'model_config': input_data['model_config'],
}, },
@@ -163,7 +162,7 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
workflow_mock.execute_activity_method = AsyncMock( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'success': False, 'object_key': 'test_object_key'}, {'success': False, 'message': 'No data returned from query'},
{'success': True, 'experiment': 'test_experiment'}, {'success': True, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -178,13 +177,12 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
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(
Activities.query_to_minio, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'query': input_data['query'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'object_prefix': f'retrain_datasets/{input_data["model_name"]}/data',
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -213,7 +211,7 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
workflow_mock.execute_activity_method = AsyncMock( workflow_mock.execute_activity_method = AsyncMock(
side_effect=[ side_effect=[
{'success': True, 'object_key': 'test_object_key'}, {'data': {'a': [1]}, 'success': True},
{'success': False, 'experiment': 'test_experiment'}, {'success': False, 'experiment': 'test_experiment'},
{ {
'success': True, 'success': True,
@@ -230,13 +228,12 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
workflow_mock.execute_activity_method.assert_has_calls( workflow_mock.execute_activity_method.assert_has_calls(
[ [
call( call(
Activities.query_to_minio, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'query': input_data['query'], 'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []), 'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'], 'model_name': input_data['model_name'],
'object_prefix': f'retrain_datasets/{input_data["model_name"]}/data',
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -250,7 +247,7 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
Activities.retrain_model, Activities.retrain_model,
{ {
**metadata, **metadata,
'object_key': 'test_object_key', 'data': {'data': {'a': [1]}, 'success': True},
'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

@@ -24,7 +24,10 @@ metadata = {
@mark.asyncio @mark.asyncio
@patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock) @patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock)
async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch): async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch):
workflow_mock.execute_local_activity_method.return_value = {'data': 'test_data'} workflow_mock.execute_activity_method.return_value = {
'success': True,
'data': {'col': ['test_data']},
}
input_data = { input_data = {
'schedule_name': 'test_schedule', 'schedule_name': 'test_schedule',
'model_name': 'test_model', 'model_name': 'test_model',
@@ -42,14 +45,16 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
await predictions_batch.run(input_data) await predictions_batch.run(input_data)
workflow_mock.execute_local_activity_method.assert_has_calls( workflow_mock.execute_activity_method.assert_has_calls(
[ [
call( call(
Activities.load_custom_query, Activities.load_query_with_minio_offload,
{ {
**metadata, **metadata,
'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'],
'key_prefix': f"predictions/{input_data['schedule_name']}",
}, },
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
@@ -58,7 +63,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
) )
prediction_input = { prediction_input = {
'metadata': metadata, 'metadata': metadata,
'data': {'data': 'test_data'}, 'data': {'success': True, 'data': {'col': ['test_data']}},
'schema': input_data['schema'], 'schema': input_data['schema'],
'table_name': input_data['table_name'], 'table_name': input_data['table_name'],
'transform_table_name': input_data['transform_table_name'], 'transform_table_name': input_data['transform_table_name'],