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:
70
README.md
70
README.md
@@ -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:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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]):
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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')),
|
||||||
}
|
}
|
||||||
|
|||||||
0
laborious/utils/models/__init__.py
Normal file
0
laborious/utils/models/__init__.py
Normal file
230
laborious/utils/models/minio_dataframe_payload.py
Normal file
230
laborious/utils/models/minio_dataframe_payload.py
Normal 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
|
||||||
26
laborious/utils/repository/minio_manager.py
Normal file
26
laborious/utils/repository/minio_manager.py
Normal 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)
|
||||||
@@ -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)
|
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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'],
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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'],
|
||||||
|
)
|
||||||
|
|||||||
54
tests/laborious/utils/models/test_minio_dataframe_payload.py
Normal file
54
tests/laborious/utils/models/test_minio_dataframe_payload.py
Normal 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']
|
||||||
@@ -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()
|
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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'],
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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'],
|
||||||
|
|||||||
Reference in New Issue
Block a user