Merge pull request #29 from Aignosi/feature/SIENTIAPDE-1325-adicionar-metricas-especificas-de-operacoes-externas
SIENTIAPDE-1325: Enhance OPC metrics and monitoring with server name and URL
This commit is contained in:
1
.github/workflows/release.yml
vendored
1
.github/workflows/release.yml
vendored
@@ -8,6 +8,7 @@ on:
|
||||
|
||||
jobs:
|
||||
release:
|
||||
if: github.event.pull_request.merged == true
|
||||
uses: Aignosi/github_workflow_templates/.github/workflows/dataops-module-release.yml@main
|
||||
permissions: write-all
|
||||
with:
|
||||
|
||||
10
README.md
10
README.md
@@ -607,18 +607,18 @@ The Laborious system exposes comprehensive Prometheus metrics for operational vi
|
||||
|
||||
### Prediction Operation Metrics
|
||||
- `laborious_predictions_written_count`: Counter for successful prediction exports
|
||||
- Labels: `pod_id`, `model_name`, `pipeline_name`
|
||||
- Labels: `pod_id`, `model_name`, `workflow_name`
|
||||
- `laborious_prediction_confidence_monitor`: Gauge for current prediction confidence levels
|
||||
- Labels: `pod_id`, `model_name`, `pipeline_name`
|
||||
- Labels: `pod_id`, `model_name`, `workflow_name`
|
||||
- `laborious_prediction_response_time_monitor`: Histogram for prediction response times
|
||||
- Labels: `pod_id`, `model_name`, `pipeline_name`
|
||||
- Labels: `pod_id`, `model_name`, `workflow_name`
|
||||
- Buckets: [0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0]
|
||||
|
||||
### OPC Export Metrics
|
||||
- `laborious_prediction_opc_writing_count`: Counter for OPC server write operations
|
||||
- Labels: `pod_id`, `model_name`, `pipeline_name`, `opc_server_id`
|
||||
- Labels: `pod_id`, `model_name`, `workflow_name`, `opc_server_id`
|
||||
- `laborious_prediction_opc_writing_response_time_monitor`: Histogram for OPC write response times
|
||||
- Labels: `pod_id`, `model_name`, `pipeline_name`, `opc_server_id`
|
||||
- Labels: `pod_id`, `model_name`, `workflow_name`, `opc_server_id`
|
||||
- Buckets: [0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0]
|
||||
|
||||
### Data Quality Metrics
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from temporalio import workflow
|
||||
|
||||
with workflow.unsafe.imports_passed_through():
|
||||
@@ -62,6 +63,8 @@ class Activities(Storage, MLFlow, Gates, OPC):
|
||||
Raises:
|
||||
Exception: If any parent class initialization fails
|
||||
"""
|
||||
metrics_controller = MetricsController(logger=logger)
|
||||
|
||||
# Initialize parent classes
|
||||
Storage.__init__(
|
||||
self,
|
||||
@@ -75,6 +78,7 @@ class Activities(Storage, MLFlow, Gates, OPC):
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
MLFlow.__init__(
|
||||
@@ -86,12 +90,22 @@ class Activities(Storage, MLFlow, Gates, OPC):
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
Gates.__init__(self, logger=logger, notification_handler=notification_handler)
|
||||
Gates.__init__(
|
||||
self,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
OPC.__init__(
|
||||
self, opc_servers=opc_config, logger=logger, notification_handler=notification_handler
|
||||
self,
|
||||
opc_servers=opc_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
async def shutdown(self):
|
||||
@@ -107,4 +121,6 @@ class Activities(Storage, MLFlow, Gates, OPC):
|
||||
proper resource cleanup and prevent resource leaks.
|
||||
"""
|
||||
Storage.close(self)
|
||||
await OPC.shutdown(self)
|
||||
MLFlow.close(self)
|
||||
Gates.close(self)
|
||||
await OPC.close(self)
|
||||
|
||||
@@ -10,7 +10,8 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.temporal.activities.base import BaseActivity
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ, now
|
||||
|
||||
from laborious import metrics
|
||||
@@ -62,7 +63,7 @@ mlflow_content_path_confidence: Mapping[str, int] = {
|
||||
}
|
||||
|
||||
|
||||
class Gates(BaseActivity):
|
||||
class Gates(SientiaMonitoring):
|
||||
"""
|
||||
Data quality gates and filtering activities for the Laborious system.
|
||||
|
||||
@@ -82,7 +83,12 @@ class Gates(BaseActivity):
|
||||
mlflow_content_filter_functions (dict): Mapping of MLFlow content filter names to functions
|
||||
"""
|
||||
|
||||
def __init__(self, logger: Logger, notification_handler: NotificationHandler):
|
||||
def __init__(
|
||||
self,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""
|
||||
Initialize data quality gates with logging and notification capabilities.
|
||||
|
||||
@@ -93,7 +99,16 @@ class Gates(BaseActivity):
|
||||
Raises:
|
||||
Exception: If BaseActivity initialization fails
|
||||
"""
|
||||
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
|
||||
def close(self) -> None:
|
||||
"""
|
||||
Close the gates activity and clean up resources.
|
||||
"""
|
||||
SientiaMonitoring.shutdown(self)
|
||||
|
||||
def __del__(self):
|
||||
self.close()
|
||||
|
||||
@activity.defn(name='input_gate')
|
||||
async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]:
|
||||
@@ -153,7 +168,7 @@ class Gates(BaseActivity):
|
||||
filter_output.append(config['policy'])
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'INTPUT_GATE_ERROR__{fil}',
|
||||
message=f'Error in filter {fil}:{config}: \n {e}',
|
||||
@@ -225,7 +240,7 @@ class Gates(BaseActivity):
|
||||
if mlflow_response_filter_functions[fil](data, config):
|
||||
filter_output.append(config['policy'])
|
||||
comments.append(data['content']['message'])
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}',
|
||||
message=data['content']['message'],
|
||||
@@ -235,7 +250,7 @@ class Gates(BaseActivity):
|
||||
)
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}',
|
||||
message=f'Error in filter {fil}:{config}: \n {e}',
|
||||
@@ -305,7 +320,7 @@ class Gates(BaseActivity):
|
||||
try:
|
||||
if mlflow_content_filter_functions[fil](data, config):
|
||||
filter_output.append(config['policy'])
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}',
|
||||
message=f'Data not passed the content filter {fil}:{config}',
|
||||
@@ -315,7 +330,7 @@ class Gates(BaseActivity):
|
||||
)
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}',
|
||||
message=f'Error in filter {fil}:{config}: \n {e}',
|
||||
@@ -608,39 +623,62 @@ class Gates(BaseActivity):
|
||||
|
||||
self.info(f'Writing metrics for model {metadata["model_name"]}', metadata)
|
||||
|
||||
metrics.PREDICTIONS_WRITTEN_COUNT.labels(
|
||||
pod_id=self.pod_id,
|
||||
model_name=metadata['model_name'],
|
||||
pipeline_name=metadata['workflow_name'],
|
||||
).inc()
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.PREDICTIONS_WRITTEN_COUNT,
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'model_name': metadata['model_name'],
|
||||
'workflow_name': metadata['workflow_name'],
|
||||
},
|
||||
)
|
||||
|
||||
metrics.PREDICTION_CONFIDENCE_MONITOR.labels(
|
||||
pod_id=self.pod_id,
|
||||
model_name=metadata['model_name'],
|
||||
pipeline_name=metadata['workflow_name'],
|
||||
).set(prediction_confidence)
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.PREDICTION_CONFIDENCE_MONITOR,
|
||||
method='set',
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'model_name': metadata['model_name'],
|
||||
'workflow_name': metadata['workflow_name'],
|
||||
},
|
||||
value=prediction_confidence,
|
||||
)
|
||||
|
||||
metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels(
|
||||
pod_id=self.pod_id,
|
||||
model_name=metadata['model_name'],
|
||||
pipeline_name=metadata['workflow_name'],
|
||||
).observe(response_time)
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.PREDICTION_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'model_name': metadata['model_name'],
|
||||
'workflow_name': metadata['workflow_name'],
|
||||
},
|
||||
value=response_time,
|
||||
)
|
||||
|
||||
for server_id, tags in opc_metrics.items():
|
||||
for tag, response_time in tags.items():
|
||||
metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels(
|
||||
pod_id=self.pod_id,
|
||||
model_name=metadata['model_name'],
|
||||
pipeline_name=metadata['workflow_name'],
|
||||
opc_server_id=server_id,
|
||||
tag=tag,
|
||||
).observe(response_time)
|
||||
metrics.PREDICTION_OPC_WRITING_COUNT.labels(
|
||||
pod_id=self.pod_id,
|
||||
model_name=metadata['model_name'],
|
||||
pipeline_name=metadata['workflow_name'],
|
||||
opc_server_id=server_id,
|
||||
tag=tag,
|
||||
).inc()
|
||||
if response_time is not None:
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'model_name': metadata['model_name'],
|
||||
'workflow_name': metadata['workflow_name'],
|
||||
'opc_server_id': server_id,
|
||||
'tag': tag,
|
||||
},
|
||||
value=response_time,
|
||||
)
|
||||
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'model_name': metadata['model_name'],
|
||||
'workflow_name': metadata['workflow_name'],
|
||||
'opc_server_id': server_id,
|
||||
'tag': tag,
|
||||
},
|
||||
)
|
||||
|
||||
self.info(f'Metrics written for model {metadata["model_name"]}', metadata)
|
||||
|
||||
@@ -10,7 +10,8 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.temporal.activities.base import BaseActivity
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.temporal.constants import (
|
||||
DATETIME_FORMAT,
|
||||
DATETIME_FORMAT_MS_WITH_TZ,
|
||||
@@ -22,7 +23,7 @@ with workflow.unsafe.imports_passed_through():
|
||||
from laborious.utils.repository.model_repository import MLFlowRepository
|
||||
|
||||
|
||||
class MLFlow(BaseActivity):
|
||||
class MLFlow(SientiaMonitoring):
|
||||
"""
|
||||
MLFlow integration activities for model inference operations.
|
||||
|
||||
@@ -50,6 +51,7 @@ class MLFlow(BaseActivity):
|
||||
mlflow_password: str,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""
|
||||
Initialize MLFlow activities with server configuration.
|
||||
@@ -65,14 +67,19 @@ class MLFlow(BaseActivity):
|
||||
Raises:
|
||||
Exception: If MLFlowRepository initialization fails
|
||||
"""
|
||||
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
self.mlflow_host = mlflow_host
|
||||
self.mlflow_port = mlflow_port
|
||||
self.mlflow_username = mlflow_username
|
||||
self.mlflow_password = mlflow_password
|
||||
|
||||
self.model_monitoring_repository = MLFlowRepository(
|
||||
f'{mlflow_host}:{mlflow_port}', mlflow_username, mlflow_password, logger
|
||||
f'{mlflow_host}:{mlflow_port}',
|
||||
mlflow_username,
|
||||
mlflow_password,
|
||||
logger,
|
||||
notification_handler,
|
||||
metrics_controller,
|
||||
)
|
||||
|
||||
if not hasattr(self, 'minio_repository'):
|
||||
@@ -87,8 +94,18 @@ class MLFlow(BaseActivity):
|
||||
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:
|
||||
"""
|
||||
Close the MLFlow activity and clean up resources.
|
||||
"""
|
||||
SientiaMonitoring.shutdown(self)
|
||||
|
||||
def __del__(self):
|
||||
self.close()
|
||||
|
||||
@activity.defn(name='request_transform')
|
||||
async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -143,7 +160,7 @@ class MLFlow(BaseActivity):
|
||||
self.debug(data.head(5).to_string(), metadata)
|
||||
|
||||
# Request transformation from MLFlow model
|
||||
response_data = self.model_monitoring_repository.transform(
|
||||
response_data = await self.model_monitoring_repository.transform(
|
||||
model_name, data, model_config, metadata
|
||||
)
|
||||
|
||||
@@ -208,7 +225,7 @@ class MLFlow(BaseActivity):
|
||||
).dt.strftime(DATETIME_FORMAT)
|
||||
|
||||
# Request prediction from MLFlow model
|
||||
response_data = self.model_monitoring_repository.predict(
|
||||
response_data = await self.model_monitoring_repository.predict(
|
||||
model_name, data, model_config, metadata
|
||||
)
|
||||
|
||||
@@ -263,12 +280,12 @@ class MLFlow(BaseActivity):
|
||||
self.info(f'Loading retrain data from Key: {object_key}', metadata)
|
||||
|
||||
try:
|
||||
data = self.minio_repository.get_parquet_as_dataframe(
|
||||
data = await self.minio_repository.get_parquet_as_dataframe(
|
||||
object_key=object_key, metadata=metadata
|
||||
)
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id='ERROR_LOADING_RETRAIN_DATA',
|
||||
message=f'Error loading retrain data: {e}',
|
||||
@@ -295,9 +312,12 @@ class MLFlow(BaseActivity):
|
||||
self.debug(f'Timestamp: {timestamp}', metadata)
|
||||
|
||||
# Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair
|
||||
if 'created_at' in data.columns:
|
||||
data = data.sort_values('created_at', ascending=False).drop_duplicates(
|
||||
subset=['variable', 'timestamp'], keep='first'
|
||||
)
|
||||
else:
|
||||
data = data.drop_duplicates(subset=['variable', 'timestamp'], keep='first')
|
||||
|
||||
data.drop(columns=['model_id'], inplace=True, errors='ignore')
|
||||
data.drop(columns=['created_at'], inplace=True, errors='ignore')
|
||||
@@ -316,13 +336,13 @@ class MLFlow(BaseActivity):
|
||||
|
||||
data.columns.name = None
|
||||
|
||||
retrain_output = self.model_monitoring_repository.retrain_model(
|
||||
retrain_output = await self.model_monitoring_repository.retrain_model(
|
||||
data=data, model_name=model_name, model_config=model_config, metadata=metadata
|
||||
)
|
||||
|
||||
if not retrain_output['success']:
|
||||
trace = retrain_output['traceback']
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id='RETRAIN_MODEL_ERROR',
|
||||
message=f'Error retraining model {model_name}: {retrain_output["message"]}',
|
||||
@@ -379,7 +399,7 @@ class MLFlow(BaseActivity):
|
||||
)
|
||||
|
||||
try:
|
||||
response = self.model_monitoring_repository.update_production_model(
|
||||
response = await self.model_monitoring_repository.update_production_model(
|
||||
experiment=experiment, model_name=model_name, metadata=metadata
|
||||
)
|
||||
|
||||
@@ -388,7 +408,7 @@ class MLFlow(BaseActivity):
|
||||
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
|
||||
message=f'Error updating production model {model_name}: {e}',
|
||||
|
||||
@@ -8,14 +8,15 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.temporal.activities.base import BaseActivity
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
|
||||
from laborious.utils.repository.opc_repository import OpcRepository
|
||||
|
||||
OPC_WRITTING_ERROR_CONFIDENCE = 12
|
||||
|
||||
|
||||
class OPC(BaseActivity):
|
||||
class OPC(SientiaMonitoring):
|
||||
"""
|
||||
OPC server integration activities for real-time data export.
|
||||
|
||||
@@ -39,15 +40,15 @@ class OPC(BaseActivity):
|
||||
opc_servers: dict[str, dict[str, Any]],
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
self.logger = logger
|
||||
self.notification_handler = notification_handler
|
||||
self.opc_servers = opc_servers
|
||||
|
||||
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
|
||||
self.opc_repository: dict[str, OpcRepository] = {}
|
||||
self.opc_servers = opc_servers
|
||||
|
||||
async def init_opc(self):
|
||||
"""
|
||||
@@ -77,6 +78,7 @@ class OPC(BaseActivity):
|
||||
for opc_id, server in self.opc_servers.items():
|
||||
self.opc_repository[opc_id] = OpcRepository(
|
||||
opc_id=server['id'],
|
||||
server_name=server['server_name'],
|
||||
url=server['url'],
|
||||
logger=self.logger,
|
||||
server_uri=server['server_uri'],
|
||||
@@ -85,10 +87,11 @@ class OPC(BaseActivity):
|
||||
server_cert_path=server['server_cert_path'],
|
||||
notification_handler=self.notification_handler,
|
||||
reconnection_interval=server['reconnection_interval'],
|
||||
metrics_controller=self.metrics_controller,
|
||||
)
|
||||
is_connected, error_data = await self.opc_repository[opc_id].connect()
|
||||
if not is_connected:
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata={
|
||||
'model_id': '-',
|
||||
'model_name': '-',
|
||||
@@ -102,7 +105,9 @@ class OPC(BaseActivity):
|
||||
attachment_content=error_data.get('attachment_content', None),
|
||||
)
|
||||
else:
|
||||
self.logger.info(f'OPC server {opc_id} connected successfully.')
|
||||
self.logger.info(
|
||||
f'OPC server {opc_id}:{server["server_name"]} connected successfully.'
|
||||
)
|
||||
|
||||
async def write_data(
|
||||
self,
|
||||
@@ -137,7 +142,7 @@ class OPC(BaseActivity):
|
||||
tag, data, data_type, self.logger, metadata
|
||||
)
|
||||
if not is_success:
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=info_data['notification_id'],
|
||||
message=info_data['message'],
|
||||
@@ -149,7 +154,7 @@ class OPC(BaseActivity):
|
||||
return info_data['response_time']
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id=f'WRITE_OPC_{tag_type.upper()}_ERROR',
|
||||
message=f'Error writing data to OPC server: {e}',
|
||||
@@ -159,7 +164,7 @@ class OPC(BaseActivity):
|
||||
)
|
||||
raise e
|
||||
|
||||
def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool:
|
||||
async def validate_server(self, server_id: str, metadata: dict[str, Any]) -> bool:
|
||||
"""
|
||||
Validate that an OPC server is available and configured for write operations.
|
||||
|
||||
@@ -182,7 +187,7 @@ class OPC(BaseActivity):
|
||||
"""
|
||||
if self.opc_repository.get(server_id) is None:
|
||||
message = f'OPC server {server_id} not found to perform write operation.'
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id='OPC_SERVER_NOT_FOUND',
|
||||
message=message,
|
||||
@@ -299,7 +304,7 @@ class OPC(BaseActivity):
|
||||
metrics: dict[str, dict[str, float | None]] = {}
|
||||
|
||||
for server_id, config in opc_output_config.items():
|
||||
if not self.validate_server(server_id, metadata):
|
||||
if not await self.validate_server(server_id, metadata):
|
||||
success = False
|
||||
continue
|
||||
|
||||
@@ -358,7 +363,7 @@ class OPC(BaseActivity):
|
||||
|
||||
return data.to_dict()
|
||||
|
||||
async def shutdown(self):
|
||||
async def close(self):
|
||||
"""
|
||||
Gracefully shutdown all OPC server connections and cleanup resources.
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.temporal.activities.postgres import Postgres
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now
|
||||
|
||||
@@ -33,6 +34,7 @@ class Storage(Postgres):
|
||||
minio_config: dict[str, Any],
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
super().__init__(
|
||||
host=host,
|
||||
@@ -44,6 +46,7 @@ class Storage(Postgres):
|
||||
max_connections=max_connections,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
if not hasattr(self, 'minio_repository'):
|
||||
@@ -58,6 +61,7 @@ class Storage(Postgres):
|
||||
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,
|
||||
)
|
||||
|
||||
@activity.defn(name='query_to_minio')
|
||||
@@ -95,14 +99,14 @@ class Storage(Postgres):
|
||||
data = pd.DataFrame(data)
|
||||
|
||||
# Write parquet to memory and upload via persistent client
|
||||
self.minio_repository.store_dataframe_as_parquet(
|
||||
await self.minio_repository.store_dataframe_as_parquet(
|
||||
dataframe=data, uri=uri, object_name=object_name, metadata=metadata
|
||||
)
|
||||
|
||||
return {'success': True, 'object_key': object_name, 'uri': uri}
|
||||
except Exception as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=metadata,
|
||||
notification_id='ERROR_STORING_QUERY_TO_MINIO',
|
||||
message=f'Error storing query to MinIO: {e}',
|
||||
|
||||
@@ -19,11 +19,12 @@ Key Metric Categories:
|
||||
Metric Labels:
|
||||
- pod_id: Kubernetes pod identifier for multi-instance deployments
|
||||
- model_name: Name of the ML model being used
|
||||
- pipeline_name: Name of the prediction pipeline
|
||||
- workflow_name: Name of the prediction pipeline
|
||||
- opc_server_id: Identifier for OPC server operations
|
||||
"""
|
||||
|
||||
from prometheus_client import Counter, Gauge, Histogram
|
||||
from sientia_do.observability.metrics import CORE_LABELS as SIENTIA_CORE_LABELS
|
||||
|
||||
# Application health metric
|
||||
APP_UP = Gauge(
|
||||
@@ -33,7 +34,7 @@ APP_UP = Gauge(
|
||||
)
|
||||
|
||||
# Core labels used across multiple metrics
|
||||
CORE_LABELS = ['pod_id', 'model_name', 'pipeline_name']
|
||||
CORE_LABELS = ['pod_id', 'model_name', 'workflow_name']
|
||||
|
||||
# Prediction operation metrics
|
||||
PREDICTIONS_WRITTEN_COUNT = Counter(
|
||||
@@ -49,7 +50,7 @@ PREDICTION_CONFIDENCE_MONITOR = Gauge(
|
||||
CORE_LABELS,
|
||||
)
|
||||
|
||||
# Performance monitoring metrics
|
||||
# Prediction total response time
|
||||
PREDICTION_RESPONSE_TIME_MONITOR = Histogram(
|
||||
'laborious_prediction_response_time_monitor',
|
||||
'Current response time of each prediction',
|
||||
@@ -57,7 +58,49 @@ PREDICTION_RESPONSE_TIME_MONITOR = Histogram(
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
# OPC export metrics
|
||||
# ================== MinIO metrics ==================
|
||||
|
||||
MINIO_READ_LAG = Histogram(
|
||||
'laborious_minio_read_lag',
|
||||
'Lag between the last write to MinIO and the last read from MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
MINIO_WRITE_LAG = Histogram(
|
||||
'laborious_minio_write_lag',
|
||||
'Lag between the last write to MinIO and the last read from MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
MINIO_READ_COUNT = Counter(
|
||||
'laborious_minio_read_count',
|
||||
'Number of reads from MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
)
|
||||
|
||||
MINIO_WRITE_COUNT = Counter(
|
||||
'laborious_minio_write_count',
|
||||
'Number of writes to MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
)
|
||||
|
||||
MINIO_READ_ERROR_COUNT = Counter(
|
||||
'laborious_minio_read_error_count',
|
||||
'Number of errors reading from MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
)
|
||||
|
||||
MINIO_WRITE_ERROR_COUNT = Counter(
|
||||
'laborious_minio_write_error_count',
|
||||
'Number of errors writing to MinIO',
|
||||
[*SIENTIA_CORE_LABELS, 'bucket_name', 'object_name'],
|
||||
)
|
||||
|
||||
# ================== OPC metrics ==================
|
||||
|
||||
|
||||
PREDICTION_OPC_WRITING_COUNT = Counter(
|
||||
'laborious_prediction_opc_writing_count',
|
||||
'Number of predictions written to the OPC server',
|
||||
@@ -70,3 +113,59 @@ PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR = Histogram(
|
||||
[*CORE_LABELS, 'opc_server_id', 'tag'],
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
OPC_CONNECTIONS_TOTAL = Counter(
|
||||
'opc_connections_initiated_total',
|
||||
'Total connection attempts to OPC servers',
|
||||
['pod_id', 'server_name'],
|
||||
)
|
||||
OPC_CONNECTIONS_FAILED = Counter(
|
||||
'opc_connections_failed_total',
|
||||
'Total failed connection attempts to OPC servers',
|
||||
['pod_id', 'server_name'],
|
||||
)
|
||||
OPC_CONNECTION_STATUS = Gauge(
|
||||
'opc_connection_status',
|
||||
'Connection status with the OPC server (1=connected, 0=disconnected)',
|
||||
['pod_id', 'server_name', 'server_url'],
|
||||
)
|
||||
|
||||
# ================== Model metrics ==================
|
||||
|
||||
MODEL_READ_LAG = Histogram(
|
||||
'laborious_model_read_lag',
|
||||
'Lag between the start and read of read operations',
|
||||
SIENTIA_CORE_LABELS,
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
MODEL_WRITE_LAG = Histogram(
|
||||
'laborious_model_write_lag',
|
||||
'Lag between the start and end of write operations',
|
||||
SIENTIA_CORE_LABELS,
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
|
||||
)
|
||||
|
||||
MODEL_READ_COUNT = Counter(
|
||||
'laborious_model_read_count',
|
||||
'Number of reads from the model',
|
||||
SIENTIA_CORE_LABELS,
|
||||
)
|
||||
|
||||
MODEL_WRITE_COUNT = Counter(
|
||||
'laborious_model_write_count',
|
||||
'Number of writes to the model',
|
||||
SIENTIA_CORE_LABELS,
|
||||
)
|
||||
|
||||
MODEL_READ_ERROR_COUNT = Counter(
|
||||
'laborious_model_read_error_count',
|
||||
'Number of errors reading from the model',
|
||||
SIENTIA_CORE_LABELS,
|
||||
)
|
||||
|
||||
MODEL_WRITE_ERROR_COUNT = Counter(
|
||||
'laborious_model_write_error_count',
|
||||
'Number of errors writing to the model',
|
||||
SIENTIA_CORE_LABELS,
|
||||
)
|
||||
|
||||
@@ -88,6 +88,7 @@ def build_opc_config() -> dict[str, Any]:
|
||||
return {
|
||||
getenv('OPC_ID', '1'): {
|
||||
'id': getenv('OPC_ID', '1'),
|
||||
'server_name': getenv('OPC_SERVER_NAME', 'default_server'),
|
||||
'url': getenv('OPC_URL', 'opc.tcp://localhost:4840'),
|
||||
'server_uri': getenv('OPC_SERVER_URI', 'opc.tcp://localhost:4840'),
|
||||
'cert_path': getenv('OPC_CERT_PATH', None),
|
||||
|
||||
@@ -6,6 +6,7 @@ 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
|
||||
|
||||
@@ -15,9 +16,13 @@ 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:
|
||||
class MinioRepository(SientiaMonitoring):
|
||||
"""
|
||||
Repository for interacting with a MinIO (S3-compatible) object storage.
|
||||
|
||||
@@ -43,6 +48,7 @@ class MinioRepository:
|
||||
minio_default_bucket: str,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""Initialize the repository and S3 client.
|
||||
|
||||
@@ -55,6 +61,7 @@ class MinioRepository:
|
||||
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,
|
||||
@@ -85,28 +92,58 @@ class MinioRepository:
|
||||
),
|
||||
)
|
||||
|
||||
self.logger = logger
|
||||
self.notification_handler = notification_handler
|
||||
|
||||
def close(self):
|
||||
"""Close the underlying S3 client."""
|
||||
self.s3_client.close()
|
||||
|
||||
def ensure_bucket_exists(self, metadata: dict[str, Any]) -> None:
|
||||
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.logger.custom_info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
|
||||
self.s3_client.head_bucket(Bucket=self.minio_bucket)
|
||||
except ClientError:
|
||||
self.logger.custom_info(f"Creating bucket '{self.minio_bucket}'", metadata)
|
||||
self.s3_client.create_bucket(Bucket=self.minio_bucket)
|
||||
await self.create_bucket(metadata)
|
||||
|
||||
def store_dataframe_as_parquet(
|
||||
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.
|
||||
@@ -117,18 +154,36 @@ class MinioRepository:
|
||||
object_name (str): Object key (path/key within the bucket).
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
"""
|
||||
self.ensure_bucket_exists(metadata)
|
||||
await self.ensure_bucket_exists(metadata)
|
||||
|
||||
self.logger.custom_info(f'Storing dataframe as parquet in {uri}', metadata)
|
||||
self.info(f'Storing dataframe as parquet in {uri}', metadata)
|
||||
|
||||
buffer = BytesIO()
|
||||
dataframe.to_parquet(buffer, engine='pyarrow', index=True)
|
||||
buffer.seek(0)
|
||||
self.s3_client.put_object(Bucket=self.minio_bucket, Key=object_name, Body=buffer.getvalue())
|
||||
|
||||
self.logger.custom_info(f'Dataframe stored as parquet in {uri}', metadata)
|
||||
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
|
||||
|
||||
def get_parquet_as_dataframe(self, object_key: str, metadata: dict[str, Any]) -> DataFrame:
|
||||
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:
|
||||
@@ -138,9 +193,22 @@ class MinioRepository:
|
||||
Returns:
|
||||
DataFrame: Loaded DataFrame.
|
||||
"""
|
||||
self.logger.custom_info(f'Getting parquet as dataframe from {object_key}', metadata)
|
||||
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())
|
||||
|
||||
@@ -17,6 +17,7 @@ Capabilities:
|
||||
import ctypes
|
||||
import gc
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from datetime import datetime, timedelta
|
||||
from os import environ, makedirs, path
|
||||
@@ -27,9 +28,14 @@ import mlflow
|
||||
import pandas as pd
|
||||
from mlflow.entities import Experiment
|
||||
from numpy import ndarray
|
||||
from sientia_do.notifications.handlers import NotificationHandler
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||
|
||||
from laborious import metrics
|
||||
|
||||
ARTIFACTS_PATH = './tmp/artifacts'
|
||||
TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl'
|
||||
PREDICTION_COMPRESSED_PATH = 'artifacts/stacking_model.pkl'
|
||||
@@ -56,8 +62,16 @@ def force_memory_release(logger: Logger):
|
||||
logger.info(f'Memory release failed: {e}')
|
||||
|
||||
|
||||
class MLFlowRepository:
|
||||
def __init__(self, host: str, username: str, password: str, logger: Logger):
|
||||
class MLFlowRepository(SientiaMonitoring):
|
||||
def __init__(
|
||||
self,
|
||||
host: str,
|
||||
username: str,
|
||||
password: str,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""Initialize MLflow client and base state.
|
||||
|
||||
Args:
|
||||
@@ -67,6 +81,7 @@ class MLFlowRepository:
|
||||
logger (Logger): Logger instance.
|
||||
"""
|
||||
# set tracking uri
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
mlflow.set_tracking_uri(host)
|
||||
|
||||
environ['MLFLOW_TRACKING_USERNAME'] = username
|
||||
@@ -204,12 +219,15 @@ class MLFlowRepository:
|
||||
Functions related to download and load models
|
||||
"""
|
||||
|
||||
def dowload_artifacts(self, model_name: str, artifact_path: str = 'data_model') -> str:
|
||||
async def dowload_artifacts(
|
||||
self, model_name: str, metadata: dict[str, Any], artifact_path: str = 'data_model'
|
||||
) -> str:
|
||||
"""
|
||||
Download artifacts from the latest production run of a model.
|
||||
|
||||
Args:
|
||||
model_name (str): Registered model name.
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
artifact_path (str): Relative path to artifacts within the run.
|
||||
|
||||
Returns:
|
||||
@@ -225,16 +243,31 @@ class MLFlowRepository:
|
||||
rmtree(full_path)
|
||||
makedirs(output_dir, exist_ok=True)
|
||||
|
||||
self.logger.info(f'Downloading artifacts from {run_id} to {output_dir}')
|
||||
self.info(f'Downloading artifacts from {run_id} to {output_dir}')
|
||||
|
||||
return self.client.download_artifacts(run_id, artifact_path, output_dir)
|
||||
core_labels = self.get_core_labels(metadata, operation_type='download_artifacts')
|
||||
|
||||
def load_predict_model(self, model_name: str, flavor: str = 'sklearn') -> Any:
|
||||
start_time = time.time()
|
||||
try:
|
||||
artifacts = self.client.download_artifacts(run_id, artifact_path, output_dir)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels)
|
||||
|
||||
return artifacts
|
||||
|
||||
async def load_predict_model(
|
||||
self, model_name: str, metadata: dict[str, Any], flavor: str = 'sklearn'
|
||||
) -> Any:
|
||||
"""
|
||||
Load a predictive model from the MLflow Model Registry.
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to download from the registry.
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
flavor (str): Model flavor ('pyfunc', 'sklearn', 'pytorch')
|
||||
artifact_path (str | None): Path to compressed artifacts if model is compressed
|
||||
|
||||
@@ -246,7 +279,11 @@ class MLFlowRepository:
|
||||
- Warnings during the model loading process are suppressed.
|
||||
"""
|
||||
model_uri = f'models:/{model_name}/production'
|
||||
self.logger.info(f'Loading prediction model {model_name} from {model_uri}')
|
||||
self.info(f'Loading prediction model {model_name} from {model_uri}')
|
||||
|
||||
core_labels = self.get_core_labels(metadata, operation_type='load_predict_model')
|
||||
start_time = time.time()
|
||||
try:
|
||||
if flavor == 'pyfunc':
|
||||
model = mlflow.pyfunc.load_model(model_uri)
|
||||
elif flavor == 'sklearn':
|
||||
@@ -255,10 +292,17 @@ class MLFlowRepository:
|
||||
model = mlflow.pytorch.load_model(model_uri)
|
||||
else:
|
||||
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels)
|
||||
return model
|
||||
|
||||
def load_transform_model(self, model_name: str, flavor: str) -> Any:
|
||||
async def load_transform_model(
|
||||
self, model_name: str, metadata: dict[str, Any], flavor: str = 'sklearn'
|
||||
) -> Any:
|
||||
"""
|
||||
Load the latest Production version of a transformation model.
|
||||
|
||||
@@ -267,6 +311,7 @@ class MLFlowRepository:
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to download.
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
|
||||
artifact_path (str | None): Path to compressed artifacts if model is compressed
|
||||
|
||||
@@ -281,7 +326,11 @@ class MLFlowRepository:
|
||||
latest_production_id = self.get_model_run_id(model_name=model_name, stage='Production')
|
||||
model_uri = self.get_model_uri(latest_production_id, prediction=False)
|
||||
|
||||
self.logger.info(f'Loading data model {model_name} from {model_uri}')
|
||||
self.info(f'Loading data model {model_name} from {model_uri}')
|
||||
|
||||
core_labels = self.get_core_labels(metadata, operation_type='load_transform_model')
|
||||
start_time = time.time()
|
||||
try:
|
||||
if flavor == 'sklearn':
|
||||
model = mlflow.sklearn.load_model(model_uri)
|
||||
elif flavor == 'pyfunc':
|
||||
@@ -290,16 +339,29 @@ class MLFlowRepository:
|
||||
model = mlflow.pytorch.load_model(model_uri)
|
||||
else:
|
||||
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_READ_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_READ_COUNT, tags=core_labels)
|
||||
|
||||
return model
|
||||
|
||||
def download_model(
|
||||
self, model_name: str, model_type: str, flavor: str, load_wrapper: bool = False
|
||||
async def download_model(
|
||||
self,
|
||||
model_name: str,
|
||||
metadata: dict[str, Any],
|
||||
model_type: str,
|
||||
flavor: str,
|
||||
load_wrapper: bool = False,
|
||||
) -> tuple[Any, str | None]:
|
||||
"""
|
||||
Download model based on type ("predict" or "transform").
|
||||
|
||||
Args:
|
||||
model_name (str): Name of the model to download
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
model_type (str): Type of model ('predict' or 'transform')
|
||||
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
|
||||
load_wrapper (bool): Whether to load wrapper
|
||||
@@ -308,7 +370,7 @@ class MLFlowRepository:
|
||||
tuple[Any, str | None]: Model object and optional artifact path.
|
||||
"""
|
||||
|
||||
self.logger.info(
|
||||
self.info(
|
||||
f'Downloading {model_type} model {model_name} with flavor {flavor} and load_wrapper {load_wrapper}'
|
||||
)
|
||||
|
||||
@@ -318,15 +380,13 @@ class MLFlowRepository:
|
||||
artifact_path = None
|
||||
|
||||
if load_wrapper:
|
||||
self.logger.info(
|
||||
f'Loading wrapper for {model_type} model {model_name} with flavor {flavor}'
|
||||
)
|
||||
self.info(f'Loading wrapper for {model_type} model {model_name} with flavor {flavor}')
|
||||
|
||||
target = 'prediction_model' if model_type == 'predict' else 'data_model'
|
||||
|
||||
artifact_path = self.dowload_artifacts(model_name, target)
|
||||
artifact_path = await self.dowload_artifacts(model_name, metadata, target)
|
||||
|
||||
self.logger.info(
|
||||
self.info(
|
||||
f'Model with type {model_type} and name {model_name} is compressed, loading from {artifact_path}'
|
||||
)
|
||||
|
||||
@@ -334,10 +394,10 @@ class MLFlowRepository:
|
||||
model = raw_model._model_impl.python_model
|
||||
else:
|
||||
if model_type == 'predict':
|
||||
model = self.load_predict_model(model_name, flavor)
|
||||
model = await self.load_predict_model(model_name, metadata, flavor)
|
||||
|
||||
else:
|
||||
model = self.load_transform_model(model_name, flavor)
|
||||
model = await self.load_transform_model(model_name, metadata, flavor)
|
||||
|
||||
return model, artifact_path
|
||||
|
||||
@@ -361,12 +421,16 @@ class MLFlowRepository:
|
||||
Returns:
|
||||
pd.DataFrame: DataFrame with converted datetime index.
|
||||
"""
|
||||
if data.empty:
|
||||
self.info('Data is empty, skipping datetime index detection and parsing', metadata)
|
||||
return data
|
||||
|
||||
index = data.index
|
||||
|
||||
# Get type of first element of index
|
||||
index_type = type(index[0])
|
||||
|
||||
self.logger.custom_info(f'Index type: {index_type}', metadata)
|
||||
self.info(f'Index type: {index_type}', metadata)
|
||||
|
||||
message = f'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}.'
|
||||
|
||||
@@ -426,7 +490,7 @@ class MLFlowRepository:
|
||||
Returns:
|
||||
dict: Model configuration.
|
||||
"""
|
||||
self.logger.debug(f'Model {model_name} is still valid, using cached version')
|
||||
self.debug(f'Model {model_name} is still valid, using cached version')
|
||||
|
||||
return cache['target']
|
||||
|
||||
@@ -441,17 +505,25 @@ class MLFlowRepository:
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
self.logger.debug(f'Model {model_name} is outdated, downloading a new one')
|
||||
self.debug(f'Model {model_name} is outdated, downloading a new one')
|
||||
|
||||
del self.model_cache[model_key]['target']
|
||||
del self.model_cache[model_key]
|
||||
|
||||
def get_model(self, model_name: str, retention: int, model_type: str, flavor: str) -> Any:
|
||||
async def get_model(
|
||||
self,
|
||||
model_name: str,
|
||||
metadata: dict[str, Any],
|
||||
retention: int,
|
||||
model_type: str,
|
||||
flavor: str,
|
||||
) -> Any:
|
||||
"""
|
||||
Retrieve a model with caching support based on retention policy.
|
||||
|
||||
Args:
|
||||
model_name (str): Name of the model to retrieve
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
retention (int): Cache retention time in minutes (0 = no cache).
|
||||
model_type (str): Type of model ('predict' or 'transform')
|
||||
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
|
||||
@@ -461,8 +533,12 @@ class MLFlowRepository:
|
||||
"""
|
||||
# Retention is 0, download a new model
|
||||
if retention <= 0:
|
||||
model, _artifact_path = self.download_model(
|
||||
model_name=model_name, model_type=model_type, flavor=flavor, load_wrapper=False
|
||||
model, _artifact_path = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type=model_type,
|
||||
flavor=flavor,
|
||||
load_wrapper=False,
|
||||
)
|
||||
return model
|
||||
|
||||
@@ -480,13 +556,17 @@ class MLFlowRepository:
|
||||
# Model is outdated, delete old model files
|
||||
self.handle_outdated_model(model_name=model_name, model_key=model_key)
|
||||
else:
|
||||
self.logger.debug(
|
||||
self.debug(
|
||||
f'Model {model_name} is not in {model_type} cache, downloading a new one'
|
||||
)
|
||||
|
||||
# Donwload new model (without lock to avoid blocking other threads)
|
||||
model, _artifact_path = self.download_model(
|
||||
model_name=model_name, model_type=model_type, flavor=flavor, load_wrapper=False
|
||||
model, _artifact_path = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type=model_type,
|
||||
flavor=flavor,
|
||||
load_wrapper=False,
|
||||
)
|
||||
|
||||
# Update cache with lock
|
||||
@@ -497,27 +577,35 @@ class MLFlowRepository:
|
||||
return model
|
||||
|
||||
@overload
|
||||
def get_cached_operation(
|
||||
async def get_cached_operation(
|
||||
self,
|
||||
model_name: str,
|
||||
data: pd.DataFrame,
|
||||
operation: Literal['transform'],
|
||||
retention: int,
|
||||
flavor: str,
|
||||
metadata: dict[str, Any],
|
||||
) -> pd.DataFrame: ...
|
||||
|
||||
@overload
|
||||
def get_cached_operation(
|
||||
async def get_cached_operation(
|
||||
self,
|
||||
model_name: str,
|
||||
data: pd.DataFrame,
|
||||
operation: Literal['predict'],
|
||||
retention: int,
|
||||
flavor: str,
|
||||
metadata: dict[str, Any],
|
||||
) -> pd.DataFrame | ndarray: ...
|
||||
|
||||
def get_cached_operation(
|
||||
self, model_name: str, data: pd.DataFrame, operation: str, retention: int, flavor: str
|
||||
async def get_cached_operation(
|
||||
self,
|
||||
model_name: str,
|
||||
data: pd.DataFrame,
|
||||
operation: str,
|
||||
retention: int,
|
||||
flavor: str,
|
||||
metadata: dict[str, Any],
|
||||
) -> pd.DataFrame | ndarray:
|
||||
"""
|
||||
Execute a cached operation using the requested model.
|
||||
@@ -527,21 +615,25 @@ class MLFlowRepository:
|
||||
data (pd.DataFrame): Input data.
|
||||
retention (int): Cache retention in minutes.
|
||||
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch').
|
||||
|
||||
metadata (dict[str, Any]): Metadata used for structured logging.
|
||||
Returns:
|
||||
pd.DataFrame | ndarray: Operation result.
|
||||
"""
|
||||
if operation not in ['transform', 'predict']:
|
||||
raise ValueError("Invalid operation. Use 'transform' or 'predict'.")
|
||||
|
||||
model = self.get_model(
|
||||
model_name=model_name, retention=retention, model_type=operation, flavor=flavor
|
||||
model = await self.get_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
retention=retention,
|
||||
model_type=operation,
|
||||
flavor=flavor,
|
||||
)
|
||||
|
||||
prediction = model.predict(data)
|
||||
|
||||
if retention == 0:
|
||||
self.logger.info(f'Deleting model {model_name}:{operation} from memory')
|
||||
self.info(f'Deleting model {model_name}:{operation} from memory')
|
||||
del model
|
||||
|
||||
force_memory_release(self.logger)
|
||||
@@ -552,7 +644,7 @@ class MLFlowRepository:
|
||||
Functions related to model retraining
|
||||
"""
|
||||
|
||||
def fit_models(
|
||||
async def fit_models(
|
||||
self,
|
||||
model_name: str,
|
||||
data: pd.DataFrame,
|
||||
@@ -585,8 +677,8 @@ class MLFlowRepository:
|
||||
`data_model`, including optional artifact paths.
|
||||
"""
|
||||
|
||||
self.logger.custom_info(f'Starting model experiment creation for {model_name}', metadata)
|
||||
self.logger.custom_debug(
|
||||
self.info(f'Starting model experiment creation for {model_name}', metadata)
|
||||
self.debug(
|
||||
f'Model configuration - transform_flavor: {transform_flavor}, predict_flavor: {predict_flavor}, target_name: {target_name}',
|
||||
metadata,
|
||||
)
|
||||
@@ -594,26 +686,26 @@ class MLFlowRepository:
|
||||
# data.to_csv(
|
||||
# f"tmp/retrain_data_{model_name}.csv", index=True)
|
||||
|
||||
self.logger.custom_info(
|
||||
f'Retrieved latest production run ID: {latest_production_id}', metadata
|
||||
)
|
||||
self.logger.custom_info(f'Loading transformation model for {model_name}', metadata)
|
||||
self.info(f'Retrieved latest production run ID: {latest_production_id}', metadata)
|
||||
self.info(f'Loading transformation model for {model_name}', metadata)
|
||||
|
||||
load_transform_wrapper = transform_flavor == 'pyfunc'
|
||||
|
||||
data_model, data_artifact_path = self.download_model(
|
||||
data_model, data_artifact_path = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type='transform',
|
||||
flavor=transform_flavor,
|
||||
load_wrapper=load_transform_wrapper,
|
||||
)
|
||||
|
||||
self.logger.custom_info(f'Loading prediction model for {model_name}', metadata)
|
||||
self.info(f'Loading prediction model for {model_name}', metadata)
|
||||
|
||||
load_predict_wrapper = predict_flavor == 'pyfunc'
|
||||
|
||||
prediction_model, prediction_artifact_path = self.download_model(
|
||||
prediction_model, prediction_artifact_path = await self.download_model(
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
model_type='predict',
|
||||
flavor=predict_flavor,
|
||||
load_wrapper=load_predict_wrapper,
|
||||
@@ -636,24 +728,22 @@ class MLFlowRepository:
|
||||
|
||||
treated_data = treated_data.drop_duplicates(subset=['timestamp'], keep='first')
|
||||
|
||||
self.logger.custom_debug(f'Treated data index: {treated_data.index}', metadata)
|
||||
self.debug(f'Treated data index: {treated_data.index}', metadata)
|
||||
|
||||
# treated_data.to_csv(
|
||||
# f"tmp/retrain_treated_data_{model_name}.csv", index=True)
|
||||
|
||||
self.logger.custom_debug(f'Transformed data shape: {treated_data.shape}', metadata)
|
||||
self.debug(f'Transformed data shape: {treated_data.shape}', metadata)
|
||||
|
||||
if target_name is None:
|
||||
target_name = data_model.target_variable
|
||||
self.logger.custom_debug(
|
||||
f'Using target variable from data model: {target_name}', metadata
|
||||
)
|
||||
self.debug(f'Using target variable from data model: {target_name}', metadata)
|
||||
else:
|
||||
self.logger.custom_debug(f'Using provided target variable: {target_name}', metadata)
|
||||
self.debug(f'Using provided target variable: {target_name}', metadata)
|
||||
|
||||
# Check if treated_data contains target variable
|
||||
if target_name not in treated_data.columns:
|
||||
self.logger.custom_debug(
|
||||
self.debug(
|
||||
f'Target variable {target_name} not found in treated data, aligning data with treated data indexes',
|
||||
metadata,
|
||||
)
|
||||
@@ -665,9 +755,7 @@ class MLFlowRepository:
|
||||
)
|
||||
else:
|
||||
# Uses target variable from treated data
|
||||
self.logger.custom_debug(
|
||||
f'Target variable {target_name} found in treated data, using it', metadata
|
||||
)
|
||||
self.debug(f'Target variable {target_name} found in treated data, using it', metadata)
|
||||
retrain_dataset = treated_data
|
||||
|
||||
# retrain_dataset.to_csv(
|
||||
@@ -675,9 +763,7 @@ class MLFlowRepository:
|
||||
|
||||
prediction_model.fit(retrain_dataset)
|
||||
|
||||
self.logger.custom_info(
|
||||
f'Model experiment creation completed successfully for {model_name}', metadata
|
||||
)
|
||||
self.info(f'Model experiment creation completed successfully for {model_name}', metadata)
|
||||
|
||||
retrain_data = {
|
||||
'prediction_model': {
|
||||
@@ -688,7 +774,7 @@ class MLFlowRepository:
|
||||
}
|
||||
return retrain_data
|
||||
|
||||
def log_model(self, model_data: dict, flavor: str, model_type: str, metadata: dict):
|
||||
async def log_model(self, model_data: dict, flavor: str, model_type: str, metadata: dict):
|
||||
"""Log a model into the active MLflow run.
|
||||
|
||||
Args:
|
||||
@@ -699,23 +785,35 @@ class MLFlowRepository:
|
||||
"""
|
||||
model = model_data['model']
|
||||
|
||||
self.logger.custom_debug(f'Logging {model_type} model to {model_type}', metadata)
|
||||
self.debug(f'Logging {model_type} model to {model_type}', metadata)
|
||||
|
||||
core_labels = self.get_core_labels(metadata, operation_type='log_model')
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if flavor == 'sklearn':
|
||||
mlflow.sklearn.log_model(model, model_type)
|
||||
elif flavor == 'pyfunc':
|
||||
code_path = [path.join(model_data['artifact_path'], 'code', 'utils')]
|
||||
|
||||
self.logger.custom_debug(f'Code path: {code_path}', metadata)
|
||||
self.debug(f'Code path: {code_path}', metadata)
|
||||
|
||||
model.store_model(artifact_path=model_type, code_path=code_path, to_disk=False)
|
||||
|
||||
self.logger.custom_debug('Model uploaded successfully', metadata)
|
||||
self.debug('Model uploaded successfully', metadata)
|
||||
elif flavor == 'pytorch':
|
||||
mlflow.pytorch.log_model(model, model_type)
|
||||
else:
|
||||
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||
|
||||
def create_new_experiment(
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels)
|
||||
|
||||
async def create_new_experiment(
|
||||
self,
|
||||
model_name: str,
|
||||
data: pd.DataFrame,
|
||||
@@ -754,7 +852,7 @@ class MLFlowRepository:
|
||||
|
||||
model_temp_path = path.join(ARTIFACTS_PATH, model_name)
|
||||
|
||||
self.logger.custom_info(f'Starting model retraining process for {model_name}', metadata)
|
||||
self.info(f'Starting model retraining process for {model_name}', metadata)
|
||||
|
||||
original_params = self.get_model_params(latest_production_id)
|
||||
retrain_params = {
|
||||
@@ -771,7 +869,7 @@ class MLFlowRepository:
|
||||
|
||||
current_run_name = self.get_next_run_name(experiment_name)
|
||||
|
||||
self.logger.custom_debug(f'Attributes: {retrain_params}', metadata)
|
||||
self.debug(f'Attributes: {retrain_params}', metadata)
|
||||
|
||||
data_path = f'{model_temp_path}/retrain_data.csv'
|
||||
|
||||
@@ -779,28 +877,31 @@ class MLFlowRepository:
|
||||
|
||||
data.to_csv(data_path, index=True)
|
||||
|
||||
self.logger.custom_info(
|
||||
self.info(
|
||||
f'Starting model upload for {experiment_name} with run name {current_run_name}',
|
||||
metadata,
|
||||
)
|
||||
|
||||
core_labels = self.get_core_labels(metadata, operation_type='create_new_experiment')
|
||||
start_time = time.time()
|
||||
try:
|
||||
with mlflow.start_run(
|
||||
experiment_id=experiment.experiment_id,
|
||||
run_name=current_run_name,
|
||||
description=experiment_description,
|
||||
) as _run:
|
||||
run_id = _run.info.run_id
|
||||
self.logger.custom_info('Logging data model', metadata)
|
||||
self.info('Logging data model', metadata)
|
||||
# dynamic parameters, including model itself
|
||||
self.log_model(data_model, transform_flavor, 'data_model', metadata)
|
||||
await self.log_model(data_model, transform_flavor, 'data_model', metadata)
|
||||
|
||||
# dynamic parameters, including model itself
|
||||
self.logger.custom_info('Logging prediction model', metadata)
|
||||
self.log_model(prediction_model, predict_flavor, 'prediction_model', metadata)
|
||||
self.info('Logging prediction model', metadata)
|
||||
await self.log_model(prediction_model, predict_flavor, 'prediction_model', metadata)
|
||||
|
||||
self.logger.custom_info(f'Model logged successfully for {model_name}', metadata)
|
||||
self.info(f'Model logged successfully for {model_name}', metadata)
|
||||
|
||||
self.logger.custom_info(f'Logging remaining parameters for {model_name}', metadata)
|
||||
self.info(f'Logging remaining parameters for {model_name}', metadata)
|
||||
|
||||
# update transfomation model
|
||||
# fixed parameters
|
||||
@@ -809,15 +910,22 @@ class MLFlowRepository:
|
||||
# log the data raw
|
||||
mlflow.log_artifact(data_path)
|
||||
|
||||
self.logger.custom_info('Deleting model from filesystem', metadata)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels)
|
||||
|
||||
self.info('Deleting model from filesystem', metadata)
|
||||
if path.exists(model_temp_path):
|
||||
rmtree(model_temp_path)
|
||||
|
||||
self.logger.custom_info('Deleting prediction model from memory', metadata)
|
||||
self.info('Deleting prediction model from memory', metadata)
|
||||
del prediction_model['model']
|
||||
del prediction_model
|
||||
|
||||
self.logger.custom_info('Deleting data model from memory', metadata)
|
||||
self.info('Deleting data model from memory', metadata)
|
||||
del data_model['model']
|
||||
del data_model
|
||||
|
||||
@@ -829,7 +937,7 @@ class MLFlowRepository:
|
||||
'experiment_name': experiment.name,
|
||||
}
|
||||
|
||||
def update_production_model_by_run_id(
|
||||
async def update_production_model_by_run_id(
|
||||
self, run_id: str, model_name: str, metadata: dict
|
||||
) -> dict:
|
||||
"""
|
||||
@@ -857,14 +965,24 @@ class MLFlowRepository:
|
||||
4. Archives existing production versions
|
||||
"""
|
||||
|
||||
self.logger.custom_info(
|
||||
self.info(
|
||||
f'Starting production model update for {model_name} with run ID: {run_id}', metadata
|
||||
)
|
||||
|
||||
# Registrar o modelo
|
||||
# Aqui estamos assumindo que você já tem um modelo salvo, caso contrário você precisará treiná-lo e salvá-lo primeiro.
|
||||
# Se o modelo já está registrado, você pode usar o método register_model() ou pyfunc.load_model() para isso.
|
||||
|
||||
core_labels = self.get_core_labels(metadata, operation_type='register_model')
|
||||
start_time = time.time()
|
||||
try:
|
||||
mlflow.register_model(f'runs:/{run_id}/prediction_model', model_name)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels)
|
||||
|
||||
# Obter a versão mais recente registrada do modelo
|
||||
model_versions = self.client.get_registered_model(model_name).latest_versions
|
||||
@@ -875,9 +993,23 @@ class MLFlowRepository:
|
||||
max_version = max(model_versions, key=lambda x: int(x.version)).version
|
||||
|
||||
# Mover a versão mais recente do modelo para o estágio de 'Production'
|
||||
self.client.transition_model_version_stage(
|
||||
name=model_name, version=max_version, stage='Production', archive_existing_versions=True
|
||||
core_labels = self.get_core_labels(
|
||||
metadata, operation_type='transition_model_version_stage'
|
||||
)
|
||||
start_time = time.time()
|
||||
try:
|
||||
self.client.transition_model_version_stage(
|
||||
name=model_name,
|
||||
version=max_version,
|
||||
stage='Production',
|
||||
archive_existing_versions=True,
|
||||
)
|
||||
except Exception as e:
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=core_labels)
|
||||
raise e
|
||||
|
||||
await self.observe_lag(start_time, metrics.MODEL_WRITE_LAG, core_labels)
|
||||
await self.emit_metric(metric_object=metrics.MODEL_WRITE_COUNT, tags=core_labels)
|
||||
|
||||
return {'model_name': model_name, 'version': max_version, 'mlflow_run_id': run_id}
|
||||
|
||||
@@ -885,7 +1017,9 @@ class MLFlowRepository:
|
||||
Functions that provide the interface to model operations
|
||||
"""
|
||||
|
||||
def transform(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict):
|
||||
async def transform(
|
||||
self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict
|
||||
):
|
||||
"""
|
||||
Transform data using a cached transformation model.
|
||||
|
||||
@@ -919,9 +1053,7 @@ class MLFlowRepository:
|
||||
and returned in the response structure rather than propagated.
|
||||
"""
|
||||
|
||||
self.logger.custom_debug(
|
||||
f'Data received for model transformation: {data.head(5).to_csv()}', metadata
|
||||
)
|
||||
self.debug(f'Data received for model transformation: {data.head(5).to_csv()}', metadata)
|
||||
|
||||
# data.to_csv(
|
||||
# f"tmp/data_{model_name}.csv", index=True)
|
||||
@@ -930,11 +1062,16 @@ class MLFlowRepository:
|
||||
flavor = model_config.get('transform_flavor', 'sklearn')
|
||||
|
||||
try:
|
||||
transformed_data: pd.DataFrame = self.get_cached_operation(
|
||||
model_name, data, 'transform', model_retention, flavor
|
||||
transformed_data: pd.DataFrame = await self.get_cached_operation(
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='transform',
|
||||
retention=model_retention,
|
||||
flavor=flavor,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
self.logger.custom_debug(
|
||||
self.debug(
|
||||
f'Data received from model transformation: {transformed_data.head(5).to_csv()}',
|
||||
metadata,
|
||||
)
|
||||
@@ -952,7 +1089,9 @@ class MLFlowRepository:
|
||||
'content': {'message': str(e), 'traceback': traceback.format_exc()},
|
||||
}
|
||||
|
||||
def predict(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict):
|
||||
async def predict(
|
||||
self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict
|
||||
):
|
||||
"""
|
||||
Generate predictions using a cached prediction model.
|
||||
|
||||
@@ -1000,21 +1139,24 @@ class MLFlowRepository:
|
||||
input_index = data.index
|
||||
start_time = datetime.now()
|
||||
|
||||
self.logger.custom_debug(
|
||||
f'Data received for model prediction: {data.head(5).to_csv()}', metadata
|
||||
)
|
||||
self.debug(f'Data received for model prediction: {data.head(5).to_csv()}', metadata)
|
||||
|
||||
# data.to_csv(
|
||||
# f"tmp/treated_data_{model_name}.csv", index=True)
|
||||
|
||||
predict_data = self.get_cached_operation(
|
||||
model_name, data, 'predict', model_retention, flavor
|
||||
predict_data = await self.get_cached_operation(
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=model_retention,
|
||||
flavor=flavor,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
end_time = datetime.now()
|
||||
|
||||
if isinstance(predict_data, pd.DataFrame):
|
||||
self.logger.custom_debug(
|
||||
self.debug(
|
||||
f'Data received from model prediction: {data.head(5).to_csv()}', metadata
|
||||
)
|
||||
|
||||
@@ -1038,7 +1180,7 @@ class MLFlowRepository:
|
||||
'content': {'message': str(e), 'traceback': traceback.format_exc()},
|
||||
}
|
||||
|
||||
def retrain_model(
|
||||
async def retrain_model(
|
||||
self, data: pd.DataFrame, model_name: str, model_config: dict, metadata: dict
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -1083,23 +1225,23 @@ class MLFlowRepository:
|
||||
Exception: Any other exception during the retraining process
|
||||
"""
|
||||
|
||||
self.logger.custom_info(f'Starting model retraining workflow for {model_name}', metadata)
|
||||
self.logger.custom_debug(f'Data received for model retraining: {data.to_csv()}', metadata)
|
||||
self.info(f'Starting model retraining workflow for {model_name}', metadata)
|
||||
self.debug(f'Data received for model retraining: {data.to_csv()}', metadata)
|
||||
|
||||
target_name = model_config.get('target', None)
|
||||
|
||||
transform_flavor = model_config.get('transform_flavor', 'sklearn')
|
||||
predict_flavor = model_config.get('predict_flavor', 'sklearn')
|
||||
|
||||
self.logger.custom_debug(
|
||||
self.debug(
|
||||
f'Model configuration - transform_flavor: {transform_flavor}, predict_flavor: {predict_flavor}, target_name: {target_name}',
|
||||
metadata,
|
||||
)
|
||||
|
||||
try:
|
||||
latest_production_id = self.get_model_run_id(model_name, stage='Production')
|
||||
self.logger.custom_info('Creating model experiment environment', metadata)
|
||||
retrain_data = self.fit_models(
|
||||
self.info('Creating model experiment environment', metadata)
|
||||
retrain_data = await self.fit_models(
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
transform_flavor=transform_flavor,
|
||||
@@ -1108,12 +1250,10 @@ class MLFlowRepository:
|
||||
metadata=metadata,
|
||||
latest_production_id=latest_production_id,
|
||||
)
|
||||
self.logger.custom_info(
|
||||
f'Model experiment created successfully: {retrain_data}', metadata
|
||||
)
|
||||
self.info(f'Model experiment created successfully: {retrain_data}', metadata)
|
||||
|
||||
self.logger.custom_info('Saving model retrain', metadata)
|
||||
experiment = self.create_new_experiment(
|
||||
self.info('Saving model retrain', metadata)
|
||||
experiment = await self.create_new_experiment(
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
retrain_data=retrain_data,
|
||||
@@ -1122,7 +1262,7 @@ class MLFlowRepository:
|
||||
metadata=metadata,
|
||||
latest_production_id=latest_production_id,
|
||||
)
|
||||
self.logger.custom_info(
|
||||
self.info(
|
||||
f'Model retraining completed successfully for experiment: {experiment}', metadata
|
||||
)
|
||||
|
||||
@@ -1133,7 +1273,7 @@ class MLFlowRepository:
|
||||
}
|
||||
except Exception as e:
|
||||
error_msg = f'Error retraining model {model_name}: {e}'
|
||||
self.logger.custom_info(error_msg, metadata)
|
||||
self.info(error_msg, metadata)
|
||||
return {
|
||||
'success': False,
|
||||
'experiment': None,
|
||||
@@ -1141,7 +1281,7 @@ class MLFlowRepository:
|
||||
'traceback': traceback.format_exc(),
|
||||
}
|
||||
|
||||
def update_production_model(
|
||||
async def update_production_model(
|
||||
self, experiment: dict[str, Any], model_name: str, metadata: dict
|
||||
) -> dict:
|
||||
"""
|
||||
@@ -1190,7 +1330,7 @@ class MLFlowRepository:
|
||||
"""
|
||||
run_id = experiment['run_id']
|
||||
experiment_id = experiment['experiment_id']
|
||||
metadata_result = self.update_production_model_by_run_id(run_id, model_name, metadata)
|
||||
metadata_result = await self.update_production_model_by_run_id(run_id, model_name, metadata)
|
||||
|
||||
metadata_result['mlflow_experiment_id'] = experiment_id
|
||||
|
||||
|
||||
@@ -8,11 +8,14 @@ from typing import Any
|
||||
|
||||
from asyncua import Client
|
||||
from asyncua.crypto.security_policies import SecurityPolicyBasic256
|
||||
from asyncua.ua import DataValue, DateTime, Variant, VariantType
|
||||
from asyncua.ua import DataValue, Variant, VariantType
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.temporal.activities.base import BaseActivity
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
|
||||
from laborious import metrics
|
||||
|
||||
data_type_map = {
|
||||
'float': {
|
||||
@@ -38,13 +41,15 @@ data_type_map = {
|
||||
}
|
||||
|
||||
|
||||
class OpcRepository(BaseActivity):
|
||||
class OpcRepository(SientiaMonitoring):
|
||||
def __init__(
|
||||
self,
|
||||
opc_id: str,
|
||||
url: str,
|
||||
server_name: str,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
reconnection_interval: int = 60,
|
||||
server_uri: str | None = None,
|
||||
cert_path: str | None = None,
|
||||
@@ -53,6 +58,7 @@ class OpcRepository(BaseActivity):
|
||||
):
|
||||
self.url = url
|
||||
self.id = opc_id
|
||||
self.server_name = server_name
|
||||
self.server_uri = server_uri
|
||||
self.cert_path = cert_path
|
||||
self.private_key_path = private_key_path
|
||||
@@ -61,10 +67,11 @@ class OpcRepository(BaseActivity):
|
||||
self.error_count = 0
|
||||
self.reconnection_interval = reconnection_interval
|
||||
self.last_reconnection_time: None | datetime = None
|
||||
self.disconnection_interval = 10.0
|
||||
self.notification_handler = notification_handler
|
||||
self.client: None | Client = None
|
||||
|
||||
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
|
||||
self.metadata = {
|
||||
'model_name': '-',
|
||||
@@ -136,7 +143,9 @@ class OpcRepository(BaseActivity):
|
||||
|
||||
if self.cert_path:
|
||||
await self.set_security()
|
||||
self.logger.custom_info(f'Starting connection to OPC server {self.id}...', self.metadata)
|
||||
self.logger.custom_info(
|
||||
f'Starting connection to OPC server {self.id}:{self.server_name}...', self.metadata
|
||||
)
|
||||
return await self.try_connect()
|
||||
|
||||
async def try_connect(self) -> tuple[bool, dict[str, Any]]:
|
||||
@@ -154,6 +163,11 @@ class OpcRepository(BaseActivity):
|
||||
- dict: Error information if connection failed
|
||||
"""
|
||||
|
||||
tags = {
|
||||
'pod_id': self.pod_id,
|
||||
'server_name': self.server_name,
|
||||
}
|
||||
await self.emit_metric(metrics.OPC_CONNECTIONS_TOTAL, tags)
|
||||
try:
|
||||
self.last_reconnection_time = datetime.now()
|
||||
if self.client is None:
|
||||
@@ -164,13 +178,26 @@ class OpcRepository(BaseActivity):
|
||||
'level': NotificationLevel.ERROR,
|
||||
}
|
||||
await self.client.connect()
|
||||
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.OPC_CONNECTION_STATUS,
|
||||
method='set',
|
||||
tags={
|
||||
**tags,
|
||||
'server_url': self.url,
|
||||
},
|
||||
value=1,
|
||||
)
|
||||
|
||||
return True, {}
|
||||
except Exception as e:
|
||||
self.disconnect()
|
||||
await self.disconnect()
|
||||
|
||||
trace = traceback.format_exc()
|
||||
self.logger.custom_error(trace, self.metadata)
|
||||
|
||||
await self.emit_metric(metrics.OPC_CONNECTIONS_FAILED, tags)
|
||||
|
||||
return False, {
|
||||
'notification_id': f'OPC_CONNECTION_ERROR_{self.id}',
|
||||
'message': f'Failed to connect to OPC server: {e}',
|
||||
@@ -202,7 +229,7 @@ class OpcRepository(BaseActivity):
|
||||
'traceback': traceback.format_exc(),
|
||||
}
|
||||
)
|
||||
await asyncio.sleep(0.1 * i)
|
||||
await asyncio.sleep(self.disconnection_interval * i)
|
||||
return error_stack
|
||||
|
||||
async def disconnect(self):
|
||||
@@ -218,7 +245,7 @@ class OpcRepository(BaseActivity):
|
||||
|
||||
errors = await self.disconnection_fallback()
|
||||
if errors:
|
||||
self.send_notification(
|
||||
await self.send_notification_async(
|
||||
metadata=self.metadata,
|
||||
notification_id=f'OPC_DISCONNECTION_ERROR_{self.id}',
|
||||
message='Failed to disconnect from OPC server in 5 attempts.',
|
||||
@@ -228,6 +255,16 @@ class OpcRepository(BaseActivity):
|
||||
)
|
||||
else:
|
||||
self.logger.warning(f'Disconnected from OPC server {self.id} successfully')
|
||||
await self.emit_metric(
|
||||
metric_object=metrics.OPC_CONNECTION_STATUS,
|
||||
method='set',
|
||||
tags={
|
||||
'pod_id': self.pod_id,
|
||||
'server_name': self.server_name,
|
||||
'server_url': self.url,
|
||||
},
|
||||
value=0,
|
||||
)
|
||||
|
||||
self.client = None
|
||||
|
||||
@@ -382,12 +419,12 @@ class OpcRepository(BaseActivity):
|
||||
|
||||
data = data_type_map[data_type]['converter'](value)
|
||||
logger.custom_info(f'Writing {data} - {type(data)} to {node}', metadata)
|
||||
now = datetime.now()
|
||||
# now = datetime.now() # NOSONAR
|
||||
ua_data = DataValue(
|
||||
Variant(data, data_type_map[data_type]['opc_type']),
|
||||
SourceTimestamp=DateTime(
|
||||
now.year, now.month, now.day, now.hour, now.minute, now.second, now.microsecond
|
||||
),
|
||||
# SourceTimestamp=DateTime( # NOSONAR
|
||||
# now.year, now.month, now.day, now.hour, now.minute, now.second, now.microsecond # NOSONAR
|
||||
# ), # NOSONAR
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -114,10 +114,6 @@ python_functions = ["test_*"]
|
||||
addopts = [
|
||||
"-v",
|
||||
"--strict-markers",
|
||||
"--cov=model_manager",
|
||||
"--cov-report=term-missing",
|
||||
"--cov-report=html",
|
||||
"--cov-report=xml",
|
||||
]
|
||||
markers = [
|
||||
"asyncio: marks tests as async",
|
||||
|
||||
@@ -3,7 +3,7 @@ psycopg2-binary
|
||||
sqlalchemy
|
||||
asyncua
|
||||
redis
|
||||
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.4.7
|
||||
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.5.2
|
||||
prometheus-client
|
||||
botocore
|
||||
boto3
|
||||
|
||||
@@ -3,7 +3,7 @@ psycopg2-binary
|
||||
sqlalchemy
|
||||
asyncua
|
||||
redis
|
||||
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.4.7
|
||||
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.5.2
|
||||
git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.39.0
|
||||
prometheus-client
|
||||
botocore
|
||||
|
||||
35
tests.ipynb
35
tests.ipynb
@@ -453,11 +453,40 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 6,
|
||||
"id": "1fbb3788",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"test\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from unittest.mock import MagicMock\n",
|
||||
"from asyncua.ua.uaerrors import BadAlreadyExists\n",
|
||||
"\n",
|
||||
"mock1 = MagicMock(\n",
|
||||
" side_effect = Exception(\"test\")\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"mock2 = MagicMock(\n",
|
||||
" side_effect = BadAlreadyExists(\"test\")\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"try:\n",
|
||||
" mock1()\n",
|
||||
"except ValueError as e:\n",
|
||||
" try:\n",
|
||||
" mock2()\n",
|
||||
" except BadAlreadyExists as e:\n",
|
||||
" print(e)\n",
|
||||
"except Exception as e:\n",
|
||||
" print(e)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
from pytest import mark
|
||||
|
||||
@@ -13,7 +13,10 @@ from laborious.activities.storage import Storage
|
||||
@patch('laborious.activities.activities.MLFlow.__init__')
|
||||
@patch('laborious.activities.activities.OPC.__init__')
|
||||
@patch('laborious.activities.activities.Gates.__init__')
|
||||
def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init):
|
||||
@patch('laborious.activities.activities.MetricsController')
|
||||
def test___init__(
|
||||
mock_metrics_controller, mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init
|
||||
):
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -70,6 +73,7 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_mlflow_init.assert_called_once_with(
|
||||
@@ -81,22 +85,32 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_opc_init.assert_called_once_with(
|
||||
ANY, opc_servers=opc_config, logger=logger, notification_handler=notification_handler
|
||||
ANY,
|
||||
opc_servers=opc_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_gates_init.assert_called_once_with(
|
||||
ANY, logger=logger, notification_handler=notification_handler
|
||||
ANY,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.activities.Storage', return_value=MagicMock())
|
||||
@patch('laborious.activities.activities.MLFlow', return_value=MagicMock())
|
||||
@patch('laborious.activities.activities.OPC', return_value=MagicMock())
|
||||
async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init):
|
||||
@patch('laborious.activities.activities.Storage')
|
||||
@patch('laborious.activities.activities.MLFlow')
|
||||
@patch('laborious.activities.activities.OPC')
|
||||
@patch('laborious.activities.activities.Gates')
|
||||
async def test_shutdown(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init):
|
||||
mock_opc_init.close = AsyncMock()
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -136,5 +150,7 @@ async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init):
|
||||
)
|
||||
|
||||
await activities.shutdown()
|
||||
mock_opc_init.shutdown.assert_called_once()
|
||||
mock_opc_init.close.assert_called_once()
|
||||
mock_storage_init.close.assert_called_once()
|
||||
mock_mlflow_init.close.assert_called_once()
|
||||
mock_gates_init.close.assert_called_once()
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
from pytest import fixture, mark
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
@@ -11,6 +11,7 @@ def gates_activity():
|
||||
gates = Gates(
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
gates.error = MagicMock()
|
||||
gates.debug = MagicMock()
|
||||
@@ -18,6 +19,8 @@ def gates_activity():
|
||||
gates.warning = MagicMock()
|
||||
gates.critical = MagicMock()
|
||||
gates.send_notification = MagicMock()
|
||||
gates.send_notification_async = AsyncMock()
|
||||
gates.emit_metric = AsyncMock()
|
||||
return gates
|
||||
|
||||
|
||||
@@ -71,7 +74,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='INTPUT_GATE_ERROR__EMPTY_DATA',
|
||||
message="Error in filter EMPTY_DATA:{'policy': 'STOP', 'config': {}}: \n Test error",
|
||||
@@ -176,7 +179,7 @@ async def test_mlflow_response_gate_filter_exception(
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER',
|
||||
message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error",
|
||||
@@ -225,7 +228,7 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'API error occurred')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@@ -298,7 +301,7 @@ async def test_mlflow_content_gate_filter_exception(
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR',
|
||||
message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test error",
|
||||
@@ -344,7 +347,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Transformed data not passed the content filter')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@@ -639,57 +642,103 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}},
|
||||
}
|
||||
await gates_activity.write_metrics(input_data)
|
||||
mock_metrics.PREDICTIONS_WRITTEN_COUNT.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTIONS_WRITTEN_COUNT.labels.return_value.inc.assert_called_once_with()
|
||||
|
||||
mock_metrics.PREDICTION_CONFIDENCE_MONITOR.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTION_CONFIDENCE_MONITOR.labels.return_value.set.assert_called_once_with(0.9)
|
||||
|
||||
mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels.return_value.observe.assert_called_once_with(
|
||||
0.1
|
||||
)
|
||||
|
||||
mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
opc_server_id='server1',
|
||||
tag='tag1',
|
||||
)
|
||||
metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
opc_server_id='server1',
|
||||
tag='tag1',
|
||||
)
|
||||
metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR,
|
||||
method='set',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
value=0.9,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
assert mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.return_value.inc.call_count == 2
|
||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.return_value.observe.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(0.1),
|
||||
call(0.2),
|
||||
],
|
||||
any_order=True,
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
value=0.1,
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag1',
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag1',
|
||||
},
|
||||
value=0.1,
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag2',
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag2',
|
||||
},
|
||||
value=0.2,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
import numpy as np
|
||||
from pytest import fixture, mark, raises
|
||||
@@ -25,6 +25,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert mlflow.mlflow_host == 'http://localhost'
|
||||
@@ -32,7 +33,9 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
assert mlflow.mlflow_username == 'admin'
|
||||
assert mlflow.mlflow_password == 'admin'
|
||||
|
||||
mock_mlflow_repository.assert_called_once_with('http://localhost:5000', 'admin', 'admin', ANY)
|
||||
mock_mlflow_repository.assert_called_once_with(
|
||||
'http://localhost:5000', 'admin', 'admin', ANY, ANY, ANY
|
||||
)
|
||||
|
||||
mock_minio_repository.assert_called_once_with(
|
||||
logger=ANY,
|
||||
@@ -42,6 +45,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=ANY,
|
||||
)
|
||||
|
||||
|
||||
@@ -63,9 +67,15 @@ def mlflow(mock_minio_repository, mock_mlflow_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
mlflow.model_monitoring_repository = AsyncMock()
|
||||
mlflow.minio_repository = AsyncMock()
|
||||
|
||||
mlflow.send_notification = MagicMock()
|
||||
mlflow.emit_metric = AsyncMock()
|
||||
mlflow.send_notification_async = AsyncMock()
|
||||
|
||||
return mlflow
|
||||
|
||||
@@ -225,6 +235,8 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
|
||||
'message': 'Model retrained successfully.',
|
||||
}
|
||||
|
||||
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock()
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
@@ -242,11 +254,9 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
|
||||
|
||||
timestamp = raw_data.__getitem__.return_value.max.return_value
|
||||
|
||||
raw_data.sort_values.assert_called_once_with('created_at', ascending=False)
|
||||
raw_data.sort_values.return_value.drop_duplicates.assert_called_once_with(
|
||||
subset=['variable', 'timestamp'], keep='first'
|
||||
)
|
||||
raw_data = raw_data.sort_values.return_value.drop_duplicates.return_value
|
||||
raw_data.sort_values.assert_not_called()
|
||||
raw_data.drop_duplicates.assert_called_once_with(subset=['variable', 'timestamp'], keep='first')
|
||||
raw_data = raw_data.drop_duplicates.return_value
|
||||
|
||||
raw_data.drop.assert_has_calls(
|
||||
[
|
||||
@@ -301,6 +311,10 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
'message': 'Model retrained failed.',
|
||||
}
|
||||
|
||||
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock(
|
||||
columns=['variable', 'timestamp', 'value', 'created_at']
|
||||
)
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
@@ -360,7 +374,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
mlflow.send_notification.assert_called_once_with(
|
||||
mlflow.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='RETRAIN_MODEL_ERROR',
|
||||
message='Error retraining model test_model: Model retrained failed.',
|
||||
@@ -464,7 +478,7 @@ async def test_update_production_model_error(mlflow):
|
||||
await mlflow.update_production_model(input_data)
|
||||
except Exception as e:
|
||||
assert str(e) == 'Error updating production model'
|
||||
mlflow.send_notification.assert_called_once_with(
|
||||
mlflow.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
|
||||
message='Error updating production model test_model: Error updating production model',
|
||||
|
||||
@@ -18,8 +18,13 @@ metadata = {
|
||||
|
||||
|
||||
def test__init__():
|
||||
servers = {'server1': 'config'}
|
||||
opc = OPC(opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock())
|
||||
servers = {'server1': {'id': 'server1'}}
|
||||
opc = OPC(
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert opc.opc_servers == servers
|
||||
assert opc.opc_repository == {}
|
||||
@@ -27,9 +32,10 @@ def test__init__():
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.opc.OpcRepository')
|
||||
@patch('laborious.activities.opc.OPC.send_notification')
|
||||
@patch('laborious.activities.opc.OPC.send_notification_async')
|
||||
async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
mock_logger = MagicMock()
|
||||
mock_metrics_controller = AsyncMock()
|
||||
server1 = MagicMock(
|
||||
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
|
||||
)
|
||||
@@ -55,6 +61,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
mock_notification_handler = MagicMock()
|
||||
servers = {
|
||||
'server1': {
|
||||
'server_name': 'server1',
|
||||
'id': 'server1',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
@@ -64,6 +71,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
'reconnection_interval': 60,
|
||||
},
|
||||
'server2': {
|
||||
'server_name': 'server2',
|
||||
'id': 'server2',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
@@ -73,6 +81,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
'reconnection_interval': 60,
|
||||
},
|
||||
'server3': {
|
||||
'server_name': 'server3',
|
||||
'id': 'server3',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
@@ -83,7 +92,10 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
},
|
||||
}
|
||||
opc = OPC(
|
||||
opc_servers=servers, logger=mock_logger, notification_handler=mock_notification_handler
|
||||
opc_servers=servers,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
await opc.init_opc()
|
||||
|
||||
@@ -97,6 +109,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
[
|
||||
call(
|
||||
opc_id='server1',
|
||||
server_name='server1',
|
||||
url='http://localhost:8080',
|
||||
logger=mock_logger,
|
||||
server_uri='opc.tcp://localhost:4840',
|
||||
@@ -105,6 +118,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
),
|
||||
]
|
||||
)
|
||||
@@ -112,6 +126,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
[
|
||||
call(
|
||||
opc_id='server2',
|
||||
server_name='server2',
|
||||
url='http://localhost:8080',
|
||||
logger=mock_logger,
|
||||
server_uri='opc.tcp://localhost:4840',
|
||||
@@ -120,6 +135,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
]
|
||||
)
|
||||
@@ -152,6 +168,7 @@ async def opc(mock_opc_repository):
|
||||
servers = {
|
||||
'server1': {
|
||||
'id': 'server1',
|
||||
'server_name': 'server1',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
'cert_path': '',
|
||||
@@ -163,9 +180,16 @@ async def opc(mock_opc_repository):
|
||||
|
||||
mock_opc_repository.return_value.write_data = AsyncMock(return_value=(True, {}))
|
||||
mock_opc_repository.return_value.connect = AsyncMock(return_value=(True, {}))
|
||||
opc = OPC(opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock())
|
||||
opc = OPC(
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
await opc.init_opc()
|
||||
opc.send_notification = MagicMock()
|
||||
opc.send_notification_async = AsyncMock()
|
||||
opc.emit_metric = AsyncMock()
|
||||
return opc
|
||||
|
||||
|
||||
@@ -219,7 +243,7 @@ async def test_write_data_failed(opc):
|
||||
)
|
||||
assert result is None
|
||||
|
||||
opc.send_notification.assert_called_once_with(
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='OPC_WRITE_DATA_ERROR_server1',
|
||||
message='Failed to write data to OPC server: Test error',
|
||||
@@ -244,7 +268,7 @@ async def test_write_data_exception(opc):
|
||||
)
|
||||
|
||||
except Exception:
|
||||
opc.send_notification.assert_called_once_with(
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='WRITE_OPC_PREDICTION_ERROR',
|
||||
message='Error writing data to OPC server: Test error',
|
||||
@@ -411,7 +435,7 @@ async def test_write_opc_data_empty_config(opc):
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_opc_data_no_validate_server(opc):
|
||||
opc.validate_server = MagicMock(return_value=False)
|
||||
opc.validate_server = AsyncMock(return_value=False)
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||
@@ -445,13 +469,14 @@ def test_process_confidence(opc, data, success, expected):
|
||||
assert result['prediction_confidence'][0] == expected
|
||||
|
||||
|
||||
def test_validate_server(opc):
|
||||
assert opc.validate_server('server1', metadata) is True
|
||||
assert opc.validate_server('server2', metadata) is False
|
||||
@mark.asyncio
|
||||
async def test_validate_server(opc):
|
||||
assert await opc.validate_server('server1', metadata) is True
|
||||
assert await opc.validate_server('server2', metadata) is False
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_shutdown(opc):
|
||||
async def test_close(opc):
|
||||
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
|
||||
await opc.shutdown()
|
||||
await opc.close()
|
||||
opc.opc_repository['server1'].disconnect.assert_called_once()
|
||||
|
||||
@@ -37,6 +37,7 @@ def storage(mock_minio_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
|
||||
@@ -44,6 +45,7 @@ def storage(mock_minio_repository):
|
||||
def test___init___not_hasattr(mock_minio_repository):
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
storage = Storage(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
@@ -61,6 +63,7 @@ def test___init___not_hasattr(mock_minio_repository):
|
||||
},
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
assert isinstance(storage, Postgres)
|
||||
|
||||
@@ -72,6 +75,7 @@ def test___init___not_hasattr(mock_minio_repository):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
|
||||
@@ -80,7 +84,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
storage.minio_repository = None
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
|
||||
metrics_controller = AsyncMock()
|
||||
storage.__init__(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
@@ -98,6 +102,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
},
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
mock_minio_repository.assert_called_once_with(
|
||||
@@ -108,6 +113,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
|
||||
@@ -130,6 +136,7 @@ def test___init___done_repository(mock_minio_repository, storage):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
mock_minio_repository.assert_not_called()
|
||||
assert storage.minio_repository is not None
|
||||
@@ -162,6 +169,7 @@ async def test_query_to_minio_success(now, dataframe, storage):
|
||||
data = [{'a': 1}, {'a': 2}, {'a': 3}]
|
||||
storage.load_custom_query = AsyncMock(return_value=data)
|
||||
now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0)
|
||||
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
|
||||
storage.minio_repository.minio_bucket = 'test'
|
||||
|
||||
result = await storage.query_to_minio({'object_prefix': 'test', **metadata})
|
||||
@@ -183,11 +191,14 @@ async def test_query_to_minio_success(now, dataframe, storage):
|
||||
@mark.asyncio
|
||||
async def test_query_to_minio_error(storage):
|
||||
storage.send_notification = MagicMock()
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
|
||||
|
||||
storage.load_custom_query = AsyncMock(side_effect=Exception('test'))
|
||||
result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'})
|
||||
assert result['success'] is False
|
||||
assert result['message'] == 'test'
|
||||
storage.send_notification.assert_called_once_with(
|
||||
storage.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='ERROR_STORING_QUERY_TO_MINIO',
|
||||
message='Error storing query to MinIO: test',
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
from botocore.utils import ClientError
|
||||
from pytest import fixture, raises
|
||||
from pytest import fixture, mark, raises
|
||||
|
||||
from laborious import metrics
|
||||
from laborious.utils.repository.minio_repository import MinioRepository
|
||||
|
||||
|
||||
@@ -17,6 +18,7 @@ def test___init___(mock_config, mock_boto3):
|
||||
minio_default_bucket='test',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert minio_repository.storage_options == {
|
||||
@@ -50,7 +52,7 @@ def test___init___(mock_config, mock_boto3):
|
||||
@patch('laborious.utils.repository.minio_repository.Config')
|
||||
@patch('laborious.utils.repository.minio_repository.boto3')
|
||||
def minio_repository(mock_boto3, mock_config):
|
||||
return MinioRepository(
|
||||
minio_repository = MinioRepository(
|
||||
minio_endpoint_url='localhost:9000',
|
||||
minio_access_key='minio',
|
||||
minio_secret_key='minio123',
|
||||
@@ -58,55 +60,98 @@ def minio_repository(mock_boto3, mock_config):
|
||||
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()
|
||||
|
||||
|
||||
def test_ensure_bucket_exists_bucket_exists(minio_repository):
|
||||
assert minio_repository.ensure_bucket_exists({}) is None
|
||||
@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
|
||||
)
|
||||
|
||||
def test_ensure_bucket_exists_bucket_not_exists_create_success(minio_repository):
|
||||
|
||||
@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 minio_repository.ensure_bucket_exists({}) is None
|
||||
assert await minio_repository.ensure_bucket_exists({}) is None
|
||||
|
||||
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
|
||||
minio_repository.s3_client.create_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()
|
||||
|
||||
|
||||
def test_ensure_bucket_exists_bucket_not_exists_create_error(minio_repository):
|
||||
minio_repository.send_notification = MagicMock()
|
||||
@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')
|
||||
|
||||
minio_repository.s3_client.head_bucket.side_effect = ClientError(
|
||||
error_response={'Error': {'Code': '404'}}, operation_name='head_bucket'
|
||||
)
|
||||
minio_repository.s3_client.create_bucket.side_effect = ClientError(
|
||||
error_response={'Error': {'Code': '404'}}, operation_name='create_bucket'
|
||||
)
|
||||
|
||||
with raises(ClientError):
|
||||
minio_repository.ensure_bucket_exists({})
|
||||
with raises(ValueError):
|
||||
await minio_repository.ensure_bucket_exists({})
|
||||
|
||||
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
|
||||
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
|
||||
@patch('laborious.utils.repository.minio_repository.BytesIO')
|
||||
def test_store_dataframe_as_parquet(mock_bytesio, minio_repository):
|
||||
async def test_store_dataframe_as_parquet_success(mock_bytesio, minio_repository):
|
||||
input_data = MagicMock()
|
||||
|
||||
minio_repository.ensure_bucket_exists = MagicMock(return_value=True)
|
||||
minio_repository.ensure_bucket_exists = AsyncMock()
|
||||
|
||||
minio_repository.store_dataframe_as_parquet(
|
||||
await minio_repository.store_dataframe_as_parquet(
|
||||
dataframe=input_data, uri='s3://test/test.parquet', object_name='test.parquet', metadata={}
|
||||
)
|
||||
|
||||
@@ -121,15 +166,53 @@ def test_store_dataframe_as_parquet(mock_bytesio, minio_repository):
|
||||
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')
|
||||
def test_get_parquet_as_dataframe(mock_read_parquet, mock_bytesio, minio_repository):
|
||||
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 = minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
|
||||
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')
|
||||
|
||||
@@ -137,3 +220,26 @@ def test_get_parquet_as_dataframe(mock_read_parquet, mock_bytesio, minio_reposit
|
||||
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()
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
from datetime import UTC, datetime
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
import mlflow as mlflow_lib
|
||||
import numpy as np
|
||||
import pytest
|
||||
from pandas import DataFrame, Timestamp
|
||||
|
||||
from laborious import metrics
|
||||
from laborious.utils.repository.model_repository import MLFlowRepository, force_memory_release
|
||||
|
||||
|
||||
@@ -38,8 +38,17 @@ def test_force_memory_release_error(gc, ctypes):
|
||||
def mlflow_repository():
|
||||
with patch('laborious.utils.repository.model_repository.mlflow'):
|
||||
repo = MLFlowRepository(
|
||||
host='http://localhost:5000', username='admin', password='admin', logger=MagicMock()
|
||||
host='http://localhost:5000',
|
||||
username='admin',
|
||||
password='admin',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
repo.emit_metric = AsyncMock()
|
||||
repo.observe_lag = AsyncMock()
|
||||
repo.send_notification = MagicMock()
|
||||
repo.send_notification_async = AsyncMock()
|
||||
return repo
|
||||
|
||||
|
||||
@@ -181,15 +190,16 @@ def test_get_model_params(mlflow, mlflow_repository):
|
||||
assert output == mlflow.get_run.return_value.data.params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
async def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = True
|
||||
|
||||
output = mlflow_repository.dowload_artifacts('test', 'path')
|
||||
output = await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
model_name='test', stage='Production'
|
||||
@@ -209,16 +219,22 @@ def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
|
||||
assert output == mlflow_repository.client.download_artifacts.return_value
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_repository):
|
||||
async def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = False
|
||||
|
||||
output = mlflow_repository.dowload_artifacts('test', 'path')
|
||||
output = await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
model_name='test', stage='Production'
|
||||
@@ -239,6 +255,26 @@ def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_re
|
||||
assert output == mlflow_repository.client.download_artifacts.return_value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
async def test_download_artifacts_error(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = True
|
||||
|
||||
mlflow_repository.client.download_artifacts.side_effect = ValueError('test')
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
def test_get_experiment_error(mlflow, mlflow_repository):
|
||||
mlflow.get_experiment_by_name.return_value = None
|
||||
|
||||
@@ -250,30 +286,61 @@ def test_get_experiment_error(mlflow, mlflow_repository):
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
def test_load_predict_model_sklearn(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'sklearn')
|
||||
def test_get_experiment_create_error(mlflow, mlflow_repository):
|
||||
mlflow.get_experiment_by_name.return_value = None
|
||||
mlflow.create_experiment.return_value = None
|
||||
mlflow.get_experiment.return_value = None
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.get_experiment('test', create_if_not_exists=True)
|
||||
|
||||
assert str(e) == 'Experiment test not found after creation, unknown reason'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_sklearn(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'sklearn')
|
||||
|
||||
assert result == mlflow.sklearn.load_model.return_value
|
||||
mlflow.sklearn.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_pyfunc(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'pyfunc')
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_pyfunc(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'pyfunc')
|
||||
assert result == mlflow.pyfunc.load_model.return_value
|
||||
mlflow.pyfunc.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_pytorch(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'pytorch')
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_pytorch(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'pytorch')
|
||||
assert result == mlflow.pytorch.load_model.return_value
|
||||
mlflow.pytorch.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_error(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_error(mlflow_repository):
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.load_predict_model('test_model', 'invalid')
|
||||
await mlflow_repository.load_predict_model('test_model', {}, 'invalid')
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
def validate_common_load_transform_model_mocks(mlflow_repository, model_name):
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
@@ -284,65 +351,96 @@ def validate_common_load_transform_model_mocks(mlflow_repository, model_name):
|
||||
)
|
||||
|
||||
|
||||
def test_load_transform_model_sklearn(mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_sklearn(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'sklearn')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'sklearn')
|
||||
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.sklearn.load_model.return_value
|
||||
mlflow.sklearn.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_pyfunc(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_pyfunc(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'pyfunc')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'pyfunc')
|
||||
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.pyfunc.load_model.return_value
|
||||
mlflow.pyfunc.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_pytorch(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_pytorch(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'pytorch')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'pytorch')
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.pytorch.load_model.return_value
|
||||
mlflow.pytorch.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_error(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_error(mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.load_transform_model('test_model', 'invalid')
|
||||
await mlflow_repository.load_transform_model('test_model', {}, 'invalid')
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
def test_download_model_invalid_model_type(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_invalid_model_type(mlflow_repository):
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.download_model('test_model', 'invalid', 'sklearn')
|
||||
await mlflow_repository.download_model('test_model', {}, 'invalid', 'sklearn')
|
||||
assert str(e) == "Invalid model_type. Use 'predict' or 'transform'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
'model_type', [('predict', 'prediction_model'), ('transform', 'data_model')]
|
||||
)
|
||||
def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
mlflow_repository.dowload_artifacts = MagicMock()
|
||||
async def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
mlflow_repository.dowload_artifacts = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', model_type[0], 'pyfunc', True)
|
||||
result = await mlflow_repository.download_model('test_model', {}, model_type[0], 'pyfunc', True)
|
||||
|
||||
mlflow_repository.dowload_artifacts.assert_called_once_with('test_model', model_type[1])
|
||||
mlflow_repository.dowload_artifacts.assert_called_once_with('test_model', {}, model_type[1])
|
||||
|
||||
mlflow.pyfunc.load_model.assert_called_once_with(
|
||||
mlflow_repository.dowload_artifacts.return_value
|
||||
@@ -354,26 +452,28 @@ def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
)
|
||||
|
||||
|
||||
def test_download_model_predict(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = MagicMock()
|
||||
mlflow_repository.load_transform_model = MagicMock()
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_predict(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = AsyncMock()
|
||||
mlflow_repository.load_transform_model = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', 'predict', 'pyfunc', False)
|
||||
result = await mlflow_repository.download_model('test_model', {}, 'predict', 'pyfunc', False)
|
||||
|
||||
mlflow_repository.load_predict_model.assert_called_once_with('test_model', 'pyfunc')
|
||||
mlflow_repository.load_predict_model.assert_called_once_with('test_model', {}, 'pyfunc')
|
||||
mlflow_repository.load_transform_model.assert_not_called()
|
||||
|
||||
assert result == (mlflow_repository.load_predict_model.return_value, None)
|
||||
|
||||
|
||||
def test_download_model_transform(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = MagicMock()
|
||||
mlflow_repository.load_transform_model = MagicMock()
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_transform(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = AsyncMock()
|
||||
mlflow_repository.load_transform_model = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', 'transform', 'pyfunc', False)
|
||||
result = await mlflow_repository.download_model('test_model', {}, 'transform', 'pyfunc', False)
|
||||
|
||||
mlflow_repository.load_predict_model.assert_not_called()
|
||||
mlflow_repository.load_transform_model.assert_called_once_with('test_model', 'pyfunc')
|
||||
mlflow_repository.load_transform_model.assert_called_once_with('test_model', {}, 'pyfunc')
|
||||
|
||||
assert result == (mlflow_repository.load_transform_model.return_value, None)
|
||||
|
||||
@@ -481,19 +581,25 @@ def test_handle_outdated_model(mlflow_repository):
|
||||
assert mlflow_repository.model_cache == {}
|
||||
|
||||
|
||||
def test_get_model_retention_0(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_retention_0(mlflow_repository):
|
||||
model = MagicMock()
|
||||
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
output = mlflow_repository.get_model('model_name', 0, 'predict', 'pyfunc')
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
output = await mlflow_repository.get_model('model_name', {}, 0, 'predict', 'pyfunc')
|
||||
|
||||
assert output == model
|
||||
mlflow_repository.download_model.assert_called_once_with(
|
||||
model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=False
|
||||
model_name='model_name',
|
||||
metadata={},
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=False,
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_cached_valid(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_cached_valid(mlflow_repository):
|
||||
mlflow_repository.check_cache_retention = MagicMock(return_value=True)
|
||||
mlflow_repository.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
@@ -503,7 +609,7 @@ def test_get_model_cached_valid(mlflow_repository):
|
||||
}
|
||||
}
|
||||
|
||||
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc')
|
||||
output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
|
||||
|
||||
assert output == mlflow_repository.handle_valid_model.return_value
|
||||
mlflow_repository.check_cache_retention.assert_called_once_with(
|
||||
@@ -517,12 +623,13 @@ def test_get_model_cached_valid(mlflow_repository):
|
||||
mlflow_repository.handle_outdated_model.assert_not_called()
|
||||
|
||||
|
||||
def test_get_model_cached_outdated(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_cached_outdated(mlflow_repository):
|
||||
mlflow_repository.check_cache_retention = MagicMock(return_value=False)
|
||||
mlflow_repository.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
cache = {
|
||||
'model_name_predict': {
|
||||
'target': 'cached_model',
|
||||
@@ -530,7 +637,7 @@ def test_get_model_cached_outdated(mlflow_repository):
|
||||
}
|
||||
mlflow_repository.model_cache = cache
|
||||
|
||||
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc')
|
||||
output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
|
||||
assert output == model
|
||||
mlflow_repository.check_cache_retention.assert_called_once_with(
|
||||
{
|
||||
@@ -544,14 +651,15 @@ def test_get_model_cached_outdated(mlflow_repository):
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_cached_not_found(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_cached_not_found(mlflow_repository):
|
||||
mlflow_repository.check_cache_retention = MagicMock(return_value=False)
|
||||
mlflow_repository.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
mlflow_repository.model_cache = {}
|
||||
model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc')
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
|
||||
assert output == model
|
||||
mlflow_repository.check_cache_retention.assert_not_called()
|
||||
mlflow_repository.handle_valid_model.assert_not_called()
|
||||
@@ -559,43 +667,53 @@ def test_get_model_cached_not_found(mlflow_repository):
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_operation_retention_0(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_operation_retention_0(force_memory_release, mlflow_repository):
|
||||
model = MagicMock()
|
||||
data = MagicMock()
|
||||
mlflow_repository.get_model = MagicMock(return_value=model)
|
||||
output = mlflow_repository.get_cached_operation('model_name', data, 'transform', 0, 'sklearn')
|
||||
mlflow_repository.get_model = AsyncMock(return_value=model)
|
||||
output = await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'transform', 0, 'sklearn', {}
|
||||
)
|
||||
assert output == model.predict.return_value
|
||||
force_memory_release.assert_called_once_with(mlflow_repository.logger)
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_predict_retention_not_0(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_predict_retention_not_0(force_memory_release, mlflow_repository):
|
||||
model = MagicMock()
|
||||
data = MagicMock()
|
||||
mlflow_repository.get_model = MagicMock(return_value=model)
|
||||
output = mlflow_repository.get_cached_operation('model_name', data, 'predict', 1, 'sklearn')
|
||||
mlflow_repository.get_model = AsyncMock(return_value=model)
|
||||
output = await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'predict', 1, 'sklearn', {}
|
||||
)
|
||||
assert output == model.predict.return_value
|
||||
force_memory_release.assert_not_called()
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_operation_invalid_operation(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_operation_invalid_operation(force_memory_release, mlflow_repository):
|
||||
data = MagicMock()
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.get_cached_operation('model_name', data, 'invalid', 0, 'sklearn')
|
||||
await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'invalid', 0, 'sklearn', {}
|
||||
)
|
||||
assert str(e) == "Invalid operation. Use 'transform' or 'predict'."
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.pd.merge')
|
||||
@patch('laborious.utils.repository.model_repository.isinstance')
|
||||
def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
@pytest.mark.asyncio
|
||||
async def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
isinstance_mock, pd_merge, mlflow_repository
|
||||
):
|
||||
isinstance_mock.return_value = False
|
||||
|
||||
data_model = MagicMock()
|
||||
prediction_model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(
|
||||
mlflow_repository.download_model = AsyncMock(
|
||||
side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
|
||||
)
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
|
||||
@@ -604,7 +722,7 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
output = mlflow_repository.fit_models(
|
||||
output = await mlflow_repository.fit_models(
|
||||
'model_name', data, 'latest_production_id', metadata['metadata'], 'sklearn', 'pyfunc', None
|
||||
)
|
||||
|
||||
@@ -612,11 +730,18 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
[
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='transform',
|
||||
flavor='sklearn',
|
||||
load_wrapper=False,
|
||||
),
|
||||
call(model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=True),
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=True,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -658,7 +783,8 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.pd.merge')
|
||||
@patch('laborious.utils.repository.model_repository.isinstance')
|
||||
def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
@pytest.mark.asyncio
|
||||
async def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
isinstance_mock, pd_merge, mlflow_repository
|
||||
):
|
||||
isinstance_mock.return_value = True
|
||||
@@ -667,7 +793,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
target_variable='feat_2',
|
||||
)
|
||||
prediction_model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(
|
||||
mlflow_repository.download_model = AsyncMock(
|
||||
side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
|
||||
)
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
|
||||
@@ -678,7 +804,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
output = mlflow_repository.fit_models(
|
||||
output = await mlflow_repository.fit_models(
|
||||
'model_name',
|
||||
data,
|
||||
'latest_production_id',
|
||||
@@ -692,11 +818,18 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
[
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='transform',
|
||||
flavor='sklearn',
|
||||
load_wrapper=False,
|
||||
),
|
||||
call(model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=True),
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=True,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -728,16 +861,27 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
}
|
||||
|
||||
|
||||
def test_log_model_sklearn(mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_sklearn(mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
mlflow_repository.log_model(model_data, 'sklearn', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'sklearn', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
mlflow.sklearn.log_model.assert_called_once_with(model_data['model'], 'prediction_model')
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
mlflow_repository.log_model(model_data, 'pyfunc', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'pyfunc', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow.pyfunc.log_model.assert_not_called()
|
||||
|
||||
@@ -747,24 +891,47 @@ def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
artifact_path='prediction_model', code_path=[path.join.return_value], to_disk=False
|
||||
)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_log_model_pytorch(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_pytorch(mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
mlflow_repository.log_model(model_data, 'pytorch', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'pytorch', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
mlflow.pytorch.log_model.assert_called_once_with(model_data['model'], 'prediction_model')
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_log_model_error(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_error(mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.log_model(model_data, 'invalid', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'invalid', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_new_experiment(
|
||||
_rmtree, path, force_memory_release, mlflow, mlflow_repository
|
||||
):
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
retrain_data = {
|
||||
@@ -782,11 +949,11 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
|
||||
|
||||
mlflow_repository.get_experiment = MagicMock()
|
||||
mlflow_repository.get_next_run_name = MagicMock()
|
||||
mlflow_repository.log_model = MagicMock()
|
||||
mlflow_repository.log_model = AsyncMock()
|
||||
path.exists.return_value = True
|
||||
path.join.return_value = './tmp/artifacts/model_name'
|
||||
|
||||
report = mlflow_repository.create_new_experiment(
|
||||
report = await mlflow_repository.create_new_experiment(
|
||||
model_name,
|
||||
data,
|
||||
retrain_data,
|
||||
@@ -843,8 +1010,60 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
|
||||
'experiment_name': mlflow_repository.get_experiment.return_value.name,
|
||||
}
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_new_experiment_error(
|
||||
_rmtree, path, force_memory_release, mlflow, mlflow_repository
|
||||
):
|
||||
mlflow.start_run.side_effect = ValueError('error')
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
retrain_data = {
|
||||
'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
}
|
||||
|
||||
mlflow_repository.get_model_params = MagicMock(
|
||||
return_value={
|
||||
'transform_flavor': 'sklearn',
|
||||
'predict_flavor': 'pyfunc',
|
||||
'target_name': 'target_name',
|
||||
}
|
||||
)
|
||||
|
||||
mlflow_repository.get_experiment = MagicMock()
|
||||
mlflow_repository.get_next_run_name = MagicMock()
|
||||
mlflow_repository.log_model = AsyncMock()
|
||||
path.exists.return_value = True
|
||||
path.join.return_value = './tmp/artifacts/model_name'
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.create_new_experiment(
|
||||
model_name,
|
||||
data,
|
||||
retrain_data,
|
||||
'latest_production_id',
|
||||
metadata['metadata'],
|
||||
'sklearn',
|
||||
'pyfunc',
|
||||
)
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
mlflow_repository.client.get_registered_model.return_value = MagicMock(
|
||||
latest_versions=[
|
||||
MagicMock(version='1'),
|
||||
@@ -852,7 +1071,9 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
MagicMock(version='3'),
|
||||
]
|
||||
)
|
||||
output = mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
output = await mlflow_repository.update_production_model_by_run_id(
|
||||
'0', 'test', metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow.register_model.assert_called_once_with(
|
||||
'runs:/0/prediction_model',
|
||||
@@ -873,21 +1094,62 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
'mlflow_run_id': '0',
|
||||
}
|
||||
|
||||
mlflow_repository.observe_lag.assert_has_calls(
|
||||
[
|
||||
call(ANY, metrics.MODEL_WRITE_LAG, ANY),
|
||||
call(ANY, metrics.MODEL_WRITE_LAG, ANY),
|
||||
]
|
||||
)
|
||||
mlflow_repository.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY),
|
||||
call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY),
|
||||
]
|
||||
)
|
||||
|
||||
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error_register_model(mlflow, mlflow_repository):
|
||||
mlflow.register_model.side_effect = ValueError('error')
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error_transition_model_version_stage(
|
||||
mlflow, mlflow_repository
|
||||
):
|
||||
mlflow_repository.client.transition_model_version_stage.side_effect = ValueError('error')
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
||||
mlflow_repository.client.get_registered_model.return_value = MagicMock(
|
||||
get_registered_model=MagicMock(return_value=MagicMock(latest_versions={}))
|
||||
)
|
||||
|
||||
try:
|
||||
mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
except Exception as e:
|
||||
assert str(e) == 'Model versions is not a list'
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
|
||||
def test_transform_success(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_success(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'model'
|
||||
model_config = {
|
||||
@@ -896,14 +1158,19 @@ def test_transform_success(mlflow_repository):
|
||||
'predict_flavor': 'pyfunc',
|
||||
}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock()
|
||||
mlflow_repository.get_cached_operation = AsyncMock(return_value=data)
|
||||
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'transform', 60, 'sklearn'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='transform',
|
||||
retention=60,
|
||||
flavor='sklearn',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
|
||||
@@ -916,7 +1183,8 @@ def test_transform_success(mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
def test_transform_error(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_error(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'model'
|
||||
model_config = {
|
||||
@@ -925,27 +1193,40 @@ def test_transform_error(mlflow_repository):
|
||||
'predict_flavor': 'pyfunc',
|
||||
}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(side_effect=Exception('error'))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error'))
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'transform', 60, 'sklearn'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='transform',
|
||||
retention=60,
|
||||
flavor='sklearn',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
|
||||
|
||||
|
||||
def test_predict_success_array(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_predict_success_array(mlflow_repository):
|
||||
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
model_name = 'model'
|
||||
mlflow_repository.get_cached_operation = MagicMock(return_value=np.array([2, 3]))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(
|
||||
return_value=DataFrame({'prediction': {'index_1': 2, 'index_2': 3}})
|
||||
)
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output['success'] is True
|
||||
@@ -955,19 +1236,25 @@ def test_predict_success_array(mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
def test_predict_success_df(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_predict_success_df(mlflow_repository):
|
||||
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
model_name = 'model'
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(
|
||||
mlflow_repository.get_cached_operation = AsyncMock(
|
||||
return_value=DataFrame({'feat_1': {'index_3': 2, 'index_4': 3}})
|
||||
)
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output['success'] is True
|
||||
@@ -977,32 +1264,41 @@ def test_predict_success_df(mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
def test_predict_error(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_predict_error(mlflow_repository):
|
||||
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
|
||||
model_name = 'model'
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(side_effect=Exception('error'))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error'))
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
|
||||
|
||||
|
||||
def test_retrain_model(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrain_model(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'test'
|
||||
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
|
||||
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.fit_models = MagicMock()
|
||||
mlflow_repository.create_new_experiment = MagicMock()
|
||||
mlflow_repository.fit_models = AsyncMock()
|
||||
mlflow_repository.create_new_experiment = AsyncMock()
|
||||
|
||||
output = mlflow_repository.retrain_model(data, model_name, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.retrain_model(
|
||||
data, model_name, model_config, metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(model_name, stage='Production')
|
||||
|
||||
@@ -1033,12 +1329,15 @@ def test_retrain_model(mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
def test_retrain_model_error(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrain_model_error(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'test'
|
||||
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
|
||||
mlflow_repository.get_model_run_id = MagicMock(side_effect=Exception('error'))
|
||||
output = mlflow_repository.retrain_model(data, model_name, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.retrain_model(
|
||||
data, model_name, model_config, metadata['metadata']
|
||||
)
|
||||
assert output == {
|
||||
'success': False,
|
||||
'experiment': None,
|
||||
@@ -1047,17 +1346,20 @@ def test_retrain_model_error(mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
def test_update_production_model(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model(mlflow_repository):
|
||||
experiment = {'run_id': '0', 'experiment_id': '0'}
|
||||
model_name = 'test'
|
||||
mlflow_repository.update_production_model_by_run_id = MagicMock()
|
||||
mlflow_repository.update_production_model_by_run_id = AsyncMock()
|
||||
mlflow_repository.update_production_model_by_run_id.return_value = {
|
||||
'model_name': 'test',
|
||||
'version': '3',
|
||||
'mlflow_run_id': '0',
|
||||
}
|
||||
|
||||
output = mlflow_repository.update_production_model(experiment, model_name, metadata['metadata'])
|
||||
output = await mlflow_repository.update_production_model(
|
||||
experiment, model_name, metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow_repository.update_production_model_by_run_id.assert_called_once_with(
|
||||
'0', 'test', metadata['metadata']
|
||||
|
||||
@@ -18,6 +18,7 @@ def mock_logger():
|
||||
def opc_repository(mock_logger):
|
||||
repository = OpcRepository(
|
||||
opc_id='test_repo',
|
||||
server_name='test_server',
|
||||
url='opc.tcp://localhost:4840',
|
||||
logger=mock_logger,
|
||||
notification_handler=Mock(),
|
||||
@@ -26,8 +27,12 @@ def opc_repository(mock_logger):
|
||||
cert_path='/path/to/cert.pem',
|
||||
private_key_path='/path/to/key.pem',
|
||||
server_cert_path='/path/to/server_cert.pem',
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
repository.disconnection_interval = 0.1
|
||||
repository.send_notification = MagicMock()
|
||||
repository.send_notification_async = AsyncMock()
|
||||
repository.emit_metric = AsyncMock()
|
||||
return repository
|
||||
|
||||
|
||||
@@ -51,6 +56,7 @@ metadata = {
|
||||
|
||||
def test_init(opc_repository):
|
||||
assert opc_repository.id == 'test_repo'
|
||||
assert opc_repository.server_name == 'test_server'
|
||||
assert opc_repository.url == 'opc.tcp://localhost:4840'
|
||||
assert opc_repository.server_uri == 'urn:test:server'
|
||||
assert opc_repository.cert_path == '/path/to/cert.pem'
|
||||
@@ -135,11 +141,13 @@ async def test_try_connect_success(opc_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_try_connect_fail(opc_repository):
|
||||
opc_repository.last_reconnection_time = None
|
||||
opc_repository.disconnect = AsyncMock()
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.client.connect.side_effect = Exception('Test error')
|
||||
|
||||
is_connected, error_data = await opc_repository.try_connect()
|
||||
|
||||
opc_repository.disconnect.assert_called_once()
|
||||
opc_repository.client.connect.assert_called_once()
|
||||
assert is_connected is False
|
||||
assert error_data['notification_id'] == f'OPC_CONNECTION_ERROR_{opc_repository.id}'
|
||||
@@ -214,7 +222,7 @@ async def test_disconnect_error(opc_repository, mock_client):
|
||||
await opc_repository.disconnect()
|
||||
|
||||
opc_repository.disconnection_fallback.assert_called_once()
|
||||
opc_repository.send_notification.assert_called_once_with(
|
||||
opc_repository.send_notification_async.assert_called_once_with(
|
||||
metadata=opc_repository.metadata,
|
||||
notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}',
|
||||
message='Failed to disconnect from OPC server in 5 attempts.',
|
||||
|
||||
@@ -11,7 +11,7 @@ image:
|
||||
# This sets the pull policy for images.
|
||||
pullPolicy: Always
|
||||
# Overrides the image tag whose default is the chart appVersion.
|
||||
tag: "1.0.1"
|
||||
tag: "1.1.0"
|
||||
|
||||
0# This is for the secrets for pulling an image from a private repository more information can be found here: https://kubernetes.io/docs/tasks/configure-pod-container/pull-image-private-registry/
|
||||
imagePullSecrets:
|
||||
@@ -151,7 +151,7 @@ env:
|
||||
- name: GITHUB_REPO_URL
|
||||
value: "git@github.com:Aignosi/sientia-dataops-laborious_temporal.git"
|
||||
- name: GITHUB_BRANCH
|
||||
value: "fix/SIENTIAPDE-1314-ajustes-nas-camadas-de-monitoramento-do-sientia"
|
||||
value: "feature/SIENTIAPDE-1325-adicionar-metricas-especificas-de-operacoes-externas"
|
||||
- name: PYTHON_APP
|
||||
value: "laborious.worker.worker"
|
||||
|
||||
@@ -182,6 +182,8 @@ env:
|
||||
|
||||
- name: OPC_ID
|
||||
value: "1"
|
||||
- name: OPC_SERVER_NAME
|
||||
value: "default_server"
|
||||
- name: OPC_URL
|
||||
value: "opc.tcp://sientia-opc-simulator-opc.sientia.svc.cluster.local:4840"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user