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:
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
from pytest import mark
|
||||
|
||||
@@ -13,7 +13,10 @@ from laborious.activities.storage import Storage
|
||||
@patch('laborious.activities.activities.MLFlow.__init__')
|
||||
@patch('laborious.activities.activities.OPC.__init__')
|
||||
@patch('laborious.activities.activities.Gates.__init__')
|
||||
def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init):
|
||||
@patch('laborious.activities.activities.MetricsController')
|
||||
def test___init__(
|
||||
mock_metrics_controller, mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init
|
||||
):
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -70,6 +73,7 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_mlflow_init.assert_called_once_with(
|
||||
@@ -81,22 +85,32 @@ def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage
|
||||
minio_config=minio_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_opc_init.assert_called_once_with(
|
||||
ANY, opc_servers=opc_config, logger=logger, notification_handler=notification_handler
|
||||
ANY,
|
||||
opc_servers=opc_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
mock_gates_init.assert_called_once_with(
|
||||
ANY, logger=logger, notification_handler=notification_handler
|
||||
ANY,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.activities.Storage', return_value=MagicMock())
|
||||
@patch('laborious.activities.activities.MLFlow', return_value=MagicMock())
|
||||
@patch('laborious.activities.activities.OPC', return_value=MagicMock())
|
||||
async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init):
|
||||
@patch('laborious.activities.activities.Storage')
|
||||
@patch('laborious.activities.activities.MLFlow')
|
||||
@patch('laborious.activities.activities.OPC')
|
||||
@patch('laborious.activities.activities.Gates')
|
||||
async def test_shutdown(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_storage_init):
|
||||
mock_opc_init.close = AsyncMock()
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -136,5 +150,7 @@ async def test_shutdown(mock_opc_init, _mock_mlflow_init, mock_storage_init):
|
||||
)
|
||||
|
||||
await activities.shutdown()
|
||||
mock_opc_init.shutdown.assert_called_once()
|
||||
mock_opc_init.close.assert_called_once()
|
||||
mock_storage_init.close.assert_called_once()
|
||||
mock_mlflow_init.close.assert_called_once()
|
||||
mock_gates_init.close.assert_called_once()
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
from pytest import fixture, mark
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
@@ -11,6 +11,7 @@ def gates_activity():
|
||||
gates = Gates(
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
gates.error = MagicMock()
|
||||
gates.debug = MagicMock()
|
||||
@@ -18,6 +19,8 @@ def gates_activity():
|
||||
gates.warning = MagicMock()
|
||||
gates.critical = MagicMock()
|
||||
gates.send_notification = MagicMock()
|
||||
gates.send_notification_async = AsyncMock()
|
||||
gates.emit_metric = AsyncMock()
|
||||
return gates
|
||||
|
||||
|
||||
@@ -71,7 +74,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='INTPUT_GATE_ERROR__EMPTY_DATA',
|
||||
message="Error in filter EMPTY_DATA:{'policy': 'STOP', 'config': {}}: \n Test error",
|
||||
@@ -176,7 +179,7 @@ async def test_mlflow_response_gate_filter_exception(
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER',
|
||||
message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error",
|
||||
@@ -225,7 +228,7 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'API error occurred')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@@ -298,7 +301,7 @@ async def test_mlflow_content_gate_filter_exception(
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called_once_with(
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MLFLOW_GATE_CONTENT_FILTER__API_ERROR',
|
||||
message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test error",
|
||||
@@ -344,7 +347,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Transformed data not passed the content filter')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@@ -639,57 +642,103 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': 0.2}},
|
||||
}
|
||||
await gates_activity.write_metrics(input_data)
|
||||
mock_metrics.PREDICTIONS_WRITTEN_COUNT.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTIONS_WRITTEN_COUNT.labels.return_value.inc.assert_called_once_with()
|
||||
|
||||
mock_metrics.PREDICTION_CONFIDENCE_MONITOR.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTION_CONFIDENCE_MONITOR.labels.return_value.set.assert_called_once_with(0.9)
|
||||
|
||||
mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels.assert_called_once_with(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
)
|
||||
mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR.labels.return_value.observe.assert_called_once_with(
|
||||
0.1
|
||||
)
|
||||
|
||||
mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
opc_server_id='server1',
|
||||
tag='tag1',
|
||||
)
|
||||
metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
pod_id=gates_activity.pod_id,
|
||||
model_name=metadata['metadata']['model_name'],
|
||||
pipeline_name=metadata['metadata']['workflow_name'],
|
||||
opc_server_id='server1',
|
||||
tag='tag1',
|
||||
)
|
||||
metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR,
|
||||
method='set',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
value=0.9,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
assert mock_metrics.PREDICTION_OPC_WRITING_COUNT.labels.return_value.inc.call_count == 2
|
||||
mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR.labels.return_value.observe.assert_has_calls(
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(0.1),
|
||||
call(0.2),
|
||||
],
|
||||
any_order=True,
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
},
|
||||
value=0.1,
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag1',
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag1',
|
||||
},
|
||||
value=0.1,
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag2',
|
||||
},
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag2',
|
||||
},
|
||||
value=0.2,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
import numpy as np
|
||||
from pytest import fixture, mark, raises
|
||||
@@ -25,6 +25,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert mlflow.mlflow_host == 'http://localhost'
|
||||
@@ -32,7 +33,9 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
assert mlflow.mlflow_username == 'admin'
|
||||
assert mlflow.mlflow_password == 'admin'
|
||||
|
||||
mock_mlflow_repository.assert_called_once_with('http://localhost:5000', 'admin', 'admin', ANY)
|
||||
mock_mlflow_repository.assert_called_once_with(
|
||||
'http://localhost:5000', 'admin', 'admin', ANY, ANY, ANY
|
||||
)
|
||||
|
||||
mock_minio_repository.assert_called_once_with(
|
||||
logger=ANY,
|
||||
@@ -42,6 +45,7 @@ def test___init__(mock_minio_repository, mock_mlflow_repository):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=ANY,
|
||||
)
|
||||
|
||||
|
||||
@@ -63,9 +67,15 @@ def mlflow(mock_minio_repository, mock_mlflow_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
mlflow.model_monitoring_repository = AsyncMock()
|
||||
mlflow.minio_repository = AsyncMock()
|
||||
|
||||
mlflow.send_notification = MagicMock()
|
||||
mlflow.emit_metric = AsyncMock()
|
||||
mlflow.send_notification_async = AsyncMock()
|
||||
|
||||
return mlflow
|
||||
|
||||
@@ -225,6 +235,8 @@ async def test_retrain_model_success_data_success_retrain(mock_to_datetime, mlfl
|
||||
'message': 'Model retrained successfully.',
|
||||
}
|
||||
|
||||
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock()
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
@@ -301,6 +313,8 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
'message': 'Model retrained failed.',
|
||||
}
|
||||
|
||||
mlflow.minio_repository.get_parquet_as_dataframe.return_value = MagicMock()
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
@@ -360,7 +374,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
mlflow.send_notification.assert_called_once_with(
|
||||
mlflow.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='RETRAIN_MODEL_ERROR',
|
||||
message='Error retraining model test_model: Model retrained failed.',
|
||||
@@ -464,7 +478,7 @@ async def test_update_production_model_error(mlflow):
|
||||
await mlflow.update_production_model(input_data)
|
||||
except Exception as e:
|
||||
assert str(e) == 'Error updating production model'
|
||||
mlflow.send_notification.assert_called_once_with(
|
||||
mlflow.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
|
||||
message='Error updating production model test_model: Error updating production model',
|
||||
|
||||
@@ -18,8 +18,13 @@ metadata = {
|
||||
|
||||
|
||||
def test__init__():
|
||||
servers = {'server1': 'config'}
|
||||
opc = OPC(opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock())
|
||||
servers = {'server1': {'id': 'server1'}}
|
||||
opc = OPC(
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert opc.opc_servers == servers
|
||||
assert opc.opc_repository == {}
|
||||
@@ -27,9 +32,10 @@ def test__init__():
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.opc.OpcRepository')
|
||||
@patch('laborious.activities.opc.OPC.send_notification')
|
||||
@patch('laborious.activities.opc.OPC.send_notification_async')
|
||||
async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
mock_logger = MagicMock()
|
||||
mock_metrics_controller = AsyncMock()
|
||||
server1 = MagicMock(
|
||||
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
|
||||
)
|
||||
@@ -83,7 +89,10 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
},
|
||||
}
|
||||
opc = OPC(
|
||||
opc_servers=servers, logger=mock_logger, notification_handler=mock_notification_handler
|
||||
opc_servers=servers,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
await opc.init_opc()
|
||||
|
||||
@@ -105,6 +114,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
),
|
||||
]
|
||||
)
|
||||
@@ -120,6 +130,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
]
|
||||
)
|
||||
@@ -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.connect = AsyncMock(return_value=(True, {}))
|
||||
opc = OPC(opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock())
|
||||
opc = OPC(
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
await opc.init_opc()
|
||||
opc.send_notification = MagicMock()
|
||||
opc.send_notification_async = AsyncMock()
|
||||
opc.emit_metric = AsyncMock()
|
||||
return opc
|
||||
|
||||
|
||||
@@ -219,7 +237,7 @@ async def test_write_data_failed(opc):
|
||||
)
|
||||
assert result is None
|
||||
|
||||
opc.send_notification.assert_called_once_with(
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='OPC_WRITE_DATA_ERROR_server1',
|
||||
message='Failed to write data to OPC server: Test error',
|
||||
@@ -244,7 +262,7 @@ async def test_write_data_exception(opc):
|
||||
)
|
||||
|
||||
except Exception:
|
||||
opc.send_notification.assert_called_once_with(
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='WRITE_OPC_PREDICTION_ERROR',
|
||||
message='Error writing data to OPC server: Test error',
|
||||
@@ -411,7 +429,7 @@ async def test_write_opc_data_empty_config(opc):
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_opc_data_no_validate_server(opc):
|
||||
opc.validate_server = MagicMock(return_value=False)
|
||||
opc.validate_server = AsyncMock(return_value=False)
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||
@@ -445,13 +463,14 @@ def test_process_confidence(opc, data, success, expected):
|
||||
assert result['prediction_confidence'][0] == expected
|
||||
|
||||
|
||||
def test_validate_server(opc):
|
||||
assert opc.validate_server('server1', metadata) is True
|
||||
assert opc.validate_server('server2', metadata) is False
|
||||
@mark.asyncio
|
||||
async def test_validate_server(opc):
|
||||
assert await opc.validate_server('server1', metadata) is True
|
||||
assert await opc.validate_server('server2', metadata) is False
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_shutdown(opc):
|
||||
async def test_close(opc):
|
||||
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
|
||||
await opc.shutdown()
|
||||
await opc.close()
|
||||
opc.opc_repository['server1'].disconnect.assert_called_once()
|
||||
|
||||
@@ -37,6 +37,7 @@ def storage(mock_minio_repository):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
|
||||
@@ -44,6 +45,7 @@ def storage(mock_minio_repository):
|
||||
def test___init___not_hasattr(mock_minio_repository):
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
storage = Storage(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
@@ -61,6 +63,7 @@ def test___init___not_hasattr(mock_minio_repository):
|
||||
},
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
assert isinstance(storage, Postgres)
|
||||
|
||||
@@ -72,6 +75,7 @@ def test___init___not_hasattr(mock_minio_repository):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
|
||||
@@ -80,7 +84,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
storage.minio_repository = None
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
|
||||
metrics_controller = AsyncMock()
|
||||
storage.__init__(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
@@ -98,6 +102,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
},
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
mock_minio_repository.assert_called_once_with(
|
||||
@@ -108,6 +113,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
minio_secret_key='minio123',
|
||||
minio_region_name='us-east-1',
|
||||
minio_default_bucket='test',
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
|
||||
@@ -130,6 +136,7 @@ def test___init___done_repository(mock_minio_repository, storage):
|
||||
},
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
mock_minio_repository.assert_not_called()
|
||||
assert storage.minio_repository is not None
|
||||
@@ -162,6 +169,7 @@ async def test_query_to_minio_success(now, dataframe, storage):
|
||||
data = [{'a': 1}, {'a': 2}, {'a': 3}]
|
||||
storage.load_custom_query = AsyncMock(return_value=data)
|
||||
now.return_value = datetime.datetime(2024, 1, 1, 0, 0, 0)
|
||||
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
|
||||
storage.minio_repository.minio_bucket = 'test'
|
||||
|
||||
result = await storage.query_to_minio({'object_prefix': 'test', **metadata})
|
||||
@@ -183,11 +191,14 @@ async def test_query_to_minio_success(now, dataframe, storage):
|
||||
@mark.asyncio
|
||||
async def test_query_to_minio_error(storage):
|
||||
storage.send_notification = MagicMock()
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.store_dataframe_as_parquet = AsyncMock()
|
||||
|
||||
storage.load_custom_query = AsyncMock(side_effect=Exception('test'))
|
||||
result = await storage.query_to_minio({**metadata, 'object_prefix': 'test'})
|
||||
assert result['success'] is False
|
||||
assert result['message'] == 'test'
|
||||
storage.send_notification.assert_called_once_with(
|
||||
storage.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='ERROR_STORING_QUERY_TO_MINIO',
|
||||
message='Error storing query to MinIO: test',
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
from botocore.utils import ClientError
|
||||
from pytest import fixture, raises
|
||||
from pytest import fixture, mark, raises
|
||||
|
||||
from laborious import metrics
|
||||
from laborious.utils.repository.minio_repository import MinioRepository
|
||||
|
||||
|
||||
@@ -17,6 +18,7 @@ def test___init___(mock_config, mock_boto3):
|
||||
minio_default_bucket='test',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
assert minio_repository.storage_options == {
|
||||
@@ -50,7 +52,7 @@ def test___init___(mock_config, mock_boto3):
|
||||
@patch('laborious.utils.repository.minio_repository.Config')
|
||||
@patch('laborious.utils.repository.minio_repository.boto3')
|
||||
def minio_repository(mock_boto3, mock_config):
|
||||
return MinioRepository(
|
||||
minio_repository = MinioRepository(
|
||||
minio_endpoint_url='localhost:9000',
|
||||
minio_access_key='minio',
|
||||
minio_secret_key='minio123',
|
||||
@@ -58,55 +60,84 @@ def minio_repository(mock_boto3, mock_config):
|
||||
minio_default_bucket='test',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
|
||||
minio_repository.emit_metric = AsyncMock()
|
||||
minio_repository.observe_lag = AsyncMock()
|
||||
minio_repository.send_notification = MagicMock()
|
||||
minio_repository.send_notification_async = AsyncMock()
|
||||
|
||||
return minio_repository
|
||||
|
||||
|
||||
def test_close(minio_repository):
|
||||
minio_repository.close()
|
||||
minio_repository.s3_client.close.assert_called_once()
|
||||
|
||||
|
||||
def test_ensure_bucket_exists_bucket_exists(minio_repository):
|
||||
assert minio_repository.ensure_bucket_exists({}) is None
|
||||
@mark.asyncio
|
||||
async def test_create_bucket_success(minio_repository):
|
||||
await minio_repository.create_bucket({})
|
||||
|
||||
minio_repository.s3_client.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.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(
|
||||
error_response={'Error': {'Code': '404'}}, operation_name='head_bucket'
|
||||
)
|
||||
minio_repository.s3_client.create_bucket.side_effect = ClientError(
|
||||
error_response={'Error': {'Code': '404'}}, operation_name='create_bucket'
|
||||
)
|
||||
minio_repository.create_bucket = AsyncMock()
|
||||
|
||||
with raises(ClientError):
|
||||
minio_repository.ensure_bucket_exists({})
|
||||
assert await minio_repository.ensure_bucket_exists({}) is None
|
||||
|
||||
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
|
||||
minio_repository.create_bucket.assert_called_once_with({})
|
||||
|
||||
minio_repository.observe_lag.assert_not_called()
|
||||
minio_repository.emit_metric.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_ensure_bucket_exists_bucket_not_exists_create_error(minio_repository):
|
||||
minio_repository.s3_client.head_bucket.side_effect = ValueError('test')
|
||||
|
||||
with raises(ValueError):
|
||||
await minio_repository.ensure_bucket_exists({})
|
||||
|
||||
minio_repository.s3_client.head_bucket.assert_called_once_with(Bucket='test')
|
||||
minio_repository.s3_client.create_bucket.assert_called_once_with(Bucket='test')
|
||||
minio_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
minio_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.utils.repository.minio_repository.BytesIO')
|
||||
def test_store_dataframe_as_parquet(mock_bytesio, minio_repository):
|
||||
async def test_store_dataframe_as_parquet_success(mock_bytesio, minio_repository):
|
||||
input_data = MagicMock()
|
||||
|
||||
minio_repository.ensure_bucket_exists = MagicMock(return_value=True)
|
||||
minio_repository.ensure_bucket_exists = AsyncMock()
|
||||
|
||||
minio_repository.store_dataframe_as_parquet(
|
||||
await minio_repository.store_dataframe_as_parquet(
|
||||
dataframe=input_data, uri='s3://test/test.parquet', object_name='test.parquet', metadata={}
|
||||
)
|
||||
|
||||
@@ -121,15 +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
|
||||
)
|
||||
|
||||
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_WRITE_LAG, ANY)
|
||||
minio_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MINIO_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.utils.repository.minio_repository.BytesIO')
|
||||
async def test_store_dataframe_as_parquet_error(mock_bytesio, minio_repository):
|
||||
input_data = MagicMock()
|
||||
|
||||
minio_repository.ensure_bucket_exists = AsyncMock()
|
||||
minio_repository.s3_client.put_object.side_effect = ValueError('test')
|
||||
|
||||
with raises(ValueError):
|
||||
await minio_repository.store_dataframe_as_parquet(
|
||||
dataframe=input_data,
|
||||
uri='s3://test/test.parquet',
|
||||
object_name='test.parquet',
|
||||
metadata={},
|
||||
)
|
||||
|
||||
minio_repository.ensure_bucket_exists.assert_called_once_with({})
|
||||
mock_bytesio.assert_called_once()
|
||||
|
||||
input_data.to_parquet.assert_called_once_with(
|
||||
mock_bytesio.return_value, engine='pyarrow', index=True
|
||||
)
|
||||
mock_bytesio.return_value.seek.assert_called_once_with(0)
|
||||
minio_repository.s3_client.put_object.assert_called_once_with(
|
||||
Bucket='test', Key='test.parquet', Body=mock_bytesio.return_value.getvalue.return_value
|
||||
)
|
||||
|
||||
minio_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MINIO_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.utils.repository.minio_repository.BytesIO')
|
||||
@patch('laborious.utils.repository.minio_repository.read_parquet')
|
||||
def test_get_parquet_as_dataframe(mock_read_parquet, mock_bytesio, minio_repository):
|
||||
async def test_get_parquet_as_dataframe_success(mock_read_parquet, mock_bytesio, minio_repository):
|
||||
input_data = {'Body': MagicMock(read=MagicMock(return_value=b'test'))}
|
||||
|
||||
minio_repository.s3_client.get_object.return_value = input_data
|
||||
|
||||
output = minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
|
||||
output = await minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
|
||||
|
||||
minio_repository.s3_client.get_object.assert_called_once_with(Bucket='test', Key='test.parquet')
|
||||
|
||||
@@ -137,3 +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)
|
||||
|
||||
assert output == mock_read_parquet.return_value
|
||||
|
||||
minio_repository.observe_lag.assert_called_once_with(ANY, metrics.MINIO_READ_LAG, ANY)
|
||||
minio_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MINIO_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.utils.repository.minio_repository.BytesIO')
|
||||
@patch('laborious.utils.repository.minio_repository.read_parquet')
|
||||
async def test_get_parquet_as_dataframe_error(mock_read_parquet, mock_bytesio, minio_repository):
|
||||
minio_repository.s3_client.get_object.side_effect = ValueError('test')
|
||||
|
||||
with raises(ValueError):
|
||||
await minio_repository.get_parquet_as_dataframe(object_key='test.parquet', metadata={})
|
||||
|
||||
minio_repository.s3_client.get_object.assert_called_once_with(
|
||||
Bucket='test', Key='test.parquet'
|
||||
)
|
||||
minio_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MINIO_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
minio_repository.observe_lag.assert_not_called()
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
from datetime import UTC, datetime
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
import mlflow as mlflow_lib
|
||||
import numpy as np
|
||||
import pytest
|
||||
from pandas import DataFrame, Timestamp
|
||||
|
||||
from laborious import metrics
|
||||
from laborious.utils.repository.model_repository import MLFlowRepository, force_memory_release
|
||||
|
||||
|
||||
@@ -38,8 +38,17 @@ def test_force_memory_release_error(gc, ctypes):
|
||||
def mlflow_repository():
|
||||
with patch('laborious.utils.repository.model_repository.mlflow'):
|
||||
repo = MLFlowRepository(
|
||||
host='http://localhost:5000', username='admin', password='admin', logger=MagicMock()
|
||||
host='http://localhost:5000',
|
||||
username='admin',
|
||||
password='admin',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
repo.emit_metric = AsyncMock()
|
||||
repo.observe_lag = AsyncMock()
|
||||
repo.send_notification = MagicMock()
|
||||
repo.send_notification_async = AsyncMock()
|
||||
return repo
|
||||
|
||||
|
||||
@@ -181,15 +190,16 @@ def test_get_model_params(mlflow, mlflow_repository):
|
||||
assert output == mlflow.get_run.return_value.data.params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
async def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = True
|
||||
|
||||
output = mlflow_repository.dowload_artifacts('test', 'path')
|
||||
output = await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
model_name='test', stage='Production'
|
||||
@@ -209,16 +219,22 @@ def test_download_artifacts_success(makedirs, rmtree, path, mlflow_repository):
|
||||
|
||||
assert output == mlflow_repository.client.download_artifacts.return_value
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_repository):
|
||||
async def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = False
|
||||
|
||||
output = mlflow_repository.dowload_artifacts('test', 'path')
|
||||
output = await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
model_name='test', stage='Production'
|
||||
@@ -239,6 +255,26 @@ def test_download_artifacts_success_path_false(makedirs, rmtree, path, mlflow_re
|
||||
assert output == mlflow_repository.client.download_artifacts.return_value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@patch('laborious.utils.repository.model_repository.makedirs')
|
||||
async def test_download_artifacts_error(makedirs, rmtree, path, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='test')
|
||||
|
||||
path.exists.return_value = True
|
||||
|
||||
mlflow_repository.client.download_artifacts.side_effect = ValueError('test')
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.dowload_artifacts('test', {}, 'path')
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
def test_get_experiment_error(mlflow, mlflow_repository):
|
||||
mlflow.get_experiment_by_name.return_value = None
|
||||
|
||||
@@ -250,30 +286,51 @@ def test_get_experiment_error(mlflow, mlflow_repository):
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
def test_load_predict_model_sklearn(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'sklearn')
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_sklearn(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'sklearn')
|
||||
|
||||
assert result == mlflow.sklearn.load_model.return_value
|
||||
mlflow.sklearn.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_pyfunc(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'pyfunc')
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_pyfunc(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'pyfunc')
|
||||
assert result == mlflow.pyfunc.load_model.return_value
|
||||
mlflow.pyfunc.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_pytorch(mlflow, mlflow_repository):
|
||||
result = mlflow_repository.load_predict_model('test_model', 'pytorch')
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_pytorch(mlflow, mlflow_repository):
|
||||
result = await mlflow_repository.load_predict_model('test_model', {}, 'pytorch')
|
||||
assert result == mlflow.pytorch.load_model.return_value
|
||||
mlflow.pytorch.load_model.assert_called_once_with('models:/test_model/production')
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
def test_load_predict_model_error(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_predict_model_error(mlflow_repository):
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.load_predict_model('test_model', 'invalid')
|
||||
await mlflow_repository.load_predict_model('test_model', {}, 'invalid')
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
def validate_common_load_transform_model_mocks(mlflow_repository, model_name):
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(
|
||||
@@ -284,65 +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_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'sklearn')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'sklearn')
|
||||
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.sklearn.load_model.return_value
|
||||
mlflow.sklearn.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_pyfunc(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_pyfunc(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'pyfunc')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'pyfunc')
|
||||
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.pyfunc.load_model.return_value
|
||||
mlflow.pyfunc.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_pytorch(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_pytorch(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
result = mlflow_repository.load_transform_model('test_model', 'pytorch')
|
||||
result = await mlflow_repository.load_transform_model('test_model', {}, 'pytorch')
|
||||
validate_common_load_transform_model_mocks(mlflow_repository, 'test_model')
|
||||
|
||||
assert result == mlflow.pytorch.load_model.return_value
|
||||
mlflow.pytorch.load_model.assert_called_once_with(mlflow_repository.get_model_uri.return_value)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_load_transform_model_error(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_transform_model_error(mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.get_model_uri = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.load_transform_model('test_model', 'invalid')
|
||||
await mlflow_repository.load_transform_model('test_model', {}, 'invalid')
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
def test_download_model_invalid_model_type(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_invalid_model_type(mlflow_repository):
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.download_model('test_model', 'invalid', 'sklearn')
|
||||
await mlflow_repository.download_model('test_model', {}, 'invalid', 'sklearn')
|
||||
assert str(e) == "Invalid model_type. Use 'predict' or 'transform'."
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
'model_type', [('predict', 'prediction_model'), ('transform', 'data_model')]
|
||||
)
|
||||
def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
mlflow_repository.dowload_artifacts = MagicMock()
|
||||
async def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
mlflow_repository.dowload_artifacts = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', model_type[0], 'pyfunc', True)
|
||||
result = await mlflow_repository.download_model('test_model', {}, model_type[0], 'pyfunc', True)
|
||||
|
||||
mlflow_repository.dowload_artifacts.assert_called_once_with('test_model', model_type[1])
|
||||
mlflow_repository.dowload_artifacts.assert_called_once_with('test_model', {}, model_type[1])
|
||||
|
||||
mlflow.pyfunc.load_model.assert_called_once_with(
|
||||
mlflow_repository.dowload_artifacts.return_value
|
||||
@@ -354,26 +442,28 @@ def test_download_model_load_wrapper(mlflow, mlflow_repository, model_type):
|
||||
)
|
||||
|
||||
|
||||
def test_download_model_predict(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = MagicMock()
|
||||
mlflow_repository.load_transform_model = MagicMock()
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_predict(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = AsyncMock()
|
||||
mlflow_repository.load_transform_model = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', 'predict', 'pyfunc', False)
|
||||
result = await mlflow_repository.download_model('test_model', {}, 'predict', 'pyfunc', False)
|
||||
|
||||
mlflow_repository.load_predict_model.assert_called_once_with('test_model', 'pyfunc')
|
||||
mlflow_repository.load_predict_model.assert_called_once_with('test_model', {}, 'pyfunc')
|
||||
mlflow_repository.load_transform_model.assert_not_called()
|
||||
|
||||
assert result == (mlflow_repository.load_predict_model.return_value, None)
|
||||
|
||||
|
||||
def test_download_model_transform(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = MagicMock()
|
||||
mlflow_repository.load_transform_model = MagicMock()
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_model_transform(mlflow_repository):
|
||||
mlflow_repository.load_predict_model = AsyncMock()
|
||||
mlflow_repository.load_transform_model = AsyncMock()
|
||||
|
||||
result = mlflow_repository.download_model('test_model', 'transform', 'pyfunc', False)
|
||||
result = await mlflow_repository.download_model('test_model', {}, 'transform', 'pyfunc', False)
|
||||
|
||||
mlflow_repository.load_predict_model.assert_not_called()
|
||||
mlflow_repository.load_transform_model.assert_called_once_with('test_model', 'pyfunc')
|
||||
mlflow_repository.load_transform_model.assert_called_once_with('test_model', {}, 'pyfunc')
|
||||
|
||||
assert result == (mlflow_repository.load_transform_model.return_value, None)
|
||||
|
||||
@@ -481,19 +571,25 @@ def test_handle_outdated_model(mlflow_repository):
|
||||
assert mlflow_repository.model_cache == {}
|
||||
|
||||
|
||||
def test_get_model_retention_0(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_retention_0(mlflow_repository):
|
||||
model = MagicMock()
|
||||
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
output = mlflow_repository.get_model('model_name', 0, 'predict', 'pyfunc')
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
output = await mlflow_repository.get_model('model_name', {}, 0, 'predict', 'pyfunc')
|
||||
|
||||
assert output == model
|
||||
mlflow_repository.download_model.assert_called_once_with(
|
||||
model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=False
|
||||
model_name='model_name',
|
||||
metadata={},
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=False,
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_cached_valid(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_cached_valid(mlflow_repository):
|
||||
mlflow_repository.check_cache_retention = MagicMock(return_value=True)
|
||||
mlflow_repository.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
@@ -503,7 +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
|
||||
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()
|
||||
|
||||
|
||||
def test_get_model_cached_outdated(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_cached_outdated(mlflow_repository):
|
||||
mlflow_repository.check_cache_retention = MagicMock(return_value=False)
|
||||
mlflow_repository.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
cache = {
|
||||
'model_name_predict': {
|
||||
'target': 'cached_model',
|
||||
@@ -530,7 +627,7 @@ def test_get_model_cached_outdated(mlflow_repository):
|
||||
}
|
||||
mlflow_repository.model_cache = cache
|
||||
|
||||
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc')
|
||||
output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
|
||||
assert output == model
|
||||
mlflow_repository.check_cache_retention.assert_called_once_with(
|
||||
{
|
||||
@@ -544,14 +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.handle_valid_model = MagicMock()
|
||||
mlflow_repository.handle_outdated_model = MagicMock()
|
||||
mlflow_repository.model_cache = {}
|
||||
model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(return_value=(model, 'artifact_path'))
|
||||
output = mlflow_repository.get_model('model_name', 1, 'predict', 'pyfunc')
|
||||
mlflow_repository.download_model = AsyncMock(return_value=(model, 'artifact_path'))
|
||||
output = await mlflow_repository.get_model('model_name', {}, 1, 'predict', 'pyfunc')
|
||||
assert output == model
|
||||
mlflow_repository.check_cache_retention.assert_not_called()
|
||||
mlflow_repository.handle_valid_model.assert_not_called()
|
||||
@@ -559,43 +657,53 @@ def test_get_model_cached_not_found(mlflow_repository):
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_operation_retention_0(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_operation_retention_0(force_memory_release, mlflow_repository):
|
||||
model = MagicMock()
|
||||
data = MagicMock()
|
||||
mlflow_repository.get_model = MagicMock(return_value=model)
|
||||
output = mlflow_repository.get_cached_operation('model_name', data, 'transform', 0, 'sklearn')
|
||||
mlflow_repository.get_model = AsyncMock(return_value=model)
|
||||
output = await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'transform', 0, 'sklearn', {}
|
||||
)
|
||||
assert output == model.predict.return_value
|
||||
force_memory_release.assert_called_once_with(mlflow_repository.logger)
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_predict_retention_not_0(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_predict_retention_not_0(force_memory_release, mlflow_repository):
|
||||
model = MagicMock()
|
||||
data = MagicMock()
|
||||
mlflow_repository.get_model = MagicMock(return_value=model)
|
||||
output = mlflow_repository.get_cached_operation('model_name', data, 'predict', 1, 'sklearn')
|
||||
mlflow_repository.get_model = AsyncMock(return_value=model)
|
||||
output = await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'predict', 1, 'sklearn', {}
|
||||
)
|
||||
assert output == model.predict.return_value
|
||||
force_memory_release.assert_not_called()
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
def test_get_cached_operation_invalid_operation(force_memory_release, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_operation_invalid_operation(force_memory_release, mlflow_repository):
|
||||
data = MagicMock()
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.get_cached_operation('model_name', data, 'invalid', 0, 'sklearn')
|
||||
await mlflow_repository.get_cached_operation(
|
||||
'model_name', data, 'invalid', 0, 'sklearn', {}
|
||||
)
|
||||
assert str(e) == "Invalid operation. Use 'transform' or 'predict'."
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.pd.merge')
|
||||
@patch('laborious.utils.repository.model_repository.isinstance')
|
||||
def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
@pytest.mark.asyncio
|
||||
async def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
isinstance_mock, pd_merge, mlflow_repository
|
||||
):
|
||||
isinstance_mock.return_value = False
|
||||
|
||||
data_model = MagicMock()
|
||||
prediction_model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(
|
||||
mlflow_repository.download_model = AsyncMock(
|
||||
side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
|
||||
)
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
|
||||
@@ -604,7 +712,7 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
output = mlflow_repository.fit_models(
|
||||
output = await mlflow_repository.fit_models(
|
||||
'model_name', data, 'latest_production_id', metadata['metadata'], 'sklearn', 'pyfunc', None
|
||||
)
|
||||
|
||||
@@ -612,11 +720,18 @@ def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
[
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='transform',
|
||||
flavor='sklearn',
|
||||
load_wrapper=False,
|
||||
),
|
||||
call(model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=True),
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=True,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -658,7 +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.isinstance')
|
||||
def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
@pytest.mark.asyncio
|
||||
async def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
isinstance_mock, pd_merge, mlflow_repository
|
||||
):
|
||||
isinstance_mock.return_value = True
|
||||
@@ -667,7 +783,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
target_variable='feat_2',
|
||||
)
|
||||
prediction_model = MagicMock()
|
||||
mlflow_repository.download_model = MagicMock(
|
||||
mlflow_repository.download_model = AsyncMock(
|
||||
side_effect=[(data_model, 'artifact_path'), (prediction_model, 'artifact_path')],
|
||||
)
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
|
||||
@@ -678,7 +794,7 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
output = mlflow_repository.fit_models(
|
||||
output = await mlflow_repository.fit_models(
|
||||
'model_name',
|
||||
data,
|
||||
'latest_production_id',
|
||||
@@ -692,11 +808,18 @@ def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
[
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='transform',
|
||||
flavor='sklearn',
|
||||
load_wrapper=False,
|
||||
),
|
||||
call(model_name='model_name', model_type='predict', flavor='pyfunc', load_wrapper=True),
|
||||
call(
|
||||
model_name='model_name',
|
||||
metadata=metadata['metadata'],
|
||||
model_type='predict',
|
||||
flavor='pyfunc',
|
||||
load_wrapper=True,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -728,16 +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'}
|
||||
mlflow_repository.log_model(model_data, 'sklearn', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'sklearn', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
mlflow.sklearn.log_model.assert_called_once_with(model_data['model'], 'prediction_model')
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
mlflow_repository.log_model(model_data, 'pyfunc', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'pyfunc', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow.pyfunc.log_model.assert_not_called()
|
||||
|
||||
@@ -747,24 +881,47 @@ def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
artifact_path='prediction_model', code_path=[path.join.return_value], to_disk=False
|
||||
)
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_log_model_pytorch(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_pytorch(mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
mlflow_repository.log_model(model_data, 'pytorch', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'pytorch', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
mlflow.pytorch.log_model.assert_called_once_with(model_data['model'], 'prediction_model')
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_log_model_error(mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_error(mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.log_model(model_data, 'invalid', 'prediction_model', metadata['metadata'])
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'invalid', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
assert str(e) == "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_new_experiment(
|
||||
_rmtree, path, force_memory_release, mlflow, mlflow_repository
|
||||
):
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
retrain_data = {
|
||||
@@ -782,11 +939,11 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
|
||||
|
||||
mlflow_repository.get_experiment = MagicMock()
|
||||
mlflow_repository.get_next_run_name = MagicMock()
|
||||
mlflow_repository.log_model = MagicMock()
|
||||
mlflow_repository.log_model = AsyncMock()
|
||||
path.exists.return_value = True
|
||||
path.join.return_value = './tmp/artifacts/model_name'
|
||||
|
||||
report = mlflow_repository.create_new_experiment(
|
||||
report = await mlflow_repository.create_new_experiment(
|
||||
model_name,
|
||||
data,
|
||||
retrain_data,
|
||||
@@ -843,8 +1000,60 @@ def test_create_new_experiment(_rmtree, path, force_memory_release, mlflow, mlfl
|
||||
'experiment_name': mlflow_repository.get_experiment.return_value.name,
|
||||
}
|
||||
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_WRITE_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.force_memory_release')
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_new_experiment_error(
|
||||
_rmtree, path, force_memory_release, mlflow, mlflow_repository
|
||||
):
|
||||
mlflow.start_run.side_effect = ValueError('error')
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
retrain_data = {
|
||||
'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
}
|
||||
|
||||
mlflow_repository.get_model_params = MagicMock(
|
||||
return_value={
|
||||
'transform_flavor': 'sklearn',
|
||||
'predict_flavor': 'pyfunc',
|
||||
'target_name': 'target_name',
|
||||
}
|
||||
)
|
||||
|
||||
mlflow_repository.get_experiment = MagicMock()
|
||||
mlflow_repository.get_next_run_name = MagicMock()
|
||||
mlflow_repository.log_model = AsyncMock()
|
||||
path.exists.return_value = True
|
||||
path.join.return_value = './tmp/artifacts/model_name'
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.create_new_experiment(
|
||||
model_name,
|
||||
data,
|
||||
retrain_data,
|
||||
'latest_production_id',
|
||||
metadata['metadata'],
|
||||
'sklearn',
|
||||
'pyfunc',
|
||||
)
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
mlflow_repository.client.get_registered_model.return_value = MagicMock(
|
||||
latest_versions=[
|
||||
MagicMock(version='1'),
|
||||
@@ -852,7 +1061,9 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
MagicMock(version='3'),
|
||||
]
|
||||
)
|
||||
output = mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
output = await mlflow_repository.update_production_model_by_run_id(
|
||||
'0', 'test', metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow.register_model.assert_called_once_with(
|
||||
'runs:/0/prediction_model',
|
||||
@@ -873,21 +1084,62 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
'mlflow_run_id': '0',
|
||||
}
|
||||
|
||||
mlflow_repository.observe_lag.assert_has_calls(
|
||||
[
|
||||
call(ANY, metrics.MODEL_WRITE_LAG, ANY),
|
||||
call(ANY, metrics.MODEL_WRITE_LAG, ANY),
|
||||
]
|
||||
)
|
||||
mlflow_repository.emit_metric.assert_has_calls(
|
||||
[
|
||||
call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY),
|
||||
call(metric_object=metrics.MODEL_WRITE_COUNT, tags=ANY),
|
||||
]
|
||||
)
|
||||
|
||||
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error_register_model(mlflow, mlflow_repository):
|
||||
mlflow.register_model.side_effect = ValueError('error')
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error_transition_model_version_stage(
|
||||
mlflow, mlflow_repository
|
||||
):
|
||||
mlflow_repository.client.transition_model_version_stage.side_effect = ValueError('error')
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_WRITE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
||||
mlflow_repository.client.get_registered_model.return_value = MagicMock(
|
||||
get_registered_model=MagicMock(return_value=MagicMock(latest_versions={}))
|
||||
)
|
||||
|
||||
try:
|
||||
mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
await mlflow_repository.update_production_model_by_run_id('0', 'test', metadata['metadata'])
|
||||
except Exception as e:
|
||||
assert str(e) == 'Model versions is not a list'
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
|
||||
def test_transform_success(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_success(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'model'
|
||||
model_config = {
|
||||
@@ -896,14 +1148,19 @@ def test_transform_success(mlflow_repository):
|
||||
'predict_flavor': 'pyfunc',
|
||||
}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock()
|
||||
mlflow_repository.get_cached_operation = AsyncMock(return_value=data)
|
||||
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'transform', 60, 'sklearn'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='transform',
|
||||
retention=60,
|
||||
flavor='sklearn',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
|
||||
@@ -916,7 +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()
|
||||
model_name = 'model'
|
||||
model_config = {
|
||||
@@ -925,27 +1183,40 @@ def test_transform_error(mlflow_repository):
|
||||
'predict_flavor': 'pyfunc',
|
||||
}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(side_effect=Exception('error'))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error'))
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.transform(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'transform', 60, 'sklearn'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='transform',
|
||||
retention=60,
|
||||
flavor='sklearn',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
|
||||
|
||||
|
||||
def test_predict_success_array(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_predict_success_array(mlflow_repository):
|
||||
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
model_name = 'model'
|
||||
mlflow_repository.get_cached_operation = MagicMock(return_value=np.array([2, 3]))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(
|
||||
return_value=DataFrame({'prediction': {'index_1': 2, 'index_2': 3}})
|
||||
)
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output['success'] is True
|
||||
@@ -955,19 +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}})
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
model_name = 'model'
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(
|
||||
mlflow_repository.get_cached_operation = AsyncMock(
|
||||
return_value=DataFrame({'feat_1': {'index_3': 2, 'index_4': 3}})
|
||||
)
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output['success'] is True
|
||||
@@ -977,32 +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}})
|
||||
model_name = 'model'
|
||||
model_config = {'retention_minutes': 60, 'predict_flavor': 'pyfunc'}
|
||||
|
||||
mlflow_repository.get_cached_operation = MagicMock(side_effect=Exception('error'))
|
||||
mlflow_repository.get_cached_operation = AsyncMock(side_effect=Exception('error'))
|
||||
|
||||
output = mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.predict(model_name, data, model_config, metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_cached_operation.assert_called_once_with(
|
||||
model_name, data, 'predict', 60, 'pyfunc'
|
||||
model_name=model_name,
|
||||
data=data,
|
||||
operation='predict',
|
||||
retention=60,
|
||||
flavor='pyfunc',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
|
||||
|
||||
|
||||
def test_retrain_model(mlflow_repository):
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrain_model(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'test'
|
||||
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
|
||||
|
||||
mlflow_repository.get_model_run_id = MagicMock()
|
||||
mlflow_repository.fit_models = MagicMock()
|
||||
mlflow_repository.create_new_experiment = MagicMock()
|
||||
mlflow_repository.fit_models = AsyncMock()
|
||||
mlflow_repository.create_new_experiment = AsyncMock()
|
||||
|
||||
output = mlflow_repository.retrain_model(data, model_name, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.retrain_model(
|
||||
data, model_name, model_config, metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(model_name, stage='Production')
|
||||
|
||||
@@ -1033,12 +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()
|
||||
model_name = 'test'
|
||||
model_config = {'target': 'target', 'transform_flavor': 'sklearn', 'predict_flavor': 'pyfunc'}
|
||||
mlflow_repository.get_model_run_id = MagicMock(side_effect=Exception('error'))
|
||||
output = mlflow_repository.retrain_model(data, model_name, model_config, metadata['metadata'])
|
||||
output = await mlflow_repository.retrain_model(
|
||||
data, model_name, model_config, metadata['metadata']
|
||||
)
|
||||
assert output == {
|
||||
'success': False,
|
||||
'experiment': None,
|
||||
@@ -1047,17 +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'}
|
||||
model_name = 'test'
|
||||
mlflow_repository.update_production_model_by_run_id = MagicMock()
|
||||
mlflow_repository.update_production_model_by_run_id = AsyncMock()
|
||||
mlflow_repository.update_production_model_by_run_id.return_value = {
|
||||
'model_name': 'test',
|
||||
'version': '3',
|
||||
'mlflow_run_id': '0',
|
||||
}
|
||||
|
||||
output = mlflow_repository.update_production_model(experiment, model_name, metadata['metadata'])
|
||||
output = await mlflow_repository.update_production_model(
|
||||
experiment, model_name, metadata['metadata']
|
||||
)
|
||||
|
||||
mlflow_repository.update_production_model_by_run_id.assert_called_once_with(
|
||||
'0', 'test', metadata['metadata']
|
||||
|
||||
@@ -26,8 +26,12 @@ def opc_repository(mock_logger):
|
||||
cert_path='/path/to/cert.pem',
|
||||
private_key_path='/path/to/key.pem',
|
||||
server_cert_path='/path/to/server_cert.pem',
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
repository.disconnection_interval = 0.1
|
||||
repository.send_notification = MagicMock()
|
||||
repository.send_notification_async = AsyncMock()
|
||||
repository.emit_metric = AsyncMock()
|
||||
return repository
|
||||
|
||||
|
||||
@@ -214,7 +218,7 @@ async def test_disconnect_error(opc_repository, mock_client):
|
||||
await opc_repository.disconnect()
|
||||
|
||||
opc_repository.disconnection_fallback.assert_called_once()
|
||||
opc_repository.send_notification.assert_called_once_with(
|
||||
opc_repository.send_notification_async.assert_called_once_with(
|
||||
metadata=opc_repository.metadata,
|
||||
notification_id=f'OPC_DISCONNECTION_ERROR_{opc_repository.id}',
|
||||
message='Failed to disconnect from OPC server in 5 attempts.',
|
||||
|
||||
Reference in New Issue
Block a user