SIENTIAPDE-1325

Refactor monitoring and metrics integration across various components

- Removed coverage options from `pyproject.toml`.
- Updated prediction metrics in `README.md` to replace `pipeline_name` with `workflow_name`.
- Upgraded `sientia-dataops-library` dependency version in `requirements-light.txt` and `requirements.txt`.
- Enhanced metrics handling in `laborious` activities, including `Activities`, `Gates`, `MLFlow`, and `OPC`, to utilize a new `MetricsController`.
- Refactored metric emission methods to improve clarity and consistency across the codebase.
- Updated tests to reflect changes in metrics handling and ensure proper functionality.
This commit is contained in:
vitor-aignosi
2025-11-04 16:49:10 -03:00
parent 77550d49a6
commit a3da800cab
22 changed files with 1350 additions and 417 deletions

View File

@@ -607,18 +607,18 @@ The Laborious system exposes comprehensive Prometheus metrics for operational vi
### Prediction Operation Metrics ### Prediction Operation Metrics
- `laborious_predictions_written_count`: Counter for successful prediction exports - `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 - `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 - `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] - Buckets: [0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0]
### OPC Export Metrics ### OPC Export Metrics
- `laborious_prediction_opc_writing_count`: Counter for OPC server write operations - `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 - `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] - Buckets: [0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0]
### Data Quality Metrics ### Data Quality Metrics

View File

@@ -1,3 +1,4 @@
from sientia_do.observability.metrics_controller import MetricsController
from temporalio import workflow from temporalio import workflow
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
@@ -62,6 +63,8 @@ class Activities(Storage, MLFlow, Gates, OPC):
Raises: Raises:
Exception: If any parent class initialization fails Exception: If any parent class initialization fails
""" """
metrics_controller = MetricsController(logger=logger)
# Initialize parent classes # Initialize parent classes
Storage.__init__( Storage.__init__(
self, self,
@@ -75,6 +78,7 @@ class Activities(Storage, MLFlow, Gates, OPC):
minio_config=minio_config, minio_config=minio_config,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller,
) )
MLFlow.__init__( MLFlow.__init__(
@@ -86,12 +90,22 @@ class Activities(Storage, MLFlow, Gates, OPC):
minio_config=minio_config, minio_config=minio_config,
logger=logger, logger=logger,
notification_handler=notification_handler, 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__( 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): async def shutdown(self):
@@ -107,4 +121,6 @@ class Activities(Storage, MLFlow, Gates, OPC):
proper resource cleanup and prevent resource leaks. proper resource cleanup and prevent resource leaks.
""" """
Storage.close(self) Storage.close(self)
await OPC.shutdown(self) MLFlow.close(self)
Gates.close(self)
await OPC.close(self)

View File

@@ -10,7 +10,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.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 sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ, now
from laborious import metrics 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. 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 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. Initialize data quality gates with logging and notification capabilities.
@@ -93,7 +99,16 @@ class Gates(BaseActivity):
Raises: Raises:
Exception: If BaseActivity initialization fails 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') @activity.defn(name='input_gate')
async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: 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']) filter_output.append(config['policy'])
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'INTPUT_GATE_ERROR__{fil}', notification_id=f'INTPUT_GATE_ERROR__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}', message=f'Error in filter {fil}:{config}: \n {e}',
@@ -225,7 +240,7 @@ class Gates(BaseActivity):
if mlflow_response_filter_functions[fil](data, config): if mlflow_response_filter_functions[fil](data, config):
filter_output.append(config['policy']) filter_output.append(config['policy'])
comments.append(data['content']['message']) comments.append(data['content']['message'])
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}', notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}',
message=data['content']['message'], message=data['content']['message'],
@@ -235,7 +250,7 @@ class Gates(BaseActivity):
) )
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}', notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}', message=f'Error in filter {fil}:{config}: \n {e}',
@@ -305,7 +320,7 @@ class Gates(BaseActivity):
try: try:
if mlflow_content_filter_functions[fil](data, config): if mlflow_content_filter_functions[fil](data, config):
filter_output.append(config['policy']) filter_output.append(config['policy'])
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}', notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}',
message=f'Data not passed the content filter {fil}:{config}', message=f'Data not passed the content filter {fil}:{config}',
@@ -315,7 +330,7 @@ class Gates(BaseActivity):
) )
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}', notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}',
message=f'Error in filter {fil}:{config}: \n {e}', message=f'Error in filter {fil}:{config}: \n {e}',
@@ -608,39 +623,61 @@ class Gates(BaseActivity):
self.info(f'Writing metrics for model {metadata["model_name"]}', metadata) self.info(f'Writing metrics for model {metadata["model_name"]}', metadata)
metrics.PREDICTIONS_WRITTEN_COUNT.labels( await self.emit_metric(
pod_id=self.pod_id, metric_object=metrics.PREDICTIONS_WRITTEN_COUNT,
model_name=metadata['model_name'], tags={
pipeline_name=metadata['workflow_name'], 'pod_id': self.pod_id,
).inc() 'model_name': metadata['model_name'],
'workflow_name': metadata['workflow_name'],
},
)
metrics.PREDICTION_CONFIDENCE_MONITOR.labels( await self.emit_metric(
pod_id=self.pod_id, metric_object=metrics.PREDICTION_CONFIDENCE_MONITOR,
model_name=metadata['model_name'], method='set',
pipeline_name=metadata['workflow_name'], tags={
).set(prediction_confidence) 'pod_id': self.pod_id,
'model_name': metadata['model_name'],
'workflow_name': metadata['workflow_name'],
},
value=prediction_confidence,
)
metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels( await self.emit_metric(
pod_id=self.pod_id, metric_object=metrics.PREDICTION_RESPONSE_TIME_MONITOR,
model_name=metadata['model_name'], method='observe',
pipeline_name=metadata['workflow_name'], tags={
).observe(response_time) '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 server_id, tags in opc_metrics.items():
for tag, response_time in tags.items(): for tag, response_time in tags.items():
metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels( await self.emit_metric(
pod_id=self.pod_id, metric_object=metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
model_name=metadata['model_name'], method='observe',
pipeline_name=metadata['workflow_name'], tags={
opc_server_id=server_id, 'pod_id': self.pod_id,
tag=tag, 'model_name': metadata['model_name'],
).observe(response_time) 'workflow_name': metadata['workflow_name'],
metrics.PREDICTION_OPC_WRITING_COUNT.labels( 'opc_server_id': server_id,
pod_id=self.pod_id, 'tag': tag,
model_name=metadata['model_name'], },
pipeline_name=metadata['workflow_name'], value=response_time,
opc_server_id=server_id, )
tag=tag,
).inc() 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) self.info(f'Metrics written for model {metadata["model_name"]}', metadata)

View File

@@ -10,7 +10,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.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 ( from sientia_do.temporal.constants import (
DATETIME_FORMAT, DATETIME_FORMAT,
DATETIME_FORMAT_MS_WITH_TZ, DATETIME_FORMAT_MS_WITH_TZ,
@@ -22,7 +23,7 @@ with workflow.unsafe.imports_passed_through():
from laborious.utils.repository.model_repository import MLFlowRepository from laborious.utils.repository.model_repository import MLFlowRepository
class MLFlow(BaseActivity): class MLFlow(SientiaMonitoring):
""" """
MLFlow integration activities for model inference operations. MLFlow integration activities for model inference operations.
@@ -50,6 +51,7 @@ class MLFlow(BaseActivity):
mlflow_password: str, mlflow_password: str,
logger: Logger, logger: Logger,
notification_handler: NotificationHandler, notification_handler: NotificationHandler,
metrics_controller: MetricsController,
): ):
""" """
Initialize MLFlow activities with server configuration. Initialize MLFlow activities with server configuration.
@@ -65,14 +67,19 @@ class MLFlow(BaseActivity):
Raises: Raises:
Exception: If MLFlowRepository initialization fails 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_host = mlflow_host
self.mlflow_port = mlflow_port self.mlflow_port = mlflow_port
self.mlflow_username = mlflow_username self.mlflow_username = mlflow_username
self.mlflow_password = mlflow_password self.mlflow_password = mlflow_password
self.model_monitoring_repository = MLFlowRepository( 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'): if not hasattr(self, 'minio_repository'):
@@ -87,8 +94,18 @@ class MLFlow(BaseActivity):
minio_secret_key=minio_config['secret_key'], minio_secret_key=minio_config['secret_key'],
minio_region_name=minio_config['region_name'], minio_region_name=minio_config['region_name'],
minio_default_bucket=minio_config['default_bucket'], 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') @activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]: async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]:
""" """
@@ -143,7 +160,7 @@ class MLFlow(BaseActivity):
self.debug(data.head(5).to_string(), metadata) self.debug(data.head(5).to_string(), metadata)
# Request transformation from MLFlow model # 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 model_name, data, model_config, metadata
) )
@@ -208,7 +225,7 @@ class MLFlow(BaseActivity):
).dt.strftime(DATETIME_FORMAT) ).dt.strftime(DATETIME_FORMAT)
# Request prediction from MLFlow model # 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 model_name, data, model_config, metadata
) )
@@ -263,12 +280,12 @@ class MLFlow(BaseActivity):
self.info(f'Loading retrain data from Key: {object_key}', metadata) self.info(f'Loading retrain data from Key: {object_key}', metadata)
try: try:
data = self.minio_repository.get_parquet_as_dataframe( data = await self.minio_repository.get_parquet_as_dataframe(
object_key=object_key, metadata=metadata object_key=object_key, metadata=metadata
) )
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id='ERROR_LOADING_RETRAIN_DATA', notification_id='ERROR_LOADING_RETRAIN_DATA',
message=f'Error loading retrain data: {e}', message=f'Error loading retrain data: {e}',
@@ -316,13 +333,13 @@ class MLFlow(BaseActivity):
data.columns.name = None 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 data=data, model_name=model_name, model_config=model_config, metadata=metadata
) )
if not retrain_output['success']: if not retrain_output['success']:
trace = retrain_output['traceback'] trace = retrain_output['traceback']
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id='RETRAIN_MODEL_ERROR', notification_id='RETRAIN_MODEL_ERROR',
message=f'Error retraining model {model_name}: {retrain_output["message"]}', message=f'Error retraining model {model_name}: {retrain_output["message"]}',
@@ -379,7 +396,7 @@ class MLFlow(BaseActivity):
) )
try: 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 experiment=experiment, model_name=model_name, metadata=metadata
) )
@@ -388,7 +405,7 @@ class MLFlow(BaseActivity):
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id='UPDATE_PRODUCTION_MODEL_ERROR', notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
message=f'Error updating production model {model_name}: {e}', message=f'Error updating production model {model_name}: {e}',

View File

@@ -8,14 +8,15 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.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 from laborious.utils.repository.opc_repository import OpcRepository
OPC_WRITTING_ERROR_CONFIDENCE = 12 OPC_WRITTING_ERROR_CONFIDENCE = 12
class OPC(BaseActivity): class OPC(SientiaMonitoring):
""" """
OPC server integration activities for real-time data export. OPC server integration activities for real-time data export.
@@ -39,12 +40,13 @@ class OPC(BaseActivity):
opc_servers: dict[str, dict[str, Any]], opc_servers: dict[str, dict[str, Any]],
logger: Logger, logger: Logger,
notification_handler: NotificationHandler, notification_handler: NotificationHandler,
metrics_controller: MetricsController,
): ):
self.logger = logger self.logger = logger
self.notification_handler = notification_handler self.notification_handler = notification_handler
self.opc_servers = opc_servers 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_repository: dict[str, OpcRepository] = {}
self.opc_servers = opc_servers self.opc_servers = opc_servers
@@ -85,10 +87,11 @@ class OPC(BaseActivity):
server_cert_path=server['server_cert_path'], server_cert_path=server['server_cert_path'],
notification_handler=self.notification_handler, notification_handler=self.notification_handler,
reconnection_interval=server['reconnection_interval'], reconnection_interval=server['reconnection_interval'],
metrics_controller=self.metrics_controller,
) )
is_connected, error_data = await self.opc_repository[opc_id].connect() is_connected, error_data = await self.opc_repository[opc_id].connect()
if not is_connected: if not is_connected:
self.send_notification( await self.send_notification_async(
metadata={ metadata={
'model_id': '-', 'model_id': '-',
'model_name': '-', 'model_name': '-',
@@ -137,7 +140,7 @@ class OPC(BaseActivity):
tag, data, data_type, self.logger, metadata tag, data, data_type, self.logger, metadata
) )
if not is_success: if not is_success:
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=info_data['notification_id'], notification_id=info_data['notification_id'],
message=info_data['message'], message=info_data['message'],
@@ -149,7 +152,7 @@ class OPC(BaseActivity):
return info_data['response_time'] return info_data['response_time']
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id=f'WRITE_OPC_{tag_type.upper()}_ERROR', notification_id=f'WRITE_OPC_{tag_type.upper()}_ERROR',
message=f'Error writing data to OPC server: {e}', message=f'Error writing data to OPC server: {e}',
@@ -159,7 +162,7 @@ class OPC(BaseActivity):
) )
raise e 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. Validate that an OPC server is available and configured for write operations.
@@ -182,7 +185,7 @@ class OPC(BaseActivity):
""" """
if self.opc_repository.get(server_id) is None: if self.opc_repository.get(server_id) is None:
message = f'OPC server {server_id} not found to perform write operation.' message = f'OPC server {server_id} not found to perform write operation.'
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id='OPC_SERVER_NOT_FOUND', notification_id='OPC_SERVER_NOT_FOUND',
message=message, message=message,
@@ -299,7 +302,7 @@ class OPC(BaseActivity):
metrics: dict[str, dict[str, float | None]] = {} metrics: dict[str, dict[str, float | None]] = {}
for server_id, config in opc_output_config.items(): 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 success = False
continue continue
@@ -358,7 +361,7 @@ class OPC(BaseActivity):
return data.to_dict() return data.to_dict()
async def shutdown(self): async def close(self):
""" """
Gracefully shutdown all OPC server connections and cleanup resources. Gracefully shutdown all OPC server connections and cleanup resources.

View File

@@ -9,6 +9,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.temporal.activities.postgres import Postgres from sientia_do.temporal.activities.postgres import Postgres
from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now from sientia_do.temporal.constants import DATETIME_FORMAT_FILENAME, now
@@ -33,6 +34,7 @@ class Storage(Postgres):
minio_config: dict[str, Any], minio_config: dict[str, Any],
logger: Logger, logger: Logger,
notification_handler: NotificationHandler, notification_handler: NotificationHandler,
metrics_controller: MetricsController,
): ):
super().__init__( super().__init__(
host=host, host=host,
@@ -44,6 +46,7 @@ class Storage(Postgres):
max_connections=max_connections, max_connections=max_connections,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller,
) )
if not hasattr(self, 'minio_repository'): if not hasattr(self, 'minio_repository'):
@@ -58,6 +61,7 @@ class Storage(Postgres):
minio_secret_key=minio_config['secret_key'], minio_secret_key=minio_config['secret_key'],
minio_region_name=minio_config['region_name'], minio_region_name=minio_config['region_name'],
minio_default_bucket=minio_config['default_bucket'], minio_default_bucket=minio_config['default_bucket'],
metrics_controller=metrics_controller,
) )
@activity.defn(name='query_to_minio') @activity.defn(name='query_to_minio')
@@ -95,14 +99,14 @@ class Storage(Postgres):
data = pd.DataFrame(data) data = pd.DataFrame(data)
# Write parquet to memory and upload via persistent client # 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 dataframe=data, uri=uri, object_name=object_name, metadata=metadata
) )
return {'success': True, 'object_key': object_name, 'uri': uri} return {'success': True, 'object_key': object_name, 'uri': uri}
except Exception as e: except Exception as e:
trace = traceback.format_exc() trace = traceback.format_exc()
self.send_notification( await self.send_notification_async(
metadata=metadata, metadata=metadata,
notification_id='ERROR_STORING_QUERY_TO_MINIO', notification_id='ERROR_STORING_QUERY_TO_MINIO',
message=f'Error storing query to MinIO: {e}', message=f'Error storing query to MinIO: {e}',

View File

@@ -19,11 +19,12 @@ Key Metric Categories:
Metric Labels: Metric Labels:
- pod_id: Kubernetes pod identifier for multi-instance deployments - pod_id: Kubernetes pod identifier for multi-instance deployments
- model_name: Name of the ML model being used - 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 - opc_server_id: Identifier for OPC server operations
""" """
from prometheus_client import Counter, Gauge, Histogram from prometheus_client import Counter, Gauge, Histogram
from sientia_do.observability.metrics import CORE_LABELS as SIENTIA_CORE_LABELS
# Application health metric # Application health metric
APP_UP = Gauge( APP_UP = Gauge(
@@ -33,7 +34,7 @@ APP_UP = Gauge(
) )
# Core labels used across multiple metrics # 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 # Prediction operation metrics
PREDICTIONS_WRITTEN_COUNT = Counter( PREDICTIONS_WRITTEN_COUNT = Counter(
@@ -49,7 +50,7 @@ PREDICTION_CONFIDENCE_MONITOR = Gauge(
CORE_LABELS, CORE_LABELS,
) )
# Performance monitoring metrics # Prediction total response time
PREDICTION_RESPONSE_TIME_MONITOR = Histogram( PREDICTION_RESPONSE_TIME_MONITOR = Histogram(
'laborious_prediction_response_time_monitor', 'laborious_prediction_response_time_monitor',
'Current response time of each prediction', '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], 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( PREDICTION_OPC_WRITING_COUNT = Counter(
'laborious_prediction_opc_writing_count', 'laborious_prediction_opc_writing_count',
'Number of predictions written to the OPC server', 'Number of predictions written to the OPC server',
@@ -70,3 +113,49 @@ PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR = Histogram(
[*CORE_LABELS, 'opc_server_id', 'tag'], [*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], buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0, 10.0],
) )
OPC_CONNECTION_STATUS = Gauge(
'laborious_opc_connection_status',
'Connection status with the OPC server (1=connected, 0=disconnected)',
['pod_id', 'opc_server_id'],
)
# ================== 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,
)

View File

@@ -6,6 +6,7 @@ object storage using boto3. It supports creating buckets on demand and
storing/loading pandas DataFrames in Parquet format. storing/loading pandas DataFrames in Parquet format.
""" """
import time
from io import BytesIO from io import BytesIO
from typing import Any from typing import Any
@@ -15,9 +16,13 @@ from botocore.exceptions import ClientError
from pandas import DataFrame, read_parquet from pandas import DataFrame, read_parquet
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.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. Repository for interacting with a MinIO (S3-compatible) object storage.
@@ -43,6 +48,7 @@ class MinioRepository:
minio_default_bucket: str, minio_default_bucket: str,
logger: Logger, logger: Logger,
notification_handler: NotificationHandler, notification_handler: NotificationHandler,
metrics_controller: MetricsController,
): ):
"""Initialize the repository and S3 client. """Initialize the repository and S3 client.
@@ -55,6 +61,7 @@ class MinioRepository:
logger (Logger): Logger instance for structured logs. logger (Logger): Logger instance for structured logs.
notification_handler (NotificationHandler): Notification handler. notification_handler (NotificationHandler): Notification handler.
""" """
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
# MinIO settings shared with pandas s3fs # MinIO settings shared with pandas s3fs
self.storage_options = { self.storage_options = {
'key': minio_access_key, 'key': minio_access_key,
@@ -85,28 +92,58 @@ class MinioRepository:
), ),
) )
self.logger = logger
self.notification_handler = notification_handler
def close(self): def close(self):
"""Close the underlying S3 client.""" """Close the underlying S3 client."""
self.s3_client.close() 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. """Ensure the default bucket exists; create it if missing.
Args: Args:
metadata (dict[str, Any]): Metadata used for structured logging. 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: try:
self.logger.custom_info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
self.s3_client.head_bucket(Bucket=self.minio_bucket) self.s3_client.head_bucket(Bucket=self.minio_bucket)
except ClientError: except ClientError:
self.logger.custom_info(f"Creating bucket '{self.minio_bucket}'", metadata) await self.create_bucket(metadata)
self.s3_client.create_bucket(Bucket=self.minio_bucket)
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] self, dataframe: DataFrame, uri: str, object_name: str, metadata: dict[str, Any]
): ):
"""Persist a DataFrame as a Parquet object in the default bucket. """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). object_name (str): Object key (path/key within the bucket).
metadata (dict[str, Any]): Metadata used for structured logging. 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() buffer = BytesIO()
dataframe.to_parquet(buffer, engine='pyarrow', index=True) dataframe.to_parquet(buffer, engine='pyarrow', index=True)
buffer.seek(0) 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. """Load a Parquet object from the default bucket into a DataFrame.
Args: Args:
@@ -138,9 +193,22 @@ class MinioRepository:
Returns: Returns:
DataFrame: Loaded DataFrame. 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)
response = self.s3_client.get_object(Bucket=self.minio_bucket, Key=object_key) 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 # Read the content into a BytesIO buffer to support seek operations
buffer = BytesIO(response['Body'].read()) buffer = BytesIO(response['Body'].read())

View File

@@ -17,6 +17,7 @@ Capabilities:
import ctypes import ctypes
import gc import gc
import threading import threading
import time
import traceback import traceback
from datetime import datetime, timedelta from datetime import datetime, timedelta
from os import environ, makedirs, path from os import environ, makedirs, path
@@ -27,9 +28,14 @@ import mlflow
import pandas as pd import pandas as pd
from mlflow.entities import Experiment from mlflow.entities import Experiment
from numpy import ndarray from numpy import ndarray
from sientia_do.notifications.handlers import NotificationHandler
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
from laborious import metrics
ARTIFACTS_PATH = './tmp/artifacts' ARTIFACTS_PATH = './tmp/artifacts'
TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl' TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl'
PREDICTION_COMPRESSED_PATH = 'artifacts/stacking_model.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}') logger.info(f'Memory release failed: {e}')
class MLFlowRepository: class MLFlowRepository(SientiaMonitoring):
def __init__(self, host: str, username: str, password: str, logger: Logger): def __init__(
self,
host: str,
username: str,
password: str,
logger: Logger,
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
"""Initialize MLflow client and base state. """Initialize MLflow client and base state.
Args: Args:
@@ -67,6 +81,7 @@ class MLFlowRepository:
logger (Logger): Logger instance. logger (Logger): Logger instance.
""" """
# set tracking uri # set tracking uri
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
mlflow.set_tracking_uri(host) mlflow.set_tracking_uri(host)
environ['MLFLOW_TRACKING_USERNAME'] = username environ['MLFLOW_TRACKING_USERNAME'] = username
@@ -204,12 +219,15 @@ class MLFlowRepository:
Functions related to download and load models 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. Download artifacts from the latest production run of a model.
Args: Args:
model_name (str): Registered model name. 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. artifact_path (str): Relative path to artifacts within the run.
Returns: Returns:
@@ -227,14 +245,29 @@ class MLFlowRepository:
self.logger.info(f'Downloading artifacts from {run_id} to {output_dir}') self.logger.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. Load a predictive model from the MLflow Model Registry.
Args: Args:
model_name (str): The name of the model to download from the registry. 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') flavor (str): Model flavor ('pyfunc', 'sklearn', 'pytorch')
artifact_path (str | None): Path to compressed artifacts if model is compressed artifact_path (str | None): Path to compressed artifacts if model is compressed
@@ -247,18 +280,29 @@ class MLFlowRepository:
""" """
model_uri = f'models:/{model_name}/production' model_uri = f'models:/{model_name}/production'
self.logger.info(f'Loading prediction model {model_name} from {model_uri}') self.logger.info(f'Loading prediction model {model_name} from {model_uri}')
if flavor == 'pyfunc':
model = mlflow.pyfunc.load_model(model_uri)
elif flavor == 'sklearn':
model = mlflow.sklearn.load_model(model_uri)
elif flavor == 'pytorch':
model = mlflow.pytorch.load_model(model_uri)
else:
raise ValueError(INVALID_FLAVOR_MESSAGE)
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':
model = mlflow.sklearn.load_model(model_uri)
elif flavor == 'pytorch':
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 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. Load the latest Production version of a transformation model.
@@ -267,6 +311,7 @@ class MLFlowRepository:
Args: Args:
model_name (str): The name of the model to download. 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') flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
artifact_path (str | None): Path to compressed artifacts if model is compressed artifact_path (str | None): Path to compressed artifacts if model is compressed
@@ -282,24 +327,41 @@ class MLFlowRepository:
model_uri = self.get_model_uri(latest_production_id, prediction=False) model_uri = self.get_model_uri(latest_production_id, prediction=False)
self.logger.info(f'Loading data model {model_name} from {model_uri}') self.logger.info(f'Loading data model {model_name} from {model_uri}')
if flavor == 'sklearn':
model = mlflow.sklearn.load_model(model_uri) core_labels = self.get_core_labels(metadata, operation_type='load_transform_model')
elif flavor == 'pyfunc': start_time = time.time()
model = mlflow.pyfunc.load_model(model_uri) try:
elif flavor == 'pytorch': if flavor == 'sklearn':
model = mlflow.pytorch.load_model(model_uri) model = mlflow.sklearn.load_model(model_uri)
else: elif flavor == 'pyfunc':
raise ValueError(INVALID_FLAVOR_MESSAGE) model = mlflow.pyfunc.load_model(model_uri)
elif flavor == 'pytorch':
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 return model
def download_model( async def download_model(
self, model_name: str, model_type: str, flavor: str, load_wrapper: bool = False self,
model_name: str,
metadata: dict[str, Any],
model_type: str,
flavor: str,
load_wrapper: bool = False,
) -> tuple[Any, str | None]: ) -> tuple[Any, str | None]:
""" """
Download model based on type ("predict" or "transform"). Download model based on type ("predict" or "transform").
Args: Args:
model_name (str): Name of the model to download 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') model_type (str): Type of model ('predict' or 'transform')
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch') flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
load_wrapper (bool): Whether to load wrapper load_wrapper (bool): Whether to load wrapper
@@ -324,7 +386,7 @@ class MLFlowRepository:
target = 'prediction_model' if model_type == 'predict' else 'data_model' 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.logger.info(
f'Model with type {model_type} and name {model_name} is compressed, loading from {artifact_path}' f'Model with type {model_type} and name {model_name} is compressed, loading from {artifact_path}'
@@ -334,10 +396,10 @@ class MLFlowRepository:
model = raw_model._model_impl.python_model model = raw_model._model_impl.python_model
else: else:
if model_type == 'predict': if model_type == 'predict':
model = self.load_predict_model(model_name, flavor) model = await self.load_predict_model(model_name, metadata, flavor)
else: else:
model = self.load_transform_model(model_name, flavor) model = await self.load_transform_model(model_name, metadata, flavor)
return model, artifact_path return model, artifact_path
@@ -446,12 +508,20 @@ class MLFlowRepository:
del self.model_cache[model_key]['target'] del self.model_cache[model_key]['target']
del self.model_cache[model_key] 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. Retrieve a model with caching support based on retention policy.
Args: Args:
model_name (str): Name of the model to retrieve 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). retention (int): Cache retention time in minutes (0 = no cache).
model_type (str): Type of model ('predict' or 'transform') model_type (str): Type of model ('predict' or 'transform')
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch') flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
@@ -461,8 +531,12 @@ class MLFlowRepository:
""" """
# Retention is 0, download a new model # Retention is 0, download a new model
if retention <= 0: if retention <= 0:
model, _artifact_path = self.download_model( model, _artifact_path = await self.download_model(
model_name=model_name, model_type=model_type, flavor=flavor, load_wrapper=False model_name=model_name,
metadata=metadata,
model_type=model_type,
flavor=flavor,
load_wrapper=False,
) )
return model return model
@@ -485,8 +559,12 @@ class MLFlowRepository:
) )
# Donwload new model (without lock to avoid blocking other threads) # Donwload new model (without lock to avoid blocking other threads)
model, _artifact_path = self.download_model( model, _artifact_path = await self.download_model(
model_name=model_name, model_type=model_type, flavor=flavor, load_wrapper=False model_name=model_name,
metadata=metadata,
model_type=model_type,
flavor=flavor,
load_wrapper=False,
) )
# Update cache with lock # Update cache with lock
@@ -497,27 +575,35 @@ class MLFlowRepository:
return model return model
@overload @overload
def get_cached_operation( async def get_cached_operation(
self, self,
model_name: str, model_name: str,
data: pd.DataFrame, data: pd.DataFrame,
operation: Literal['transform'], operation: Literal['transform'],
retention: int, retention: int,
flavor: str, flavor: str,
metadata: dict[str, Any],
) -> pd.DataFrame: ... ) -> pd.DataFrame: ...
@overload @overload
def get_cached_operation( async def get_cached_operation(
self, self,
model_name: str, model_name: str,
data: pd.DataFrame, data: pd.DataFrame,
operation: Literal['predict'], operation: Literal['predict'],
retention: int, retention: int,
flavor: str, flavor: str,
metadata: dict[str, Any],
) -> pd.DataFrame | ndarray: ... ) -> pd.DataFrame | ndarray: ...
def get_cached_operation( async def get_cached_operation(
self, model_name: str, data: pd.DataFrame, operation: str, retention: int, flavor: str self,
model_name: str,
data: pd.DataFrame,
operation: str,
retention: int,
flavor: str,
metadata: dict[str, Any],
) -> pd.DataFrame | ndarray: ) -> pd.DataFrame | ndarray:
""" """
Execute a cached operation using the requested model. Execute a cached operation using the requested model.
@@ -527,15 +613,19 @@ class MLFlowRepository:
data (pd.DataFrame): Input data. data (pd.DataFrame): Input data.
retention (int): Cache retention in minutes. retention (int): Cache retention in minutes.
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch'). flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch').
metadata (dict[str, Any]): Metadata used for structured logging.
Returns: Returns:
pd.DataFrame | ndarray: Operation result. pd.DataFrame | ndarray: Operation result.
""" """
if operation not in ['transform', 'predict']: if operation not in ['transform', 'predict']:
raise ValueError("Invalid operation. Use 'transform' or 'predict'.") raise ValueError("Invalid operation. Use 'transform' or 'predict'.")
model = self.get_model( model = await self.get_model(
model_name=model_name, retention=retention, model_type=operation, flavor=flavor model_name=model_name,
metadata=metadata,
retention=retention,
model_type=operation,
flavor=flavor,
) )
prediction = model.predict(data) prediction = model.predict(data)
@@ -552,7 +642,7 @@ class MLFlowRepository:
Functions related to model retraining Functions related to model retraining
""" """
def fit_models( async def fit_models(
self, self,
model_name: str, model_name: str,
data: pd.DataFrame, data: pd.DataFrame,
@@ -601,8 +691,9 @@ class MLFlowRepository:
load_transform_wrapper = transform_flavor == 'pyfunc' 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, model_name=model_name,
metadata=metadata,
model_type='transform', model_type='transform',
flavor=transform_flavor, flavor=transform_flavor,
load_wrapper=load_transform_wrapper, load_wrapper=load_transform_wrapper,
@@ -612,8 +703,9 @@ class MLFlowRepository:
load_predict_wrapper = predict_flavor == 'pyfunc' 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, model_name=model_name,
metadata=metadata,
model_type='predict', model_type='predict',
flavor=predict_flavor, flavor=predict_flavor,
load_wrapper=load_predict_wrapper, load_wrapper=load_predict_wrapper,
@@ -688,7 +780,7 @@ class MLFlowRepository:
} }
return retrain_data 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. """Log a model into the active MLflow run.
Args: Args:
@@ -700,22 +792,34 @@ class MLFlowRepository:
model = model_data['model'] model = model_data['model']
self.logger.custom_debug(f'Logging {model_type} model to {model_type}', metadata) self.logger.custom_debug(f'Logging {model_type} model to {model_type}', metadata)
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) core_labels = self.get_core_labels(metadata, operation_type='log_model')
start_time = time.time()
model.store_model(artifact_path=model_type, code_path=code_path, to_disk=False) 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('Model uploaded successfully', metadata) self.logger.custom_debug(f'Code path: {code_path}', metadata)
elif flavor == 'pytorch':
mlflow.pytorch.log_model(model, model_type)
else:
raise ValueError(INVALID_FLAVOR_MESSAGE)
def create_new_experiment( model.store_model(artifact_path=model_type, code_path=code_path, to_disk=False)
self.logger.custom_debug('Model uploaded successfully', metadata)
elif flavor == 'pytorch':
mlflow.pytorch.log_model(model, model_type)
else:
raise ValueError(INVALID_FLAVOR_MESSAGE)
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, self,
model_name: str, model_name: str,
data: pd.DataFrame, data: pd.DataFrame,
@@ -784,30 +888,40 @@ class MLFlowRepository:
metadata, metadata,
) )
with mlflow.start_run( core_labels = self.get_core_labels(metadata, operation_type='create_new_experiment')
experiment_id=experiment.experiment_id, start_time = time.time()
run_name=current_run_name, try:
description=experiment_description, with mlflow.start_run(
) as _run: experiment_id=experiment.experiment_id,
run_id = _run.info.run_id run_name=current_run_name,
self.logger.custom_info('Logging data model', metadata) description=experiment_description,
# dynamic parameters, including model itself ) as _run:
self.log_model(data_model, transform_flavor, 'data_model', metadata) run_id = _run.info.run_id
self.logger.custom_info('Logging data model', metadata)
# dynamic parameters, including model itself
await self.log_model(data_model, transform_flavor, 'data_model', metadata)
# dynamic parameters, including model itself # dynamic parameters, including model itself
self.logger.custom_info('Logging prediction model', metadata) self.logger.custom_info('Logging prediction model', metadata)
self.log_model(prediction_model, predict_flavor, '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.logger.custom_info(f'Model logged successfully for {model_name}', metadata)
self.logger.custom_info(f'Logging remaining parameters for {model_name}', metadata) self.logger.custom_info(f'Logging remaining parameters for {model_name}', metadata)
# update transfomation model # update transfomation model
# fixed parameters # fixed parameters
mlflow.log_params(retrain_params) mlflow.log_params(retrain_params)
# log the data raw # log the data raw
mlflow.log_artifact(data_path) mlflow.log_artifact(data_path)
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.logger.custom_info('Deleting model from filesystem', metadata) self.logger.custom_info('Deleting model from filesystem', metadata)
if path.exists(model_temp_path): if path.exists(model_temp_path):
@@ -829,7 +943,7 @@ class MLFlowRepository:
'experiment_name': experiment.name, '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 self, run_id: str, model_name: str, metadata: dict
) -> dict: ) -> dict:
""" """
@@ -864,7 +978,17 @@ class MLFlowRepository:
# Registrar o modelo # Registrar o modelo
# Aqui estamos assumindo que você já tem um modelo salvo, caso contrário você precisará treiná-lo e salvá-lo primeiro. # 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. # Se o modelo já está registrado, você pode usar o método register_model() ou pyfunc.load_model() para isso.
mlflow.register_model(f'runs:/{run_id}/prediction_model', model_name)
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 # Obter a versão mais recente registrada do modelo
model_versions = self.client.get_registered_model(model_name).latest_versions model_versions = self.client.get_registered_model(model_name).latest_versions
@@ -875,9 +999,23 @@ class MLFlowRepository:
max_version = max(model_versions, key=lambda x: int(x.version)).version 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' # Mover a versão mais recente do modelo para o estágio de 'Production'
self.client.transition_model_version_stage( core_labels = self.get_core_labels(
name=model_name, version=max_version, stage='Production', archive_existing_versions=True 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} return {'model_name': model_name, 'version': max_version, 'mlflow_run_id': run_id}
@@ -885,7 +1023,9 @@ class MLFlowRepository:
Functions that provide the interface to model operations 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. Transform data using a cached transformation model.
@@ -930,8 +1070,13 @@ class MLFlowRepository:
flavor = model_config.get('transform_flavor', 'sklearn') flavor = model_config.get('transform_flavor', 'sklearn')
try: try:
transformed_data: pd.DataFrame = self.get_cached_operation( transformed_data: pd.DataFrame = await self.get_cached_operation(
model_name, data, 'transform', model_retention, flavor model_name=model_name,
data=data,
operation='transform',
retention=model_retention,
flavor=flavor,
metadata=metadata,
) )
self.logger.custom_debug( self.logger.custom_debug(
@@ -952,7 +1097,9 @@ class MLFlowRepository:
'content': {'message': str(e), 'traceback': traceback.format_exc()}, '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. Generate predictions using a cached prediction model.
@@ -1007,8 +1154,13 @@ class MLFlowRepository:
# data.to_csv( # data.to_csv(
# f"tmp/treated_data_{model_name}.csv", index=True) # f"tmp/treated_data_{model_name}.csv", index=True)
predict_data = self.get_cached_operation( predict_data = await self.get_cached_operation(
model_name, data, 'predict', model_retention, flavor model_name=model_name,
data=data,
operation='predict',
retention=model_retention,
flavor=flavor,
metadata=metadata,
) )
end_time = datetime.now() end_time = datetime.now()
@@ -1038,7 +1190,7 @@ class MLFlowRepository:
'content': {'message': str(e), 'traceback': traceback.format_exc()}, '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 self, data: pd.DataFrame, model_name: str, model_config: dict, metadata: dict
) -> dict[str, Any]: ) -> dict[str, Any]:
""" """
@@ -1099,7 +1251,7 @@ class MLFlowRepository:
try: try:
latest_production_id = self.get_model_run_id(model_name, stage='Production') latest_production_id = self.get_model_run_id(model_name, stage='Production')
self.logger.custom_info('Creating model experiment environment', metadata) self.logger.custom_info('Creating model experiment environment', metadata)
retrain_data = self.fit_models( retrain_data = await self.fit_models(
model_name=model_name, model_name=model_name,
data=data, data=data,
transform_flavor=transform_flavor, transform_flavor=transform_flavor,
@@ -1113,7 +1265,7 @@ class MLFlowRepository:
) )
self.logger.custom_info('Saving model retrain', metadata) self.logger.custom_info('Saving model retrain', metadata)
experiment = self.create_new_experiment( experiment = await self.create_new_experiment(
model_name=model_name, model_name=model_name,
data=data, data=data,
retrain_data=retrain_data, retrain_data=retrain_data,
@@ -1141,7 +1293,7 @@ class MLFlowRepository:
'traceback': traceback.format_exc(), 'traceback': traceback.format_exc(),
} }
def update_production_model( async def update_production_model(
self, experiment: dict[str, Any], model_name: str, metadata: dict self, experiment: dict[str, Any], model_name: str, metadata: dict
) -> dict: ) -> dict:
""" """
@@ -1190,7 +1342,7 @@ class MLFlowRepository:
""" """
run_id = experiment['run_id'] run_id = experiment['run_id']
experiment_id = experiment['experiment_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 metadata_result['mlflow_experiment_id'] = experiment_id

View File

@@ -8,11 +8,14 @@ from typing import Any
from asyncua import Client from asyncua import Client
from asyncua.crypto.security_policies import SecurityPolicyBasic256 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.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
from sientia_do.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 = { data_type_map = {
'float': { 'float': {
@@ -38,13 +41,14 @@ data_type_map = {
} }
class OpcRepository(BaseActivity): class OpcRepository(SientiaMonitoring):
def __init__( def __init__(
self, self,
opc_id: str, opc_id: str,
url: str, url: str,
logger: Logger, logger: Logger,
notification_handler: NotificationHandler, notification_handler: NotificationHandler,
metrics_controller: MetricsController,
reconnection_interval: int = 60, reconnection_interval: int = 60,
server_uri: str | None = None, server_uri: str | None = None,
cert_path: str | None = None, cert_path: str | None = None,
@@ -61,10 +65,11 @@ class OpcRepository(BaseActivity):
self.error_count = 0 self.error_count = 0
self.reconnection_interval = reconnection_interval self.reconnection_interval = reconnection_interval
self.last_reconnection_time: None | datetime = None self.last_reconnection_time: None | datetime = None
self.disconnection_interval = 10.0
self.notification_handler = notification_handler self.notification_handler = notification_handler
self.client: None | Client = None 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 = { self.metadata = {
'model_name': '-', 'model_name': '-',
@@ -164,6 +169,17 @@ class OpcRepository(BaseActivity):
'level': NotificationLevel.ERROR, 'level': NotificationLevel.ERROR,
} }
await self.client.connect() await self.client.connect()
await self.emit_metric(
metric_object=metrics.OPC_CONNECTION_STATUS,
method='set',
tags={
'pod_id': self.pod_id,
'opc_server_id': self.id,
},
value=1,
)
return True, {} return True, {}
except Exception as e: except Exception as e:
self.disconnect() self.disconnect()
@@ -202,7 +218,7 @@ class OpcRepository(BaseActivity):
'traceback': traceback.format_exc(), 'traceback': traceback.format_exc(),
} }
) )
await asyncio.sleep(0.1 * i) await asyncio.sleep(self.disconnection_interval * i)
return error_stack return error_stack
async def disconnect(self): async def disconnect(self):
@@ -218,7 +234,7 @@ class OpcRepository(BaseActivity):
errors = await self.disconnection_fallback() errors = await self.disconnection_fallback()
if errors: if errors:
self.send_notification( await self.send_notification_async(
metadata=self.metadata, metadata=self.metadata,
notification_id=f'OPC_DISCONNECTION_ERROR_{self.id}', notification_id=f'OPC_DISCONNECTION_ERROR_{self.id}',
message='Failed to disconnect from OPC server in 5 attempts.', message='Failed to disconnect from OPC server in 5 attempts.',
@@ -228,6 +244,15 @@ class OpcRepository(BaseActivity):
) )
else: else:
self.logger.warning(f'Disconnected from OPC server {self.id} successfully') 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,
'opc_server_id': self.id,
},
value=0,
)
self.client = None self.client = None
@@ -382,12 +407,12 @@ class OpcRepository(BaseActivity):
data = data_type_map[data_type]['converter'](value) data = data_type_map[data_type]['converter'](value)
logger.custom_info(f'Writing {data} - {type(data)} to {node}', metadata) logger.custom_info(f'Writing {data} - {type(data)} to {node}', metadata)
now = datetime.now() # now = datetime.now()
ua_data = DataValue( ua_data = DataValue(
Variant(data, data_type_map[data_type]['opc_type']), Variant(data, data_type_map[data_type]['opc_type']),
SourceTimestamp=DateTime( # SourceTimestamp=DateTime(
now.year, now.month, now.day, now.hour, now.minute, now.second, now.microsecond # now.year, now.month, now.day, now.hour, now.minute, now.second, now.microsecond
), # ),
) )
try: try:

View File

@@ -114,10 +114,6 @@ python_functions = ["test_*"]
addopts = [ addopts = [
"-v", "-v",
"--strict-markers", "--strict-markers",
"--cov=model_manager",
"--cov-report=term-missing",
"--cov-report=html",
"--cov-report=xml",
] ]
markers = [ markers = [
"asyncio: marks tests as async", "asyncio: marks tests as async",

View File

@@ -3,7 +3,7 @@ psycopg2-binary
sqlalchemy sqlalchemy
asyncua asyncua
redis redis
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.4.7 git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.5.2
prometheus-client prometheus-client
botocore botocore
boto3 boto3

View File

@@ -3,7 +3,7 @@ psycopg2-binary
sqlalchemy sqlalchemy
asyncua asyncua
redis redis
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.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 git+ssh://git@github.com/Aignosi/sientia-mlops-library.git@0.39.0
prometheus-client prometheus-client
botocore botocore

View File

@@ -453,11 +453,40 @@
}, },
{ {
"cell_type": "code", "cell_type": "code",
"execution_count": null, "execution_count": 6,
"id": "1fbb3788", "id": "1fbb3788",
"metadata": {}, "metadata": {},
"outputs": [], "outputs": [
"source": [] {
"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": { "metadata": {

View File

@@ -1,4 +1,4 @@
from unittest.mock import ANY, MagicMock, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
from pytest import mark from pytest import mark
@@ -13,7 +13,10 @@ from laborious.activities.storage import Storage
@patch('laborious.activities.activities.MLFlow.__init__') @patch('laborious.activities.activities.MLFlow.__init__')
@patch('laborious.activities.activities.OPC.__init__') @patch('laborious.activities.activities.OPC.__init__')
@patch('laborious.activities.activities.Gates.__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 = { postgres_config = {
'host': 'localhost', 'host': 'localhost',
'port': 5432, 'port': 5432,
@@ -70,6 +73,7 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage
minio_config=minio_config, minio_config=minio_config,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
) )
mock_mlflow_init.assert_called_once_with( 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, minio_config=minio_config,
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
) )
mock_opc_init.assert_called_once_with( 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( 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 @mark.asyncio
@patch('laborious.activities.activities.Storage', return_value=MagicMock()) @patch('laborious.activities.activities.Storage')
@patch('laborious.activities.activities.MLFlow', return_value=MagicMock()) @patch('laborious.activities.activities.MLFlow')
@patch('laborious.activities.activities.OPC', return_value=MagicMock()) @patch('laborious.activities.activities.OPC')
async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init): @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 = { postgres_config = {
'host': 'localhost', 'host': 'localhost',
'port': 5432, 'port': 5432,
@@ -136,5 +150,7 @@ async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init):
) )
await activities.shutdown() 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_storage_init.close.assert_called_once()
mock_mlflow_init.close.assert_called_once()
mock_gates_init.close.assert_called_once()

View File

@@ -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 pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
@@ -11,6 +11,7 @@ def gates_activity():
gates = Gates( gates = Gates(
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
gates.error = MagicMock() gates.error = MagicMock()
gates.debug = MagicMock() gates.debug = MagicMock()
@@ -18,6 +19,8 @@ def gates_activity():
gates.warning = MagicMock() gates.warning = MagicMock()
gates.critical = MagicMock() gates.critical = MagicMock()
gates.send_notification = MagicMock() gates.send_notification = MagicMock()
gates.send_notification_async = AsyncMock()
gates.emit_metric = AsyncMock()
return gates return gates
@@ -71,7 +74,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
# Assert # Assert
assert result == (None, 0, '') assert result == (None, 0, '')
gates_activity.send_notification.assert_called_once_with( gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'], metadata=metadata['metadata'],
notification_id='INTPUT_GATE_ERROR__EMPTY_DATA', notification_id='INTPUT_GATE_ERROR__EMPTY_DATA',
message="Error in filter EMPTY_DATA:{'policy': 'STOP', 'config': {}}: \n Test error", 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
assert result == (None, 0, '') assert result == (None, 0, '')
gates_activity.send_notification.assert_called_once_with( gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'], metadata=metadata['metadata'],
notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER', notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER',
message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error", 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
assert result == ('STOP', -1, 'API error occurred') assert result == ('STOP', -1, 'API error occurred')
gates_activity.debug.assert_called() gates_activity.debug.assert_called()
gates_activity.send_notification.assert_called() gates_activity.send_notification_async.assert_called()
@mark.asyncio @mark.asyncio
@@ -298,7 +301,7 @@ async def test_mlflow_content_gate_filter_exception(
# Assert # Assert
assert result == (None, 0, '') assert result == (None, 0, '')
gates_activity.debug.assert_called() 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'], metadata=metadata['metadata'],
notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR', notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR',
message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test 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
assert result == ('STOP', -1, 'Transformed data not passed the content filter') assert result == ('STOP', -1, 'Transformed data not passed the content filter')
gates_activity.debug.assert_called() gates_activity.debug.assert_called()
gates_activity.send_notification.assert_called() gates_activity.send_notification_async.assert_called()
@mark.asyncio @mark.asyncio
@@ -639,57 +642,103 @@ async def test_write_metrics(mock_metrics, gates_activity):
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}}, 'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}},
} }
await gates_activity.write_metrics(input_data) await gates_activity.write_metrics(input_data)
mock_metrics.PREDICTIONS_WRITTEN_COUNT.labels.assert_called_once_with( gates_activity.emit_metric.assert_has_calls(
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(
[ [
call( call(
pod_id=gates_activity.pod_id, metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT,
model_name=metadata['metadata']['model_name'], tags={
pipeline_name=metadata['metadata']['workflow_name'], 'pod_id': gates_activity.pod_id,
opc_server_id='server1', 'model_name': metadata['metadata']['model_name'],
tag='tag1', '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( call(
pod_id=gates_activity.pod_id, metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR,
model_name=metadata['metadata']['model_name'], method='set',
pipeline_name=metadata['metadata']['workflow_name'], tags={
opc_server_id='server1', 'pod_id': gates_activity.pod_id,
tag='tag1', 'model_name': metadata['metadata']['model_name'],
) 'workflow_name': metadata['metadata']['workflow_name'],
},
value=0.9,
),
] ]
) )
gates_activity.emit_metric.assert_has_calls(
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(
[ [
call(0.1), call(
call(0.2), metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR,
], method='observe',
any_order=True, 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,
),
]
) )

View File

@@ -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 import numpy as np
from pytest import fixture, mark, raises from pytest import fixture, mark, raises
@@ -25,6 +25,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
}, },
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
assert mlflow.mlflow_host == 'http://localhost' 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_username == 'admin'
assert mlflow.mlflow_password == '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( mock_minio_repository.assert_called_once_with(
logger=ANY, logger=ANY,
@@ -42,6 +45,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
minio_secret_key='minio123', minio_secret_key='minio123',
minio_region_name='us-east-1', minio_region_name='us-east-1',
minio_default_bucket='test', minio_default_bucket='test',
metrics_controller=ANY,
) )
@@ -63,9 +67,15 @@ def mlflow(mock_minio_repository, mock_mlflow_repository):
}, },
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
mlflow.model_monitoring_repository = AsyncMock()
mlflow.minio_repository = AsyncMock()
mlflow.send_notification = MagicMock() mlflow.send_notification = MagicMock()
mlflow.emit_metric = AsyncMock()
mlflow.send_notification_async = AsyncMock()
return mlflow return mlflow
@@ -225,6 +235,8 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
'message': 'Model retrained successfully.', 'message': 'Model retrained successfully.',
} }
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock()
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
@@ -301,6 +313,8 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
'message': 'Model retrained failed.', 'message': 'Model retrained failed.',
} }
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock()
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
@@ -360,7 +374,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
metadata=metadata['metadata'], metadata=metadata['metadata'],
) )
mlflow.send_notification.assert_called_once_with( mlflow.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'], metadata=metadata['metadata'],
notification_id='RETRAIN_MODEL_ERROR', notification_id='RETRAIN_MODEL_ERROR',
message='Error retraining model test_model: Model retrained failed.', 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) await mlflow.update_production_model(input_data)
except Exception as e: except Exception as e:
assert str(e) == 'Error updating production model' 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'], metadata=metadata['metadata'],
notification_id='UPDATE_PRODUCTION_MODEL_ERROR', notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
message='Error updating production model test_model: Error updating production model', message='Error updating production model test_model: Error updating production model',

View File

@@ -18,8 +18,13 @@ metadata = {
def test__init__(): def test__init__():
servers = {'server1': 'config'} servers = {'server1': {'id': 'server1'}}
opc = OPC(opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock()) opc = OPC(
opc_servers=servers,
logger=MagicMock(),
notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
)
assert opc.opc_servers == servers assert opc.opc_servers == servers
assert opc.opc_repository == {} assert opc.opc_repository == {}
@@ -27,9 +32,10 @@ def test__init__():
@mark.asyncio @mark.asyncio
@patch('laborious.activities.opc.OpcRepository') @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): async def test_init_opc(mock_send_notification, mock_opc_repository):
mock_logger = MagicMock() mock_logger = MagicMock()
mock_metrics_controller = AsyncMock()
server1 = MagicMock( server1 = MagicMock(
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {})) connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
) )
@@ -83,7 +89,10 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
}, },
} }
opc = OPC( 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() await opc.init_opc()
@@ -105,6 +114,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
server_cert_path='', server_cert_path='',
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
reconnection_interval=60, reconnection_interval=60,
metrics_controller=mock_metrics_controller,
), ),
] ]
) )
@@ -120,6 +130,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
server_cert_path='', server_cert_path='',
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
reconnection_interval=60, reconnection_interval=60,
metrics_controller=mock_metrics_controller,
) )
] ]
) )
@@ -163,9 +174,16 @@ async def opc(mock_opc_repository):
mock_opc_repository.return_value.write_data = AsyncMock(return_value=(True, {})) mock_opc_repository.return_value.write_data = AsyncMock(return_value=(True, {}))
mock_opc_repository.return_value.connect = 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() await opc.init_opc()
opc.send_notification = MagicMock() opc.send_notification = MagicMock()
opc.send_notification_async = AsyncMock()
opc.emit_metric = AsyncMock()
return opc return opc
@@ -219,7 +237,7 @@ async def test_write_data_failed(opc):
) )
assert result is None assert result is None
opc.send_notification.assert_called_once_with( opc.send_notification_async.assert_called_once_with(
metadata=metadata, metadata=metadata,
notification_id='OPC_WRITE_DATA_ERROR_server1', notification_id='OPC_WRITE_DATA_ERROR_server1',
message='Failed to write data to OPC server: Test error', message='Failed to write data to OPC server: Test error',
@@ -244,7 +262,7 @@ async def test_write_data_exception(opc):
) )
except Exception: except Exception:
opc.send_notification.assert_called_once_with( opc.send_notification_async.assert_called_once_with(
metadata=metadata, metadata=metadata,
notification_id='WRITE_OPC_PREDICTION_ERROR', notification_id='WRITE_OPC_PREDICTION_ERROR',
message='Error writing data to OPC server: Test error', message='Error writing data to OPC server: Test error',
@@ -411,7 +429,7 @@ async def test_write_opc_data_empty_config(opc):
@mark.asyncio @mark.asyncio
async def test_write_opc_data_no_validate_server(opc): 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 = { input_data = {
**metadata, **metadata,
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]}, 'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
@@ -445,13 +463,14 @@ def test_process_confidence(opc, data, success, expected):
assert result['prediction_confidence'][0] == expected assert result['prediction_confidence'][0] == expected
def test_validate_server(opc): @mark.asyncio
assert opc.validate_server('server1', metadata) is True async def test_validate_server(opc):
assert opc.validate_server('server2', metadata) is False assert await opc.validate_server('server1', metadata) is True
assert await opc.validate_server('server2', metadata) is False
@mark.asyncio @mark.asyncio
async def test_shutdown(opc): async def test_close(opc):
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True) opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
await opc.shutdown() await opc.close()
opc.opc_repository['server1'].disconnect.assert_called_once() opc.opc_repository['server1'].disconnect.assert_called_once()

View File

@@ -37,6 +37,7 @@ def storage(mock_minio_repository):
}, },
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
@@ -44,6 +45,7 @@ def storage(mock_minio_repository):
def test___init___not_hasattr(mock_minio_repository): def test___init___not_hasattr(mock_minio_repository):
logger = MagicMock() logger = MagicMock()
notification_handler = MagicMock() notification_handler = MagicMock()
metrics_controller = AsyncMock()
storage = Storage( storage = Storage(
host='localhost', host='localhost',
port=5432, port=5432,
@@ -61,6 +63,7 @@ def test___init___not_hasattr(mock_minio_repository):
}, },
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller,
) )
assert isinstance(storage, Postgres) assert isinstance(storage, Postgres)
@@ -72,6 +75,7 @@ def test___init___not_hasattr(mock_minio_repository):
minio_secret_key='minio123', minio_secret_key='minio123',
minio_region_name='us-east-1', minio_region_name='us-east-1',
minio_default_bucket='test', 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 storage.minio_repository = None
logger = MagicMock() logger = MagicMock()
notification_handler = MagicMock() notification_handler = MagicMock()
metrics_controller = AsyncMock()
storage.__init__( storage.__init__(
host='localhost', host='localhost',
port=5432, port=5432,
@@ -98,6 +102,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
}, },
logger=logger, logger=logger,
notification_handler=notification_handler, notification_handler=notification_handler,
metrics_controller=metrics_controller,
) )
mock_minio_repository.assert_called_once_with( 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_secret_key='minio123',
minio_region_name='us-east-1', minio_region_name='us-east-1',
minio_default_bucket='test', minio_default_bucket='test',
metrics_controller=metrics_controller,
) )
@@ -130,6 +136,7 @@ def test___init___done_repository(mock_minio_repository, storage):
}, },
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
mock_minio_repository.assert_not_called() mock_minio_repository.assert_not_called()
assert storage.minio_repository is not None 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}] data = [{'a': 1}, {'a': 2}, {'a': 3}]
storage.load_custom_query = AsyncMock(return_value=data) storage.load_custom_query = AsyncMock(return_value=data)
now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0) now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0)
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
storage.minio_repository.minio_bucket = 'test' storage.minio_repository.minio_bucket = 'test'
result = await storage.query_to_minio({'object_prefix': 'test', **metadata}) 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 @mark.asyncio
async def test_query_to_minio_error(storage): async def test_query_to_minio_error(storage):
storage.send_notification = MagicMock() storage.send_notification = MagicMock()
storage.send_notification_async = AsyncMock()
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
storage.load_custom_query = AsyncMock(side_effect=Exception('test')) storage.load_custom_query = AsyncMock(side_effect=Exception('test'))
result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'}) result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'})
assert result['success'] is False assert result['success'] is False
assert result['message'] == 'test' assert result['message'] == 'test'
storage.send_notification.assert_called_once_with( storage.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'], metadata=metadata['metadata'],
notification_id='ERROR_STORING_QUERY_TO_MINIO', notification_id='ERROR_STORING_QUERY_TO_MINIO',
message='Error storing query to MinIO: test', message='Error storing query to MinIO: test',

View File

@@ -1,8 +1,9 @@
from unittest.mock import MagicMock, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
from botocore.utils import ClientError 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 from laborious.utils.repository.minio_repository import MinioRepository
@@ -17,6 +18,7 @@ def test___init___(mock_config, mock_boto3):
minio_default_bucket='test', minio_default_bucket='test',
logger=MagicMock(), logger=MagicMock(),
notification_handler=MagicMock(), notification_handler=MagicMock(),
metrics_controller=AsyncMock(),
) )
assert minio_repository.storage_options == { 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.Config')
@patch('laborious.utils.repository.minio_repository.boto3') @patch('laborious.utils.repository.minio_repository.boto3')
def minio_repository(mock_boto3, mock_config): def minio_repository(mock_boto3, mock_config):
return MinioRepository( minio_repository = MinioRepository(
minio_endpoint_url='localhost:9000', minio_endpoint_url='localhost:9000',
minio_access_key='minio', minio_access_key='minio',
minio_secret_key='minio123', minio_secret_key='minio123',
@@ -58,55 +60,84 @@ def minio_repository(mock_boto3, mock_config):
minio_default_bucket='test', minio_default_bucket='test',
logger=MagicMock(), logger=MagicMock(),
notification_handler=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): def test_close(minio_repository):
minio_repository.close() minio_repository.close()
minio_repository.s3_client.close.assert_called_once() minio_repository.s3_client.close.assert_called_once()
def test_ensure_bucket_exists_bucket_exists(minio_repository): @mark.asyncio
assert minio_repository.ensure_bucket_exists({}) is None async def test_create_bucket_success(minio_repository):
await minio_repository.create_bucket({})
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
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'
)
assert 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.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
)
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_exists(minio_repository):
assert await minio_repository.ensure_bucket_exists({}) is None
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_READ_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_COUNT, tags=ANY
)
@mark.asyncio
async def test_ensure_bucket_exists_bucket_not_exists_create_success(minio_repository):
minio_repository.s3_client.head_bucket.side_effect = ClientError( minio_repository.s3_client.head_bucket.side_effect = ClientError(
error_response={'Error': {'Code': '404'}}, operation_name='head_bucket' error_response={'Error': {'Code': '404'}}, operation_name='head_bucket'
) )
minio_repository.s3_client.create_bucket.side_effect = ClientError( minio_repository.create_bucket = AsyncMock()
error_response={'Error': {'Code': '404'}}, operation_name='create_bucket'
)
with raises(ClientError): assert await minio_repository.ensure_bucket_exists({}) is None
minio_repository.ensure_bucket_exists({})
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
minio_repository.create_bucket.assert_called_once_with({})
minio_repository.observe_lag.assert_not_called()
minio_repository.emit_metric.assert_not_called()
@mark.asyncio
async def test_ensure_bucket_exists_bucket_not_exists_create_error(minio_repository):
minio_repository.s3_client.head_bucket.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.ensure_bucket_exists({})
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test') minio_repository.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') @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() 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={} dataframe=input_data, uri='s3://test/test.parquet', object_name='test.parquet', metadata={}
) )
@@ -121,15 +152,53 @@ def test_store_dataframe_as_parquet(mock_bytesio, minio_repository):
Bucket='test', Key='test.parquet', Body=mock_bytesio.return_value.getvalue.return_value 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.BytesIO')
@patch('laborious.utils.repository.minio_repository.read_parquet') @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'))} input_data = {'Body': MagicMock(read=MagicMock(return_value=b'test'))}
minio_repository.s3_client.get_object.return_value = input_data 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') minio_repository.s3_client.get_object.assert_called_once_with(Bucket='test', Key='test.parquet')
@@ -137,3 +206,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) mock_read_parquet.assert_called_once_with(mock_bytesio.return_value)
assert output == mock_read_parquet.return_value assert output == mock_read_parquet.return_value
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_READ_LAG, ANY)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_COUNT, tags=ANY
)
@mark.asyncio
@patch('laborious.utils.repository.minio_repository.BytesIO')
@patch('laborious.utils.repository.minio_repository.read_parquet')
async def test_get_parquet_as_dataframe_error(mock_read_parquet, mock_bytesio, minio_repository):
minio_repository.s3_client.get_object.side_effect = ValueError('test')
with raises(ValueError):
await minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
minio_repository.s3_client.get_object.assert_called_once_with(
Bucket='test', Key='test.parquet'
)
minio_repository.emit_metric.assert_called_once_with(
metric_object=metrics.MINIO_READ_ERROR_COUNT, tags=ANY
)
minio_repository.observe_lag.assert_not_called()

View File

@@ -1,11 +1,11 @@
from datetime import UTC, datetime 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 mlflow as mlflow_lib
import numpy as np
import pytest import pytest
from pandas import DataFrame, Timestamp from pandas import DataFrame, Timestamp
from laborious import metrics
from laborious.utils.repository.model_repository import MLFlowRepository, force_memory_release 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(): def mlflow_repository():
with patch('laborious.utils.repository.model_repository.mlflow'): with patch('laborious.utils.repository.model_repository.mlflow'):
repo = MLFlowRepository( 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 return repo
@@ -181,15 +190,16 @@ def test_get_model_params(mlflow, mlflow_repository):
assert output == mlflow.get_run.return_value.data.params 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.path')
@patch('laborious.utils.repository.model_repository.rmtree') @patch('laborious.utils.repository.model_repository.rmtree')
@patch('laborious.utils.repository.model_repository.makedirs') @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') mlflow_repository.get_model_run_id = MagicMock(return_value='test')
path.exists.return_value = True 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( mlflow_repository.get_model_run_id.assert_called_once_with(
model_name='test', stage='Production' 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 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.path')
@patch('laborious.utils.repository.model_repository.rmtree') @patch('laborious.utils.repository.model_repository.rmtree')
@patch('laborious.utils.repository.model_repository.makedirs') @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') mlflow_repository.get_model_run_id = MagicMock(return_value='test')
path.exists.return_value = False 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( mlflow_repository.get_model_run_id.assert_called_once_with(
model_name='test', stage='Production' 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 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): def test_get_experiment_error(mlflow, mlflow_repository):
mlflow.get_experiment_by_name.return_value = None mlflow.get_experiment_by_name.return_value = None
@@ -250,30 +286,51 @@ def test_get_experiment_error(mlflow, mlflow_repository):
raise AssertionError('Expected ValueError') raise AssertionError('Expected ValueError')
def test_load_predict_model_sklearn(mlflow, mlflow_repository): @pytest.mark.asyncio
result = mlflow_repository.load_predict_model('test_model', 'sklearn') 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 assert result == mlflow.sklearn.load_model.return_value
mlflow.sklearn.load_model.assert_called_once_with('models:/test_model/production') 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): @pytest.mark.asyncio
result = mlflow_repository.load_predict_model('test_model', 'pyfunc') 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 assert result == mlflow.pyfunc.load_model.return_value
mlflow.pyfunc.load_model.assert_called_once_with('models:/test_model/production') 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): @pytest.mark.asyncio
result = mlflow_repository.load_predict_model('test_model', 'pytorch') 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 assert result == mlflow.pytorch.load_model.return_value
mlflow.pytorch.load_model.assert_called_once_with('models:/test_model/production') 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: 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'." 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): def validate_common_load_transform_model_mocks(mlflow_repository, model_name):
mlflow_repository.get_model_run_id.assert_called_once_with( mlflow_repository.get_model_run_id.assert_called_once_with(
@@ -284,65 +341,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_run_id = MagicMock()
mlflow_repository.get_model_uri = 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') validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
assert result == mlflow.sklearn.load_model.return_value assert result == mlflow.sklearn.load_model.return_value
mlflow.sklearn.load_model.assert_called_once_with(mlflow_repository.get_model_uri.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_run_id = MagicMock()
mlflow_repository.get_model_uri = 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') validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
assert result == mlflow.pyfunc.load_model.return_value assert result == mlflow.pyfunc.load_model.return_value
mlflow.pyfunc.load_model.assert_called_once_with(mlflow_repository.get_model_uri.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_run_id = MagicMock()
mlflow_repository.get_model_uri = 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') validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
assert result == mlflow.pytorch.load_model.return_value assert result == mlflow.pytorch.load_model.return_value
mlflow.pytorch.load_model.assert_called_once_with(mlflow_repository.get_model_uri.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_run_id = MagicMock()
mlflow_repository.get_model_uri = MagicMock() mlflow_repository.get_model_uri = MagicMock()
with pytest.raises(ValueError) as e: 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'." 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: 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'." 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( @pytest.mark.parametrize(
'model_type', [('predict', 'prediction_model'), ('transform', 'data_model')] 'model_type', [('predict', 'prediction_model'), ('transform', 'data_model')]
) )
def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type): async def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
mlflow_repository.dowload_artifacts = MagicMock() 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.pyfunc.load_model.assert_called_once_with(
mlflow_repository.dowload_artifacts.return_value mlflow_repository.dowload_artifacts.return_value
@@ -354,26 +442,28 @@ def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
) )
def test_download_model_predict(mlflow_repository): @pytest.mark.asyncio
mlflow_repository.load_predict_model = MagicMock() async def test_download_model_predict(mlflow_repository):
mlflow_repository.load_transform_model = MagicMock() 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() mlflow_repository.load_transform_model.assert_not_called()
assert result == (mlflow_repository.load_predict_model.return_value, None) assert result == (mlflow_repository.load_predict_model.return_value, None)
def test_download_model_transform(mlflow_repository): @pytest.mark.asyncio
mlflow_repository.load_predict_model = MagicMock() async def test_download_model_transform(mlflow_repository):
mlflow_repository.load_transform_model = MagicMock() 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_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) assert result == (mlflow_repository.load_transform_model.return_value, None)
@@ -481,19 +571,25 @@ def test_handle_outdated_model(mlflow_repository):
assert mlflow_repository.model_cache == {} 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() model = MagicMock()
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path')) mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
output = mlflow_repository.get_model('model_name', 0, 'predict', 'pyfunc') output = await mlflow_repository.get_model('model_name', {}, 0, 'predict', 'pyfunc')
assert output == model assert output == model
mlflow_repository.download_model.assert_called_once_with( 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.check_cache_retention = MagicMock(return_value=True)
mlflow_repository.handle_valid_model = MagicMock() mlflow_repository.handle_valid_model = MagicMock()
mlflow_repository.handle_outdated_model = MagicMock() mlflow_repository.handle_outdated_model = MagicMock()
@@ -503,7 +599,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 assert output == mlflow_repository.handle_valid_model.return_value
mlflow_repository.check_cache_retention.assert_called_once_with( mlflow_repository.check_cache_retention.assert_called_once_with(
@@ -517,12 +613,13 @@ def test_get_model_cached_valid(mlflow_repository):
mlflow_repository.handle_outdated_model.assert_not_called() 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.check_cache_retention = MagicMock(return_value=False)
mlflow_repository.handle_valid_model = MagicMock() mlflow_repository.handle_valid_model = MagicMock()
mlflow_repository.handle_outdated_model = MagicMock() mlflow_repository.handle_outdated_model = MagicMock()
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 = { cache = {
'model_name_predict': { 'model_name_predict': {
'target': 'cached_model', 'target': 'cached_model',
@@ -530,7 +627,7 @@ def test_get_model_cached_outdated(mlflow_repository):
} }
mlflow_repository.model_cache = cache 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 assert output == model
mlflow_repository.check_cache_retention.assert_called_once_with( mlflow_repository.check_cache_retention.assert_called_once_with(
{ {
@@ -544,14 +641,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.check_cache_retention = MagicMock(return_value=False)
mlflow_repository.handle_valid_model = MagicMock() mlflow_repository.handle_valid_model = MagicMock()
mlflow_repository.handle_outdated_model = MagicMock() mlflow_repository.handle_outdated_model = MagicMock()
mlflow_repository.model_cache = {} mlflow_repository.model_cache = {}
model = MagicMock() model = MagicMock()
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path')) mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc') output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
assert output == model assert output == model
mlflow_repository.check_cache_retention.assert_not_called() mlflow_repository.check_cache_retention.assert_not_called()
mlflow_repository.handle_valid_model.assert_not_called() mlflow_repository.handle_valid_model.assert_not_called()
@@ -559,43 +657,53 @@ def test_get_model_cached_not_found(mlflow_repository):
@patch('laborious.utils.repository.model_repository.force_memory_release') @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() model = MagicMock()
data = MagicMock() data = MagicMock()
mlflow_repository.get_model = MagicMock(return_value=model) mlflow_repository.get_model = AsyncMock(return_value=model)
output = mlflow_repository.get_cached_operation('model_name', data, 'transform', 0, 'sklearn') output = await mlflow_repository.get_cached_operation(
'model_name', data, 'transform', 0, 'sklearn', {}
)
assert output == model.predict.return_value assert output == model.predict.return_value
force_memory_release.assert_called_once_with(mlflow_repository.logger) force_memory_release.assert_called_once_with(mlflow_repository.logger)
@patch('laborious.utils.repository.model_repository.force_memory_release') @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() model = MagicMock()
data = MagicMock() data = MagicMock()
mlflow_repository.get_model = MagicMock(return_value=model) mlflow_repository.get_model = AsyncMock(return_value=model)
output = mlflow_repository.get_cached_operation('model_name', data, 'predict', 1, 'sklearn') output = await mlflow_repository.get_cached_operation(
'model_name', data, 'predict', 1, 'sklearn', {}
)
assert output == model.predict.return_value assert output == model.predict.return_value
force_memory_release.assert_not_called() force_memory_release.assert_not_called()
@patch('laborious.utils.repository.model_repository.force_memory_release') @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() data = MagicMock()
with pytest.raises(ValueError) as e: 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'." assert str(e) == "Invalid operation. Use 'transform' or 'predict'."
@patch('laborious.utils.repository.model_repository.pd.merge') @patch('laborious.utils.repository.model_repository.pd.merge')
@patch('laborious.utils.repository.model_repository.isinstance') @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, pd_merge, mlflow_repository
): ):
isinstance_mock.return_value = False isinstance_mock.return_value = False
data_model = MagicMock() data_model = MagicMock()
prediction_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')], side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
) )
mlflow_repository.detect_and_parse_datetime_index = MagicMock( mlflow_repository.detect_and_parse_datetime_index = MagicMock(
@@ -604,7 +712,7 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
data = MagicMock() data = MagicMock()
output = mlflow_repository.fit_models( output = await mlflow_repository.fit_models(
'model_name', data, 'latest_production_id', metadata['metadata'], 'sklearn', 'pyfunc', None 'model_name', data, 'latest_production_id', metadata['metadata'], 'sklearn', 'pyfunc', None
) )
@@ -612,11 +720,18 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
[ [
call( call(
model_name='model_name', model_name='model_name',
metadata=metadata['metadata'],
model_type='transform', model_type='transform',
flavor='sklearn', flavor='sklearn',
load_wrapper=False, 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 +773,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.pd.merge')
@patch('laborious.utils.repository.model_repository.isinstance') @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, pd_merge, mlflow_repository
): ):
isinstance_mock.return_value = True isinstance_mock.return_value = True
@@ -667,7 +783,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
target_variable='feat_2', target_variable='feat_2',
) )
prediction_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')], side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
) )
mlflow_repository.detect_and_parse_datetime_index = MagicMock( mlflow_repository.detect_and_parse_datetime_index = MagicMock(
@@ -678,7 +794,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
data = MagicMock() data = MagicMock()
output = mlflow_repository.fit_models( output = await mlflow_repository.fit_models(
'model_name', 'model_name',
data, data,
'latest_production_id', 'latest_production_id',
@@ -692,11 +808,18 @@ def test_fit_models_df_target_name_not_none_and_in_model(
[ [
call( call(
model_name='model_name', model_name='model_name',
metadata=metadata['metadata'],
model_type='transform', model_type='transform',
flavor='sklearn', flavor='sklearn',
load_wrapper=False, 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 +851,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'} 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.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') @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'} 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() mlflow.pyfunc.log_model.assert_not_called()
@@ -747,24 +881,47 @@ def test_log_model_pyfunc(path, mlflow, mlflow_repository):
artifact_path='prediction_model', code_path=[path.join.return_value], to_disk=False 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'} 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.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'} model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
with pytest.raises(ValueError) as e: 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'." 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.force_memory_release')
@patch('laborious.utils.repository.model_repository.path') @patch('laborious.utils.repository.model_repository.path')
@patch('laborious.utils.repository.model_repository.rmtree') @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' model_name = 'model_name'
data = MagicMock() data = MagicMock()
retrain_data = { retrain_data = {
@@ -782,11 +939,11 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
mlflow_repository.get_experiment = MagicMock() mlflow_repository.get_experiment = MagicMock()
mlflow_repository.get_next_run_name = MagicMock() mlflow_repository.get_next_run_name = MagicMock()
mlflow_repository.log_model = MagicMock() mlflow_repository.log_model = AsyncMock()
path.exists.return_value = True path.exists.return_value = True
path.join.return_value = './tmp/artifacts/model_name' path.join.return_value = './tmp/artifacts/model_name'
report = mlflow_repository.create_new_experiment( report = await mlflow_repository.create_new_experiment(
model_name, model_name,
data, data,
retrain_data, retrain_data,
@@ -843,8 +1000,60 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
'experiment_name': mlflow_repository.get_experiment.return_value.name, '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( mlflow_repository.client.get_registered_model.return_value = MagicMock(
latest_versions=[ latest_versions=[
MagicMock(version='1'), MagicMock(version='1'),
@@ -852,7 +1061,9 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
MagicMock(version='3'), 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( mlflow.register_model.assert_called_once_with(
'runs:/0/prediction_model', 'runs:/0/prediction_model',
@@ -873,21 +1084,62 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
'mlflow_run_id': '0', '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( mlflow_repository.client.get_registered_model.return_value = MagicMock(
get_registered_model=MagicMock(return_value=MagicMock(latest_versions={})) get_registered_model=MagicMock(return_value=MagicMock(latest_versions={}))
) )
try: 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: except Exception as e:
assert str(e) == 'Model versions is not a list' assert str(e) == 'Model versions is not a list'
else: else:
raise AssertionError('Expected Exception') raise AssertionError('Expected Exception')
def test_transform_success(mlflow_repository): @pytest.mark.asyncio
async def test_transform_success(mlflow_repository):
data = MagicMock() data = MagicMock()
model_name = 'model' model_name = 'model'
model_config = { model_config = {
@@ -896,14 +1148,19 @@ def test_transform_success(mlflow_repository):
'predict_flavor': 'pyfunc', '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() 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( 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( mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
@@ -916,7 +1173,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() data = MagicMock()
model_name = 'model' model_name = 'model'
model_config = { model_config = {
@@ -925,27 +1183,40 @@ def test_transform_error(mlflow_repository):
'predict_flavor': 'pyfunc', '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( 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}} 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}}) data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
model_name = 'model' 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( 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 assert output['success'] is True
@@ -955,19 +1226,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}}) data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
model_name = 'model' 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}}) 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( 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 assert output['success'] is True
@@ -977,32 +1254,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}}) data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_name = 'model' model_name = 'model'
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'} 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( 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}} 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() data = MagicMock()
model_name = 'test' model_name = 'test'
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'} model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
mlflow_repository.get_model_run_id = MagicMock() mlflow_repository.get_model_run_id = MagicMock()
mlflow_repository.fit_models = MagicMock() mlflow_repository.fit_models = AsyncMock()
mlflow_repository.create_new_experiment = MagicMock() 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') mlflow_repository.get_model_run_id.assert_called_once_with(model_name, stage='Production')
@@ -1033,12 +1319,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() data = MagicMock()
model_name = 'test' model_name = 'test'
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'} model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
mlflow_repository.get_model_run_id = MagicMock(side_effect=Exception('error')) 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 == { assert output == {
'success': False, 'success': False,
'experiment': None, 'experiment': None,
@@ -1047,17 +1336,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'} experiment = {'run_id': '0', 'experiment_id': '0'}
model_name = 'test' 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 = { mlflow_repository.update_production_model_by_run_id.return_value = {
'model_name': 'test', 'model_name': 'test',
'version': '3', 'version': '3',
'mlflow_run_id': '0', '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( mlflow_repository.update_production_model_by_run_id.assert_called_once_with(
'0', 'test', metadata['metadata'] '0', 'test', metadata['metadata']

View File

@@ -26,8 +26,12 @@ def opc_repository(mock_logger):
cert_path='/path/to/cert.pem', cert_path='/path/to/cert.pem',
private_key_path='/path/to/key.pem', private_key_path='/path/to/key.pem',
server_cert_path='/path/to/server_cert.pem', server_cert_path='/path/to/server_cert.pem',
metrics_controller=AsyncMock(),
) )
repository.disconnection_interval = 0.1
repository.send_notification = MagicMock() repository.send_notification = MagicMock()
repository.send_notification_async = AsyncMock()
repository.emit_metric = AsyncMock()
return repository return repository
@@ -214,7 +218,7 @@ async def test_disconnect_error(opc_repository, mock_client):
await opc_repository.disconnect() await opc_repository.disconnect()
opc_repository.disconnection_fallback.assert_called_once() 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, metadata=opc_repository.metadata,
notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}', notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}',
message='Failed to disconnect from OPC server in 5 attempts.', message='Failed to disconnect from OPC server in 5 attempts.',