SIENTIAPDE-1646
Update README, requirements, and E2E tests for improved configuration and functionality - Enhanced the README with updated model configuration examples, including the addition of an alias for production. - Removed the `requirements-light.txt` file and updated `requirements-local.txt` and `requirements.txt` to replace `asyncua` with `opcua`. - Refactored E2E test scenarios to utilize scenario input files for better maintainability and clarity. - Improved test coverage for MinIO offload functionality and added new helper functions for loading scenario inputs. - Updated `values.yaml` to reflect new global configurations and environment variables for the laborious worker.
This commit is contained in:
@@ -2,6 +2,19 @@ import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from sientia_do.temporal.activities.postgres_sync import Postgres
|
||||
|
||||
|
||||
def _noop_postgres_del(_self):
|
||||
"""
|
||||
Unit tests use MagicMock metrics controllers; postgres_sync.Postgres.__del__ calls
|
||||
close() during GC and triggers async shutdown. Explicit ``close()`` is covered in tests.
|
||||
"""
|
||||
return None
|
||||
|
||||
|
||||
Postgres.__del__ = _noop_postgres_del # type: ignore[method-assign]
|
||||
|
||||
# The production code converts SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES to int at import-time.
|
||||
# Tests must set it to a valid integer string to avoid import errors.
|
||||
os.environ.setdefault('SIENTIA_MINIO_OFFLOAD_THRESHOLD_MEGABYTES', '1')
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
from pytest import mark
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
from laborious.activities.activities import Activities
|
||||
from laborious.activities.api import API
|
||||
@@ -156,7 +154,6 @@ def test___init__(
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.activities.Storage')
|
||||
@patch('laborious.activities.activities.MLFlow')
|
||||
@patch('laborious.activities.activities.OPC')
|
||||
@@ -164,7 +161,7 @@ def test___init__(
|
||||
@patch('laborious.activities.activities.ModelMetrics')
|
||||
@patch('laborious.activities.activities.API')
|
||||
@patch('laborious.activities.activities.MinioRepository')
|
||||
async def test_shutdown(
|
||||
def test_shutdown(
|
||||
_mock_minio_repository,
|
||||
mock_api_init,
|
||||
mock_model_metrics_init,
|
||||
@@ -173,7 +170,7 @@ async def test_shutdown(
|
||||
mock_mlflow_init,
|
||||
mock_storage_init,
|
||||
):
|
||||
mock_opc_init.close = AsyncMock()
|
||||
mock_opc_init.close = MagicMock()
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -222,7 +219,7 @@ async def test_shutdown(
|
||||
mlflow_repository=mlflow_repository,
|
||||
)
|
||||
|
||||
await activities.shutdown()
|
||||
activities.shutdown()
|
||||
mock_opc_init.close.assert_called_once()
|
||||
mock_storage_init.close.assert_called_once()
|
||||
mock_mlflow_init.close.assert_called_once()
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
|
||||
import pytest_asyncio
|
||||
from pytest import fixture, mark
|
||||
from pytest import fixture
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
from laborious.activities.api import API, PI_WEB_API_PREDICTION_ERROR_CONFIDENCE
|
||||
@@ -71,7 +70,7 @@ def test_get_pi_web_api_core_labels_without_operation_type(mock_pi_web_api_clien
|
||||
auth_token='test_token',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
with patch.object(
|
||||
SientiaMonitoring,
|
||||
@@ -105,7 +104,7 @@ def test_get_pi_web_api_core_labels_with_operation_type(mock_pi_web_api_client):
|
||||
auth_token='test_token',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
with patch.object(
|
||||
SientiaMonitoring,
|
||||
@@ -132,17 +131,17 @@ def test__init__():
|
||||
auth_token='test_token',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
|
||||
assert api.pi_web_api_client is not None
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
@fixture
|
||||
@patch('laborious.activities.api.PIWebAPIClient')
|
||||
def api(mock_pi_web_api_client):
|
||||
mock_client = MagicMock()
|
||||
mock_client.write_value = AsyncMock()
|
||||
mock_client.write_value = MagicMock()
|
||||
mock_client.close = MagicMock()
|
||||
mock_client.base_url = 'https://test-pi-server.com'
|
||||
mock_pi_web_api_client.return_value = mock_client
|
||||
@@ -153,12 +152,12 @@ def api(mock_pi_web_api_client):
|
||||
auth_token='test_token',
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
api_instance.send_notification_async = AsyncMock()
|
||||
api_instance.send_notification = MagicMock()
|
||||
api_instance.info = MagicMock()
|
||||
api_instance.error = MagicMock()
|
||||
api_instance.emit_metric = AsyncMock()
|
||||
api_instance.emit_metric_sync = MagicMock()
|
||||
api_instance.get_core_labels = MagicMock(
|
||||
return_value={
|
||||
'pod_id': 'test_pod',
|
||||
@@ -170,9 +169,8 @@ def api(mock_pi_web_api_client):
|
||||
return api_instance
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
|
||||
def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_data):
|
||||
input_data = {
|
||||
**base_input_data,
|
||||
'pi_web_api_output_config': {
|
||||
@@ -190,7 +188,7 @@ async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_dat
|
||||
[{'WebId': 'web_id_3', 'Errors': []}, {'WebId': 'web_id_4', 'Errors': []}],
|
||||
]
|
||||
|
||||
result = await api.write_pi_web_api_data(input_data)
|
||||
result = api.write_pi_web_api_data(input_data)
|
||||
|
||||
api.pi_web_api_client.write_value.assert_has_calls(
|
||||
[
|
||||
@@ -220,9 +218,8 @@ async def test_write_pi_web_api_data_success(mock_dataframe, api, base_input_dat
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data):
|
||||
def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_input_data):
|
||||
mock_dataframe.return_value = _create_mock_dataframe(
|
||||
{
|
||||
'prediction': [0.75],
|
||||
@@ -233,9 +230,9 @@ async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_
|
||||
|
||||
api.pi_web_api_client.write_value.side_effect = Exception('Prediction write failed')
|
||||
|
||||
result = await api.write_pi_web_api_data(base_input_data)
|
||||
result = api.write_pi_web_api_data(base_input_data)
|
||||
|
||||
api.send_notification_async.assert_called_once_with(
|
||||
api.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='WRITE_PI_WEB_API_PREDICTION_ERROR',
|
||||
message="Error writing prediction data to PI Web API: Prediction write failed\n Tags: {'tag1': 'web_id_1'}",
|
||||
@@ -248,9 +245,8 @@ async def test_write_pi_web_api_data_prediction_error(mock_dataframe, api, base_
|
||||
assert api.pi_web_api_client.write_value.call_count == 1
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data):
|
||||
def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_input_data):
|
||||
mock_dataframe.return_value = _create_mock_dataframe()
|
||||
|
||||
# First call succeeds, second fails
|
||||
@@ -259,9 +255,9 @@ async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_
|
||||
Exception('Confidence write failed'),
|
||||
]
|
||||
|
||||
result = await api.write_pi_web_api_data(base_input_data)
|
||||
result = api.write_pi_web_api_data(base_input_data)
|
||||
|
||||
api.send_notification_async.assert_called_once_with(
|
||||
api.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='WRITE_PI_WEB_API_CONFIDENCE_ERROR',
|
||||
message="Error writing confidence data to PI Web API: Confidence write failed\n Tags: {'tag2': 'web_id_2'}",
|
||||
@@ -278,9 +274,8 @@ async def test_write_pi_web_api_data_confidence_error(mock_dataframe, api, base_
|
||||
assert api.pi_web_api_client.write_value.call_count == 2
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
|
||||
def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_data):
|
||||
input_data = {
|
||||
**base_input_data,
|
||||
'pi_web_api_output_config': {
|
||||
@@ -298,7 +293,7 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_
|
||||
[],
|
||||
]
|
||||
|
||||
result = await api.write_pi_web_api_data(input_data)
|
||||
result = api.write_pi_web_api_data(input_data)
|
||||
|
||||
api.pi_web_api_client.write_value.assert_has_calls(
|
||||
[
|
||||
@@ -328,9 +323,8 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_updates_confidence_and_comments(
|
||||
def test_write_pi_web_api_data_updates_confidence_and_comments(
|
||||
mock_dataframe, api, base_input_data
|
||||
):
|
||||
mock_dataframe.return_value = _create_mock_dataframe()
|
||||
@@ -341,23 +335,23 @@ async def test_write_pi_web_api_data_updates_confidence_and_comments(
|
||||
with patch.object(
|
||||
api,
|
||||
'process_pi_web_api_response',
|
||||
new=AsyncMock(side_effect=[(0.33, 'PI warning'), (0, '')]),
|
||||
new=MagicMock(side_effect=[(0.33, 'PI warning'), (0, '')]),
|
||||
) as process_mock:
|
||||
result = await api.write_pi_web_api_data(base_input_data)
|
||||
result = api.write_pi_web_api_data(base_input_data)
|
||||
|
||||
assert process_mock.await_count == 2
|
||||
assert process_mock.call_count == 2
|
||||
assert result is not None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_close(api):
|
||||
@patch('laborious.activities.api.SientiaMonitoring.shutdown')
|
||||
def test_close(mock_shutdown, api):
|
||||
api.close()
|
||||
|
||||
api.pi_web_api_client.close.assert_called_once()
|
||||
mock_shutdown.assert_called_once_with(api)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_process_pi_web_api_response_success(api):
|
||||
def test_process_pi_web_api_response_success(api):
|
||||
"""Test successful processing of PI Web API response with all tags written."""
|
||||
response_data = [
|
||||
{'WebId': 'web_id_1', 'Errors': []},
|
||||
@@ -371,7 +365,7 @@ async def test_process_pi_web_api_response_success(api):
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
|
||||
confidence, message = await api.process_pi_web_api_response(
|
||||
confidence, message = api.process_pi_web_api_response(
|
||||
response_data=response_data,
|
||||
tags=tags,
|
||||
core_labels=core_labels,
|
||||
@@ -380,9 +374,9 @@ async def test_process_pi_web_api_response_success(api):
|
||||
|
||||
assert confidence == 0
|
||||
assert message == ''
|
||||
assert api.emit_metric.call_count == 2
|
||||
# Verify that emit_metric was called with correct tags structure
|
||||
call_args_list = api.emit_metric.call_args_list
|
||||
assert api.emit_metric_sync.call_count == 2
|
||||
# Verify that emit_metric_sync was called with correct tags structure
|
||||
call_args_list = api.emit_metric_sync.call_args_list
|
||||
assert len(call_args_list) == 2
|
||||
# Check that all calls include core_labels and tag_name
|
||||
for call_args in call_args_list:
|
||||
@@ -390,8 +384,7 @@ async def test_process_pi_web_api_response_success(api):
|
||||
assert call_args.kwargs['tags']['tag_name'] in ['tag1', 'tag2']
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_process_pi_web_api_response_with_errors(api):
|
||||
def test_process_pi_web_api_response_with_errors(api):
|
||||
"""Test processing response with errors in some tags."""
|
||||
response_data = [
|
||||
{'WebId': 'web_id_1', 'Errors': ['Error writing tag']},
|
||||
@@ -405,7 +398,7 @@ async def test_process_pi_web_api_response_with_errors(api):
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
|
||||
confidence, message = await api.process_pi_web_api_response(
|
||||
confidence, message = api.process_pi_web_api_response(
|
||||
response_data=response_data,
|
||||
tags=tags,
|
||||
core_labels=core_labels,
|
||||
@@ -417,11 +410,10 @@ async def test_process_pi_web_api_response_with_errors(api):
|
||||
message
|
||||
== "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag2'] tags were written."
|
||||
)
|
||||
assert api.emit_metric.call_count == 2
|
||||
assert api.emit_metric_sync.call_count == 2
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_process_pi_web_api_response_missing_tags(api):
|
||||
def test_process_pi_web_api_response_missing_tags(api):
|
||||
"""Test processing response when number of written tags doesn't match expected."""
|
||||
response_data = [
|
||||
{'WebId': 'web_id_1', 'Errors': []},
|
||||
@@ -434,7 +426,7 @@ async def test_process_pi_web_api_response_missing_tags(api):
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
|
||||
confidence, message = await api.process_pi_web_api_response(
|
||||
confidence, message = api.process_pi_web_api_response(
|
||||
response_data=response_data,
|
||||
tags=tags,
|
||||
core_labels=core_labels,
|
||||
@@ -446,14 +438,13 @@ async def test_process_pi_web_api_response_missing_tags(api):
|
||||
message
|
||||
== "The number of written tags does not match the number of tag names: Expected ['tag1', 'tag2'] tags, but ['tag1'] tags were written."
|
||||
)
|
||||
api.send_notification_async.assert_called_once()
|
||||
call_args = api.send_notification_async.call_args
|
||||
api.send_notification.assert_called_once()
|
||||
call_args = api.send_notification.call_args
|
||||
assert call_args.kwargs['notification_id'] == 'WRITE_PI_WEB_API_PREDICTION_ERROR'
|
||||
assert call_args.kwargs['level'] == NotificationLevel.ERROR
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_process_pi_web_api_response_missing_webid(api):
|
||||
def test_process_pi_web_api_response_missing_webid(api):
|
||||
"""Test processing response when WebId is missing in response item."""
|
||||
response_data = [
|
||||
{'Errors': []},
|
||||
@@ -467,7 +458,7 @@ async def test_process_pi_web_api_response_missing_webid(api):
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
|
||||
confidence, message = await api.process_pi_web_api_response(
|
||||
confidence, message = api.process_pi_web_api_response(
|
||||
response_data=response_data,
|
||||
tags=tags,
|
||||
core_labels=core_labels,
|
||||
@@ -482,8 +473,7 @@ async def test_process_pi_web_api_response_missing_webid(api):
|
||||
api.error.assert_any_call('The response did not contain some WebIds', metadata['metadata'])
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_process_pi_web_api_response_missing_tag_name(api):
|
||||
def test_process_pi_web_api_response_missing_tag_name(api):
|
||||
"""Test processing response when tag name is not found for WebId."""
|
||||
response_data = [
|
||||
{'WebId': 'unknown_web_id', 'Errors': []},
|
||||
@@ -496,7 +486,7 @@ async def test_process_pi_web_api_response_missing_tag_name(api):
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
|
||||
confidence, message = await api.process_pi_web_api_response(
|
||||
confidence, message = api.process_pi_web_api_response(
|
||||
response_data=response_data,
|
||||
tags=tags,
|
||||
core_labels=core_labels,
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
|
||||
from pandas import DataFrame
|
||||
from pytest import fixture, mark
|
||||
from pytest import fixture
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
|
||||
from laborious.activities.gates import Gates
|
||||
|
||||
@@ -17,17 +18,17 @@ def _passthrough_from_dict():
|
||||
|
||||
def _minio_payload(retrieve_return, status=None):
|
||||
"""
|
||||
Build a MinioDataFramePayload-like test double with async retrieve.
|
||||
Build a MinioDataFramePayload-like test double with retrieve.
|
||||
|
||||
Args:
|
||||
retrieve_return: Value returned from await retrieve(minio_repo, metadata).
|
||||
retrieve_return: Value returned from retrieve(minio_repo, metadata).
|
||||
status: Optional status dict for MLflow response gate (payload.status).
|
||||
|
||||
Return:
|
||||
MagicMock: Object with async retrieve and optional status.
|
||||
"""
|
||||
p = MagicMock()
|
||||
p.retrieve = AsyncMock(return_value=retrieve_return)
|
||||
p.retrieve = MagicMock(return_value=retrieve_return)
|
||||
p.status = status
|
||||
return p
|
||||
|
||||
@@ -37,7 +38,7 @@ def gates_activity():
|
||||
gates = Gates(
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
gates.error = MagicMock()
|
||||
gates.debug = MagicMock()
|
||||
@@ -45,8 +46,7 @@ def gates_activity():
|
||||
gates.warning = MagicMock()
|
||||
gates.critical = MagicMock()
|
||||
gates.send_notification = MagicMock()
|
||||
gates.send_notification_async = AsyncMock()
|
||||
gates.emit_metric = AsyncMock()
|
||||
gates.emit_metric_sync = MagicMock()
|
||||
return gates
|
||||
|
||||
|
||||
@@ -60,8 +60,7 @@ metadata = {
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_invalid_filter(gates_activity):
|
||||
def test_input_gate_invalid_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -71,7 +70,7 @@ async def test_input_gate_invalid_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
@@ -80,9 +79,8 @@ async def test_input_gate_invalid_filter(gates_activity):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.input_filter_functions')
|
||||
async def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity):
|
||||
def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity):
|
||||
# Arrange
|
||||
mock_input_filter_functions.__contains__.return_value = True
|
||||
mock_input_filter_functions.__getitem__.return_value = MagicMock(
|
||||
@@ -96,11 +94,11 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
gates_activity.send_notification.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",
|
||||
@@ -110,8 +108,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_no_filters(gates_activity):
|
||||
def test_input_gate_no_filters(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -121,15 +118,14 @@ async def test_input_gate_no_filters(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_with_filter(gates_activity):
|
||||
def test_input_gate_with_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -139,15 +135,14 @@ async def test_input_gate_with_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Input data with bad quality')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_with_filter_lowercase_keys(gates_activity):
|
||||
def test_input_gate_with_filter_lowercase_keys(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -157,14 +152,13 @@ async def test_input_gate_with_filter_lowercase_keys(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Input data with bad quality')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_with_filter_capitalized_keys(gates_activity):
|
||||
def test_input_gate_with_filter_capitalized_keys(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -174,14 +168,13 @@ async def test_input_gate_with_filter_capitalized_keys(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Input data with bad quality')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_input_gate_with_filter_not_caught(gates_activity):
|
||||
def test_input_gate_with_filter_not_caught(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -191,15 +184,14 @@ async def test_input_gate_with_filter_not_caught(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.input_gate(input_data)
|
||||
result = gates_activity.input_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_response_gate_invalid_filter(gates_activity):
|
||||
def test_mlflow_response_gate_invalid_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -213,15 +205,14 @@ async def test_mlflow_response_gate_invalid_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.mlflow_response_filter_functions')
|
||||
async def test_mlflow_response_gate_filter_exception(
|
||||
def test_mlflow_response_gate_filter_exception(
|
||||
mock_mlflow_response_filter_functions, gates_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -241,11 +232,11 @@ async def test_mlflow_response_gate_filter_exception(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
gates_activity.send_notification.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",
|
||||
@@ -255,8 +246,7 @@ async def test_mlflow_response_gate_filter_exception(
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_response_gate_no_filters(gates_activity):
|
||||
def test_mlflow_response_gate_no_filters(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -270,15 +260,14 @@ async def test_mlflow_response_gate_no_filters(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_response_gate_with_filter(gates_activity):
|
||||
def test_mlflow_response_gate_with_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -292,16 +281,15 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'API error occurred')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
|
||||
def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -315,14 +303,13 @@ async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity)
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'API error occurred')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
|
||||
def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -336,15 +323,14 @@ async def test_mlflow_response_gate_with_filter_not_caught(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_response_gate(input_data)
|
||||
result = gates_activity.mlflow_response_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_content_gate_invalid_filter(gates_activity):
|
||||
def test_mlflow_content_gate_invalid_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -355,17 +341,14 @@ async def test_mlflow_content_gate_invalid_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.mlflow_content_filter_functions')
|
||||
async def test_mlflow_content_gate_filter_exception(
|
||||
mock_mlflow_content_filter_functions, gates_activity
|
||||
):
|
||||
def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functions, gates_activity):
|
||||
# Arrange
|
||||
mock_mlflow_content_filter_functions.__contains__.return_value = True
|
||||
mock_mlflow_content_filter_functions.__getitem__.return_value = MagicMock(
|
||||
@@ -380,12 +363,12 @@ async def test_mlflow_content_gate_filter_exception(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification_async.assert_called_once_with(
|
||||
gates_activity.send_notification.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",
|
||||
@@ -395,8 +378,7 @@ async def test_mlflow_content_gate_filter_exception(
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_content_gate_no_filters(gates_activity):
|
||||
def test_mlflow_content_gate_no_filters(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -407,15 +389,14 @@ async def test_mlflow_content_gate_no_filters(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_content_gate_with_filter(gates_activity):
|
||||
def test_mlflow_content_gate_with_filter(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -426,16 +407,15 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == ('STOP', -1, 'Transformed data not passed the content filter')
|
||||
gates_activity.debug.assert_called()
|
||||
gates_activity.send_notification_async.assert_called()
|
||||
gates_activity.send_notification.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
|
||||
def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -446,15 +426,14 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_mlflow_content_gate_filter_returns_false(gates_activity):
|
||||
def test_mlflow_content_gate_filter_returns_false(gates_activity):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}},
|
||||
@@ -463,7 +442,7 @@ async def test_mlflow_content_gate_filter_returns_false(gates_activity):
|
||||
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
|
||||
}
|
||||
|
||||
result = await gates_activity.mlflow_content_gate(input_data)
|
||||
result = gates_activity.mlflow_content_gate(input_data)
|
||||
|
||||
assert result == (None, 0, '')
|
||||
gates_activity.debug.assert_called()
|
||||
@@ -525,8 +504,7 @@ def test_get_prediction_store_policy_valid_policy(gates_activity):
|
||||
assert policy_value == 1
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_prediction_no_timestamp(gates_activity):
|
||||
def test_format_prediction_no_timestamp(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -545,7 +523,7 @@ async def test_format_prediction_no_timestamp(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_prediction(input_data)
|
||||
result = gates_activity.format_prediction(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['prediction'] == {0: 1}
|
||||
@@ -557,8 +535,7 @@ async def test_format_prediction_no_timestamp(gates_activity):
|
||||
assert result['comments'] == {0: ''}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_prediction_with_timestamp_erl(gates_activity):
|
||||
def test_format_prediction_with_timestamp_erl(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -585,7 +562,7 @@ async def test_format_prediction_with_timestamp_erl(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_prediction(input_data)
|
||||
result = gates_activity.format_prediction(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['prediction'] == {0: 2, 1: 1}
|
||||
@@ -597,8 +574,7 @@ async def test_format_prediction_with_timestamp_erl(gates_activity):
|
||||
assert result['comments'] == {0: '', 1: ''}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_prediction_with_timestamp_lts(gates_activity):
|
||||
def test_format_prediction_with_timestamp_lts(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -625,7 +601,7 @@ async def test_format_prediction_with_timestamp_lts(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_prediction(input_data)
|
||||
result = gates_activity.format_prediction(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['prediction'] == {0: 3, 1: 2}
|
||||
@@ -637,8 +613,7 @@ async def test_format_prediction_with_timestamp_lts(gates_activity):
|
||||
assert result['comments'] == {0: '', 1: ''}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
|
||||
def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -663,16 +638,15 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
|
||||
gates_activity.get_prediction_store_policy = MagicMock(return_value=('invalid', 1))
|
||||
|
||||
try:
|
||||
await gates_activity.format_prediction(input_data)
|
||||
gates_activity.format_prediction(input_data)
|
||||
except ValueError as e:
|
||||
assert str(e) == 'Invalid policy type: invalid'
|
||||
else:
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
|
||||
async def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity):
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
|
||||
def test_format_transformed_data_single_row(mock_from_dataframe, gates_activity):
|
||||
# Arrange
|
||||
payload_result = MagicMock()
|
||||
mock_from_dataframe.return_value = payload_result
|
||||
@@ -691,7 +665,7 @@ async def test_format_transformed_data_single_row(mock_from_dataframe, gates_act
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
result = gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert result is payload_result
|
||||
@@ -705,9 +679,8 @@ async def test_format_transformed_data_single_row(mock_from_dataframe, gates_act
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
|
||||
async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity):
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
|
||||
def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_activity):
|
||||
# Arrange
|
||||
payload_result = MagicMock()
|
||||
mock_from_dataframe.return_value = payload_result
|
||||
@@ -732,7 +705,7 @@ async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
result = gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert result is payload_result
|
||||
@@ -746,9 +719,8 @@ async def test_format_transformed_data_multiple_rows(mock_from_dataframe, gates_
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=AsyncMock)
|
||||
async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity):
|
||||
@patch('laborious.activities.gates.MinioDataFramePayload.from_dataframe', new_callable=MagicMock)
|
||||
def test_format_transformed_data_empty_data(mock_from_dataframe, gates_activity):
|
||||
# Arrange
|
||||
payload_result = MagicMock()
|
||||
mock_from_dataframe.return_value = payload_result
|
||||
@@ -760,7 +732,7 @@ async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_act
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
result = gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert result is payload_result
|
||||
@@ -774,8 +746,7 @@ async def test_format_transformed_data_empty_data(mock_from_dataframe, gates_act
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_default_prediction(gates_activity):
|
||||
def test_format_default_prediction(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -786,7 +757,7 @@ async def test_format_default_prediction(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_default_prediction(input_data)
|
||||
result = gates_activity.format_default_prediction(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['prediction'] == {0: 0}
|
||||
@@ -799,8 +770,7 @@ async def test_format_default_prediction(gates_activity):
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_retrain_report(gates_activity):
|
||||
def test_format_retrain_report(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -819,7 +789,7 @@ async def test_format_retrain_report(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_retrain_report(input_data)
|
||||
result = gates_activity.format_retrain_report(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['model_id'] == {0: 'test_model'}
|
||||
@@ -831,8 +801,7 @@ async def test_format_retrain_report(gates_activity):
|
||||
assert result['mlflow_experiment_id'] == {0: 'test_mlflow_experiment_id'}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_retrain_report_failure(gates_activity):
|
||||
def test_format_retrain_report_failure(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -851,7 +820,7 @@ async def test_format_retrain_report_failure(gates_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_retrain_report(input_data)
|
||||
result = gates_activity.format_retrain_report(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['model_id'] == {0: 'test_model'}
|
||||
@@ -865,9 +834,8 @@ async def test_format_retrain_report_failure(gates_activity):
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.metrics')
|
||||
async def test_write_metrics(mock_metrics, gates_activity):
|
||||
def test_write_metrics(mock_metrics, gates_activity):
|
||||
"""Test write_metrics method."""
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -878,7 +846,7 @@ 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)
|
||||
gates_activity.write_metrics(input_data)
|
||||
core_tags = {
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'runtime': gates_activity.runtime,
|
||||
@@ -886,7 +854,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
}
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTIONS_WRITTEN_COUNT,
|
||||
@@ -894,7 +862,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_CONFIDENCE_MONITOR,
|
||||
@@ -904,7 +872,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_RESPONSE_TIME_MONITOR,
|
||||
@@ -914,7 +882,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
@@ -926,7 +894,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
@@ -940,7 +908,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_COUNT,
|
||||
@@ -952,7 +920,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
gates_activity.emit_metric.assert_has_calls(
|
||||
gates_activity.emit_metric_sync.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
@@ -968,9 +936,8 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.metrics')
|
||||
async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity):
|
||||
def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity):
|
||||
"""Test write_metrics method with None response_time in opc_metrics."""
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -981,10 +948,10 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act
|
||||
},
|
||||
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': None}},
|
||||
}
|
||||
await gates_activity.write_metrics(input_data)
|
||||
gates_activity.write_metrics(input_data)
|
||||
|
||||
# Verify that metrics for tag1 are emitted
|
||||
gates_activity.emit_metric.assert_any_call(
|
||||
gates_activity.emit_metric_sync.assert_any_call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
@@ -1002,7 +969,38 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act
|
||||
# Verify that metrics for tag2 (with None response_time) are NOT emitted
|
||||
calls = [
|
||||
c
|
||||
for c in gates_activity.emit_metric.call_args_list
|
||||
for c in gates_activity.emit_metric_sync.call_args_list
|
||||
if len(c[1].get('tags', {})) > 0 and c[1]['tags'].get('tag') == 'tag2'
|
||||
]
|
||||
assert len(calls) == 0, 'Metrics should not be emitted for None response_time'
|
||||
|
||||
|
||||
@patch.object(SientiaMonitoring, 'shutdown')
|
||||
def test_close_disposes_minio_repository(mock_shutdown):
|
||||
"""
|
||||
``Gates.close`` should close the optional MinIO client and clear the repository reference.
|
||||
"""
|
||||
minio = MagicMock()
|
||||
gates = Gates(
|
||||
minio_repository=minio,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
gates.close()
|
||||
minio.close.assert_called_once()
|
||||
assert gates.minio_repository is None
|
||||
mock_shutdown.assert_called_once_with(gates)
|
||||
|
||||
|
||||
@patch.object(SientiaMonitoring, 'shutdown')
|
||||
def test_close_without_minio_repository(mock_shutdown):
|
||||
"""When no MinIO repository is configured, ``close`` only shuts down monitoring."""
|
||||
gates = Gates(
|
||||
minio_repository=None,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
gates.close()
|
||||
mock_shutdown.assert_called_once_with(gates)
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from pytest import fixture, mark, raises
|
||||
from pytest import fixture, raises
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||
|
||||
@@ -22,7 +22,7 @@ def _passthrough_from_dict():
|
||||
def test___init__(mock_minio_repository):
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
metrics_controller = MagicMock()
|
||||
mlflow_repo = MagicMock()
|
||||
plugin_store = MagicMock()
|
||||
|
||||
@@ -63,7 +63,7 @@ def test___init__(mock_minio_repository):
|
||||
def mlflow(mock_minio_repository):
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
metrics_controller = MagicMock()
|
||||
mlflow_repo = MagicMock()
|
||||
plugin_store = MagicMock()
|
||||
|
||||
@@ -85,11 +85,10 @@ def mlflow(mock_minio_repository):
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
mlflow.minio_repository = AsyncMock()
|
||||
mlflow.minio_repository = MagicMock()
|
||||
|
||||
mlflow.send_notification = MagicMock()
|
||||
mlflow.emit_metric = AsyncMock()
|
||||
mlflow.send_notification_async = AsyncMock()
|
||||
mlflow.emit_metric = MagicMock()
|
||||
mlflow.error = MagicMock()
|
||||
mlflow.debug = MagicMock()
|
||||
mlflow.info = MagicMock()
|
||||
@@ -150,12 +149,11 @@ def test_detect_and_parse_datetime_index_timestamp_with_tz_success(mlflow):
|
||||
assert out.index[0].endswith('+0000')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
async def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw = pd.DataFrame(
|
||||
{
|
||||
@@ -180,8 +178,8 @@ async def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
wrapper.transform.return_value = (out_df, {'meta': True})
|
||||
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
|
||||
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -190,7 +188,7 @@ async def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
response_data = await mlflow.request_transform(input_data)
|
||||
response_data = mlflow.request_transform(input_data)
|
||||
|
||||
mlflow.mlflow_repository.get_cached_model.assert_called_once_with(
|
||||
model_name='test_model',
|
||||
@@ -202,12 +200,11 @@ async def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
assert response_data == mock_from_dataframe.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
async def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow):
|
||||
def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow):
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw = pd.DataFrame(
|
||||
{
|
||||
@@ -223,8 +220,8 @@ async def test_request_transform_success_without_transform_meta(mock_from_datafr
|
||||
wrapper.transform.return_value = (out_df, {})
|
||||
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
|
||||
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -233,21 +230,20 @@ async def test_request_transform_success_without_transform_meta(mock_from_datafr
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
await mlflow.request_transform(input_data)
|
||||
mlflow.request_transform(input_data)
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
async def test_request_transform_failure(mock_from_dataframe, mlflow):
|
||||
def test_request_transform_failure(mock_from_dataframe, mlflow):
|
||||
mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('boom')
|
||||
|
||||
data_mock = MagicMock()
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=data_mock)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=data_mock)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -260,7 +256,7 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow):
|
||||
data_mock.drop_duplicates.return_value = data_mock
|
||||
data_mock.pivot.return_value = data_mock
|
||||
|
||||
await mlflow.request_transform(input_data)
|
||||
mlflow.request_transform(input_data)
|
||||
|
||||
mock_from_dataframe.assert_called_once_with(
|
||||
dataframe=None,
|
||||
@@ -274,13 +270,12 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
wrapper = MagicMock()
|
||||
pred_df = MagicMock()
|
||||
wrapper.predict.return_value = (pred_df, {})
|
||||
@@ -288,8 +283,8 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')])
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=data_mock)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=data_mock)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -301,7 +296,7 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
pred_df.columns = MagicMock()
|
||||
pred_df.__setitem__ = MagicMock()
|
||||
|
||||
response_data = await mlflow.request_predict(input_data)
|
||||
response_data = mlflow.request_predict(input_data)
|
||||
|
||||
data_mock.replace.assert_called_once_with(np.nan, None, inplace=True)
|
||||
mock_to_datetime.assert_called()
|
||||
@@ -310,15 +305,12 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
assert response_data == mock_from_dataframe.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_request_predict_success_dataframe_and_meta(
|
||||
mock_to_datetime, mock_from_dataframe, mlflow
|
||||
):
|
||||
def test_request_predict_success_dataframe_and_meta(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
wrapper = MagicMock()
|
||||
pred_df = pd.DataFrame({'raw': [0.3]})
|
||||
wrapper.predict.return_value = (pred_df, {'m': 1})
|
||||
@@ -326,8 +318,8 @@ async def test_request_predict_success_dataframe_and_meta(
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')])
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=data_mock)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=data_mock)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -336,25 +328,24 @@ async def test_request_predict_success_dataframe_and_meta(
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
await mlflow.request_predict(input_data)
|
||||
mlflow.request_predict(input_data)
|
||||
|
||||
assert list(pred_df.columns) == ['prediction', 'response_time']
|
||||
mlflow.info.assert_any_call("Wrapper predict metadata: {'m': 1}", metadata['metadata'])
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
)
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('predict boom')
|
||||
|
||||
data_mock = MagicMock()
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=data_mock)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=data_mock)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -363,7 +354,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
await mlflow.request_predict(input_data)
|
||||
mlflow.request_predict(input_data)
|
||||
|
||||
mock_from_dataframe.assert_called_once_with(
|
||||
dataframe=None,
|
||||
@@ -377,12 +368,11 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.mlflow.log_artifact')
|
||||
@patch('laborious.activities.mlflow.tempfile.mkdtemp')
|
||||
@patch('laborious.activities.mlflow.rmtree')
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_success_data_success_retrain(
|
||||
def test_retrain_model_success_data_success_retrain(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
@@ -408,10 +398,10 @@ async def test_retrain_model_success_data_success_retrain(
|
||||
}
|
||||
)
|
||||
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
@@ -428,12 +418,11 @@ async def test_retrain_model_success_data_success_retrain(
|
||||
assert response['experiment']['run_id'] == 'new-run'
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.mlflow.log_artifact')
|
||||
@patch('laborious.activities.mlflow.tempfile.mkdtemp')
|
||||
@patch('laborious.activities.mlflow.rmtree')
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_success_with_payload_data(
|
||||
def test_retrain_model_success_with_payload_data(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
@@ -448,8 +437,8 @@ async def test_retrain_model_success_with_payload_data(
|
||||
|
||||
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
|
||||
raw_data.__getitem__.return_value.max.return_value = 'ts'
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
pivoted = MagicMock()
|
||||
raw_data.sort_values.return_value = raw_data
|
||||
@@ -460,7 +449,7 @@ async def test_retrain_model_success_with_payload_data(
|
||||
pivoted.index = MagicMock()
|
||||
pivoted.__setitem__ = MagicMock()
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
@@ -472,12 +461,11 @@ async def test_retrain_model_success_with_payload_data(
|
||||
assert response['success'] is True
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.mlflow.log_artifact')
|
||||
@patch('laborious.activities.mlflow.tempfile.mkdtemp')
|
||||
@patch('laborious.activities.mlflow.rmtree')
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag(
|
||||
def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
@@ -498,10 +486,10 @@ async def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag(
|
||||
'value': [1.0, 2.0, 3.0, 4.0],
|
||||
}
|
||||
)
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
@@ -515,17 +503,16 @@ async def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag(
|
||||
assert response['success'] is True
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow):
|
||||
def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow):
|
||||
mv_alias = MagicMock(run_id='src')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
||||
mlflow.mlflow_repository.get_cached_model.side_effect = RuntimeError('retrain failed')
|
||||
|
||||
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
|
||||
raw_data.__getitem__.return_value.max.return_value = 'tsmax'
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
pivoted = MagicMock()
|
||||
raw_data.sort_values.return_value = raw_data
|
||||
@@ -536,7 +523,7 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
pivoted.index = MagicMock()
|
||||
pivoted.__setitem__ = MagicMock()
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
@@ -549,9 +536,8 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
assert 'retrain failed' in response['message']
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_retrain_model_data_error(mlflow):
|
||||
response = await mlflow.retrain_model(
|
||||
def test_retrain_model_data_error(mlflow):
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -565,8 +551,7 @@ async def test_retrain_model_data_error(mlflow):
|
||||
assert 'data' in response['message'].lower() or 'loading' in response['message'].lower()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_retrain_model_missing_target(mlflow):
|
||||
def test_retrain_model_missing_target(mlflow):
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw_data = pd.DataFrame(
|
||||
{
|
||||
@@ -575,10 +560,10 @@ async def test_retrain_model_missing_target(mlflow):
|
||||
'value': [1.0, 2.0],
|
||||
}
|
||||
)
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
response = mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
@@ -591,12 +576,11 @@ async def test_retrain_model_missing_target(mlflow):
|
||||
assert 'target' in response['message']
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_retrain_model_data_error_no_minio_repository(mlflow):
|
||||
def test_retrain_model_data_error_no_minio_repository(mlflow):
|
||||
mlflow.minio_repository = None
|
||||
|
||||
with raises(ValueError) as e:
|
||||
await mlflow.retrain_model(
|
||||
mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'object_key': 'test_object_key',
|
||||
@@ -610,8 +594,7 @@ async def test_retrain_model_data_error_no_minio_repository(mlflow):
|
||||
assert str(e.value) == 'Minio repository not initialized'
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_update_production_model(mlflow):
|
||||
def test_update_production_model(mlflow):
|
||||
mlflow.mlflow_repository._client.search_model_versions.return_value = [
|
||||
MagicMock(version='3', run_id='run-x'),
|
||||
MagicMock(version='2', run_id='run-x'),
|
||||
@@ -626,7 +609,7 @@ async def test_update_production_model(mlflow):
|
||||
'status': 'success',
|
||||
}
|
||||
|
||||
response = await mlflow.update_production_model(input_data)
|
||||
response = mlflow.update_production_model(input_data)
|
||||
|
||||
mlflow.mlflow_repository.promote_to_alias.assert_called_once_with(
|
||||
model_name='test_model',
|
||||
@@ -639,8 +622,7 @@ async def test_update_production_model(mlflow):
|
||||
assert response['version'] == '3'
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_update_production_model_error(mlflow):
|
||||
def test_update_production_model_error(mlflow):
|
||||
mlflow.mlflow_repository._client.search_model_versions.return_value = []
|
||||
|
||||
input_data = {
|
||||
@@ -653,9 +635,9 @@ async def test_update_production_model_error(mlflow):
|
||||
}
|
||||
|
||||
try:
|
||||
await mlflow.update_production_model(input_data)
|
||||
mlflow.update_production_model(input_data)
|
||||
except Exception:
|
||||
mlflow.send_notification_async.assert_called_once_with(
|
||||
mlflow.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
|
||||
message=ANY,
|
||||
@@ -667,9 +649,8 @@ async def test_update_production_model_error(mlflow):
|
||||
raise AssertionError('Expected exception')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -690,14 +671,13 @@ async def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
with patch('laborious.activities.mlflow.rmtree'):
|
||||
with patch('laborious.activities.mlflow.Path') as mp:
|
||||
mp.return_value.rglob.return_value = [MagicMock()]
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
result = mlflow.get_reference_data(input_data)
|
||||
|
||||
mock_reference_data.to_dict.assert_called_once_with(orient='records')
|
||||
assert result == mock_reference_data.to_dict.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_not_found(mlflow):
|
||||
def test_get_reference_data_not_found(mlflow):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -705,14 +685,13 @@ async def test_get_reference_data_not_found(mlflow):
|
||||
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.side_effect = Exception('missing')
|
||||
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
result = mlflow.get_reference_data(input_data)
|
||||
|
||||
mlflow.warning.assert_called()
|
||||
assert result is None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
||||
def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -724,13 +703,12 @@ async def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
||||
with patch('laborious.activities.mlflow.rmtree'):
|
||||
with patch('laborious.activities.mlflow.Path') as mp:
|
||||
mp.return_value.rglob.return_value = []
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
result = mlflow.get_reference_data(input_data)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_exception(mlflow):
|
||||
def test_get_reference_data_exception(mlflow):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -740,6 +718,6 @@ async def test_get_reference_data_exception(mlflow):
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail')
|
||||
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
result = mlflow.get_reference_data(input_data)
|
||||
|
||||
assert result is None
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
from pandas import DataFrame
|
||||
from pytest import fixture, mark
|
||||
from pytest import fixture
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
from laborious.activities.model_metrics import ModelMetrics
|
||||
@@ -12,7 +12,7 @@ def model_metrics_activity():
|
||||
model_metrics = ModelMetrics(
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
model_metrics.error = MagicMock()
|
||||
model_metrics.debug = MagicMock()
|
||||
@@ -20,8 +20,7 @@ def model_metrics_activity():
|
||||
model_metrics.warning = MagicMock()
|
||||
model_metrics.critical = MagicMock()
|
||||
model_metrics.send_notification = MagicMock()
|
||||
model_metrics.send_notification_async = AsyncMock()
|
||||
model_metrics.emit_metric = AsyncMock()
|
||||
model_metrics.emit_metric_sync = MagicMock()
|
||||
model_metrics.get_core_labels = MagicMock(
|
||||
return_value={
|
||||
'pod_id': 'test_pod',
|
||||
@@ -29,7 +28,7 @@ def model_metrics_activity():
|
||||
'workflow_name': 'test_workflow',
|
||||
}
|
||||
)
|
||||
model_metrics.observe_lag = AsyncMock()
|
||||
model_metrics.observe_lag_sync = MagicMock()
|
||||
model_metrics.pod_id = 'test_pod'
|
||||
return model_metrics
|
||||
|
||||
@@ -44,8 +43,7 @@ metadata = {
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||
def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
@@ -64,7 +62,7 @@ async def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||
|
||||
# Act & Assert
|
||||
try:
|
||||
await model_metrics_activity.calculate_drift(input_data)
|
||||
model_metrics_activity.calculate_drift(input_data)
|
||||
except ValueError as e:
|
||||
assert str(e) == 'Invalid chunk period: invalid, must be "min" or "s"'
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
@@ -74,10 +72,9 @@ async def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_with_reference_data(
|
||||
def test_calculate_drift_with_reference_data(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -107,7 +104,7 @@ async def test_calculate_drift_with_reference_data(
|
||||
}
|
||||
]
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame(
|
||||
{
|
||||
@@ -142,7 +139,7 @@ async def test_calculate_drift_with_reference_data(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, list)
|
||||
@@ -150,10 +147,18 @@ async def test_calculate_drift_with_reference_data(
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True)
|
||||
mock_drift_df.drop.assert_called_once_with(
|
||||
columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore'
|
||||
)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(
|
||||
columns={'metric': 'method', 'statistic': 'value'}, inplace=True
|
||||
columns={
|
||||
'metric': 'method',
|
||||
'statistic': 'value',
|
||||
'alert': 'drift',
|
||||
'chunk_index': 'chunk',
|
||||
'chunk_end_date': 'timestamp_end',
|
||||
}
|
||||
)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(
|
||||
subset=['timestamp', 'method', 'feature'], keep='first', inplace=True
|
||||
@@ -161,10 +166,9 @@ async def test_calculate_drift_with_reference_data(
|
||||
mock_drift_df.to_dict.assert_called_once_with(orient='records')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_without_reference_data(
|
||||
def test_calculate_drift_without_reference_data(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -194,7 +198,7 @@ async def test_calculate_drift_without_reference_data(
|
||||
}
|
||||
]
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df)
|
||||
|
||||
target_data_dict = {
|
||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||
@@ -232,13 +236,13 @@ async def test_calculate_drift_without_reference_data(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, list)
|
||||
assert result == mock_drift_df.to_dict.return_value
|
||||
model_metrics_activity.warning.assert_called()
|
||||
model_metrics_activity.send_notification_async.assert_called_once_with(
|
||||
model_metrics_activity.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING',
|
||||
message='Using 30% first rows of target data as reference data',
|
||||
@@ -247,10 +251,18 @@ async def test_calculate_drift_without_reference_data(
|
||||
attachment_content=ANY,
|
||||
)
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True)
|
||||
mock_drift_df.drop.assert_called_once_with(
|
||||
columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore'
|
||||
)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(
|
||||
columns={'metric': 'method', 'statistic': 'value'}, inplace=True
|
||||
columns={
|
||||
'metric': 'method',
|
||||
'statistic': 'value',
|
||||
'alert': 'drift',
|
||||
'chunk_index': 'chunk',
|
||||
'chunk_end_date': 'timestamp_end',
|
||||
}
|
||||
)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(
|
||||
subset=['timestamp', 'method', 'feature'], keep='first', inplace=True
|
||||
@@ -258,16 +270,13 @@ async def test_calculate_drift_without_reference_data(
|
||||
mock_drift_df.to_dict.assert_called_once_with(orient='records')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_empty_drift_df(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
def test_calculate_drift_empty_drift_df(mock_to_datetime, mock_dataframe, model_metrics_activity):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=DataFrame())
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=DataFrame())
|
||||
|
||||
reference_data = DataFrame(
|
||||
{
|
||||
@@ -302,7 +311,7 @@ async def test_calculate_drift_empty_drift_df(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
@@ -311,10 +320,9 @@ async def test_calculate_drift_empty_drift_df(
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_empty_after_timestamp_filter(
|
||||
def test_calculate_drift_empty_after_timestamp_filter(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -340,7 +348,7 @@ async def test_calculate_drift_empty_after_timestamp_filter(
|
||||
|
||||
mock_drift_df.__getitem__.side_effect = getitem_side_effect
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame(
|
||||
{
|
||||
@@ -375,7 +383,7 @@ async def test_calculate_drift_empty_after_timestamp_filter(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
@@ -383,17 +391,16 @@ async def test_calculate_drift_empty_after_timestamp_filter(
|
||||
'No drift metrics found after dropping rows where timestamp is not in target data',
|
||||
metadata['metadata'],
|
||||
)
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True)
|
||||
# When the timestamp filter empties the dataframe, the rename/drop pipeline
|
||||
# is short-circuited, so neither ``drop`` nor ``rename`` should run.
|
||||
mock_drift_df.drop.assert_not_called()
|
||||
mock_drift_df.rename.assert_not_called()
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_success_min(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
def test_calculate_drift_success_min(mock_to_datetime, mock_dataframe, model_metrics_activity):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = (
|
||||
@@ -421,7 +428,7 @@ async def test_calculate_drift_success_min(
|
||||
}
|
||||
]
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame(
|
||||
{
|
||||
@@ -456,7 +463,7 @@ async def test_calculate_drift_success_min(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, list)
|
||||
@@ -464,10 +471,18 @@ async def test_calculate_drift_success_min(
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True)
|
||||
mock_drift_df.drop.assert_called_once_with(
|
||||
columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore'
|
||||
)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(
|
||||
columns={'metric': 'method', 'statistic': 'value'}, inplace=True
|
||||
columns={
|
||||
'metric': 'method',
|
||||
'statistic': 'value',
|
||||
'alert': 'drift',
|
||||
'chunk_index': 'chunk',
|
||||
'chunk_end_date': 'timestamp_end',
|
||||
}
|
||||
)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(
|
||||
subset=['timestamp', 'method', 'feature'], keep='first', inplace=True
|
||||
@@ -475,10 +490,9 @@ async def test_calculate_drift_success_min(
|
||||
mock_drift_df.to_dict.assert_called_once_with(orient='records')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model_metrics_activity):
|
||||
def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model_metrics_activity):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = (
|
||||
@@ -506,7 +520,7 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model
|
||||
}
|
||||
]
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame(
|
||||
{
|
||||
@@ -541,7 +555,7 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, list)
|
||||
@@ -549,10 +563,18 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value'], inplace=True)
|
||||
mock_drift_df.drop.assert_called_once_with(
|
||||
columns=['p_value', 'chunk_start_date'], inplace=True, errors='ignore'
|
||||
)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(
|
||||
columns={'metric': 'method', 'statistic': 'value'}, inplace=True
|
||||
columns={
|
||||
'metric': 'method',
|
||||
'statistic': 'value',
|
||||
'alert': 'drift',
|
||||
'chunk_index': 'chunk',
|
||||
'chunk_end_date': 'timestamp_end',
|
||||
}
|
||||
)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(
|
||||
subset=['timestamp', 'method', 'feature'], keep='first', inplace=True
|
||||
@@ -560,16 +582,15 @@ async def test_calculate_drift_success_s(mock_to_datetime, mock_dataframe, model
|
||||
mock_drift_df.to_dict.assert_called_once_with(orient='records')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_get_drift_metrics_error(
|
||||
def test_calculate_drift_get_drift_metrics_error(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(
|
||||
model_metrics_activity.get_drift_metrics = MagicMock(
|
||||
side_effect=Exception('Get drift metrics error')
|
||||
)
|
||||
|
||||
@@ -606,14 +627,14 @@ async def test_calculate_drift_get_drift_metrics_error(
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
result = model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error getting drift metrics: Get drift metrics error', metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.send_notification_async.assert_called_once_with(
|
||||
model_metrics_activity.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR',
|
||||
message='Error getting drift metrics: Get drift metrics error',
|
||||
@@ -623,12 +644,11 @@ async def test_calculate_drift_get_drift_metrics_error(
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_success(
|
||||
def test_get_drift_metrics_success(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -668,7 +688,7 @@ async def test_get_drift_metrics_success(
|
||||
).columns
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.get_drift_metrics(
|
||||
result = model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
@@ -681,16 +701,15 @@ async def test_get_drift_metrics_success(
|
||||
# Assert
|
||||
assert isinstance(result, DataFrame)
|
||||
model_metrics_activity.debug.assert_called()
|
||||
model_metrics_activity.observe_lag.assert_called()
|
||||
model_metrics_activity.emit_metric.assert_called()
|
||||
model_metrics_activity.observe_lag_sync.assert_called()
|
||||
model_metrics_activity.emit_metric_sync.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_univariate_error(
|
||||
def test_get_drift_metrics_univariate_error(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
@@ -722,7 +741,7 @@ async def test_get_drift_metrics_univariate_error(
|
||||
|
||||
# Act & Assert
|
||||
try:
|
||||
await model_metrics_activity.get_drift_metrics(
|
||||
model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
@@ -736,19 +755,18 @@ async def test_get_drift_metrics_univariate_error(
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error detecting univariate drift: Univariate drift error', metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.emit_metric.assert_called_with(
|
||||
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_multivariate_error(
|
||||
def test_get_drift_metrics_multivariate_error(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
mock_time.return_value = 1000.0
|
||||
@@ -769,7 +787,7 @@ async def test_get_drift_metrics_multivariate_error(
|
||||
).columns
|
||||
|
||||
try:
|
||||
await model_metrics_activity.get_drift_metrics(
|
||||
model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
@@ -783,19 +801,18 @@ async def test_get_drift_metrics_multivariate_error(
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error detecting multivariate drift: Multivariate drift error', metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.emit_metric.assert_called_with(
|
||||
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_dataframe_error(
|
||||
def test_get_drift_metrics_dataframe_error(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
mock_time.return_value = 1000.0
|
||||
@@ -817,7 +834,7 @@ async def test_get_drift_metrics_dataframe_error(
|
||||
).columns
|
||||
|
||||
try:
|
||||
await model_metrics_activity.get_drift_metrics(
|
||||
model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
@@ -831,15 +848,14 @@ async def test_get_drift_metrics_dataframe_error(
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error getting drift metrics: Dataframe error', metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.emit_metric.assert_called_with(
|
||||
model_metrics_activity.emit_metric_sync.assert_called_with(
|
||||
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -858,7 +874,7 @@ async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activi
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 4
|
||||
@@ -877,8 +893,7 @@ async def test_calculate_simple_metrics_success_all_metrics(model_metrics_activi
|
||||
model_metrics_activity.debug.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -897,7 +912,7 @@ async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 1
|
||||
@@ -911,8 +926,7 @@ async def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -931,7 +945,7 @@ async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity)
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 1
|
||||
@@ -945,8 +959,7 @@ async def test_calculate_simple_metrics_success_mse_only(model_metrics_activity)
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -965,7 +978,7 @@ async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity)
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 1
|
||||
@@ -979,8 +992,7 @@ async def test_calculate_simple_metrics_success_mae_only(model_metrics_activity)
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -999,7 +1011,7 @@ async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 1
|
||||
@@ -1013,8 +1025,7 @@ async def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
||||
# Arrange
|
||||
# All target values are the same, so ss_tot will be 0
|
||||
target_data = DataFrame(
|
||||
@@ -1034,7 +1045,7 @@ async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 1
|
||||
@@ -1049,8 +1060,7 @@ async def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity):
|
||||
# Arrange
|
||||
target_data = DataFrame(
|
||||
{
|
||||
@@ -1069,7 +1079,7 @@ async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_me
|
||||
}
|
||||
|
||||
# Act
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
# Assert
|
||||
assert len(result['metric']) == 2
|
||||
@@ -1084,8 +1094,7 @@ async def test_calculate_simple_metrics_success_multiple_metrics_subset(model_me
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity):
|
||||
def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity):
|
||||
target_data = DataFrame(
|
||||
{
|
||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||
@@ -1102,7 +1111,7 @@ async def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_act
|
||||
'interval_minutes': 5,
|
||||
}
|
||||
|
||||
result = DataFrame(await model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||
|
||||
assert len(result['metric']) == 1
|
||||
assert result['metric'].values[0] == 'rmse'
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
from unittest.mock import ANY, MagicMock, call, patch
|
||||
|
||||
import pytest_asyncio
|
||||
import pytest
|
||||
from pandas import DataFrame
|
||||
from pytest import mark
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
@@ -23,39 +23,32 @@ def test__init__():
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
|
||||
assert opc.opc_servers == servers
|
||||
assert opc.opc_repository == {}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.opc.OpcRepository')
|
||||
@patch('laborious.activities.opc.OPC.send_notification_async')
|
||||
async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
@patch('laborious.activities.opc.OPC.send_notification')
|
||||
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, {}))
|
||||
)
|
||||
server2 = MagicMock(
|
||||
connect=AsyncMock(return_value=(True, {})), write_data=AsyncMock(return_value=(True, {}))
|
||||
)
|
||||
server3 = MagicMock(
|
||||
connect=AsyncMock(
|
||||
return_value=(
|
||||
False,
|
||||
{
|
||||
'notification_id': 'OPC_CONNECTION_ERROR_server3',
|
||||
'message': 'Failed to connect to OPC server: Test error',
|
||||
'block': 'opc_repository',
|
||||
'level': NotificationLevel.ERROR,
|
||||
'attachment_content': 'Test error',
|
||||
},
|
||||
)
|
||||
),
|
||||
write_data=AsyncMock(return_value=(True, {})),
|
||||
mock_metrics_controller = MagicMock()
|
||||
server1 = MagicMock()
|
||||
server1.connect.return_value = (True, {})
|
||||
server2 = MagicMock()
|
||||
server2.connect.return_value = (True, {})
|
||||
server3 = MagicMock()
|
||||
server3.connect.return_value = (
|
||||
False,
|
||||
{
|
||||
'notification_id': 'OPC_CONNECTION_ERROR_server3',
|
||||
'message': 'Failed to connect to OPC server: Test error',
|
||||
'block': 'opc_repository',
|
||||
'level': NotificationLevel.ERROR,
|
||||
'attachment_content': 'Test error',
|
||||
},
|
||||
)
|
||||
mock_opc_repository.side_effect = [server1, server2, server3]
|
||||
mock_notification_handler = MagicMock()
|
||||
@@ -97,7 +90,7 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
await opc.init_opc()
|
||||
opc.init_opc()
|
||||
|
||||
assert opc.opc_servers == servers
|
||||
assert opc.logger == mock_logger
|
||||
@@ -112,13 +105,12 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_name='server1',
|
||||
url='http://localhost:8080',
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
server_uri='opc.tcp://localhost:4840',
|
||||
cert_path='',
|
||||
private_key_path='',
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
),
|
||||
]
|
||||
)
|
||||
@@ -129,13 +121,12 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
server_name='server2',
|
||||
url='http://localhost:8080',
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
server_uri='opc.tcp://localhost:4840',
|
||||
cert_path='',
|
||||
private_key_path='',
|
||||
server_cert_path='',
|
||||
notification_handler=mock_notification_handler,
|
||||
reconnection_interval=60,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
]
|
||||
)
|
||||
@@ -162,51 +153,49 @@ async def test_init_opc(mock_send_notification, mock_opc_repository):
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
@patch('laborious.activities.opc.OpcRepository')
|
||||
async def opc(mock_opc_repository):
|
||||
servers = {
|
||||
'server1': {
|
||||
'id': 'server1',
|
||||
'server_name': 'server1',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
'cert_path': '',
|
||||
'private_key_path': '',
|
||||
'server_cert_path': '',
|
||||
'reconnection_interval': 60,
|
||||
@pytest.fixture
|
||||
def opc():
|
||||
with patch('laborious.activities.opc.OpcRepository') as mock_opc_repository:
|
||||
mock_opc_repository.return_value.write_data = MagicMock(return_value=(True, {}))
|
||||
mock_opc_repository.return_value.connect = MagicMock(return_value=(True, {}))
|
||||
servers = {
|
||||
'server1': {
|
||||
'id': 'server1',
|
||||
'server_name': 'server1',
|
||||
'url': 'http://localhost:8080',
|
||||
'server_uri': 'opc.tcp://localhost:4840',
|
||||
'cert_path': '',
|
||||
'private_key_path': '',
|
||||
'server_cert_path': '',
|
||||
'reconnection_interval': 60,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
await opc.init_opc()
|
||||
opc.send_notification = MagicMock()
|
||||
opc.send_notification_async = AsyncMock()
|
||||
opc.emit_metric = AsyncMock()
|
||||
return opc
|
||||
opc_instance = OPC(
|
||||
opc_servers=servers,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
opc_instance.init_opc()
|
||||
opc_instance.send_notification = MagicMock()
|
||||
opc_instance.emit_metric_sync = MagicMock()
|
||||
yield opc_instance
|
||||
|
||||
|
||||
WRITE_DATA_CASES = [
|
||||
('tag1', 'int', 50),
|
||||
('tag2', 'float', 50.5),
|
||||
('tag3', 'bool', True),
|
||||
('tag4', 'string', 'test'),
|
||||
('tag4', 'str', 'test'),
|
||||
]
|
||||
|
||||
|
||||
@mark.parametrize('tag,data_type,data', WRITE_DATA_CASES)
|
||||
@mark.asyncio
|
||||
async def test_write_data_success(opc, tag, data_type, data):
|
||||
def test_write_data_success(opc, tag, data_type, data):
|
||||
opc.opc_repository['server1'].write_data.return_value = (True, {'response_time': 0.1})
|
||||
|
||||
result = await opc.write_data(
|
||||
result = opc.write_data(
|
||||
server_id='server1',
|
||||
tag=tag,
|
||||
data=data,
|
||||
@@ -220,8 +209,7 @@ async def test_write_data_success(opc, tag, data_type, data):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_data_failed(opc):
|
||||
def test_write_data_failed(opc):
|
||||
opc.opc_repository['server1'].write_data.return_value = (
|
||||
False,
|
||||
{
|
||||
@@ -233,7 +221,7 @@ async def test_write_data_failed(opc):
|
||||
},
|
||||
)
|
||||
|
||||
result = await opc.write_data(
|
||||
result = opc.write_data(
|
||||
server_id='server1',
|
||||
tag='tag1',
|
||||
data=50,
|
||||
@@ -243,7 +231,7 @@ async def test_write_data_failed(opc):
|
||||
)
|
||||
assert result is None
|
||||
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
opc.send_notification.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='OPC_WRITE_DATA_ERROR_server1',
|
||||
message='Failed to write data to OPC server: Test error',
|
||||
@@ -253,12 +241,11 @@ async def test_write_data_failed(opc):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_data_exception(opc):
|
||||
def test_write_data_exception(opc):
|
||||
opc.opc_repository['server1'].write_data.side_effect = Exception('Test error')
|
||||
|
||||
try:
|
||||
await opc.write_data(
|
||||
opc.write_data(
|
||||
server_id='server1',
|
||||
tag='tag1',
|
||||
data=50,
|
||||
@@ -268,7 +255,7 @@ async def test_write_data_exception(opc):
|
||||
)
|
||||
|
||||
except Exception:
|
||||
opc.send_notification_async.assert_called_once_with(
|
||||
opc.send_notification.assert_called_once_with(
|
||||
metadata=metadata,
|
||||
notification_id='WRITE_OPC_PREDICTION_ERROR',
|
||||
message='Error writing data to OPC server: Test error',
|
||||
@@ -281,9 +268,8 @@ async def test_write_data_exception(opc):
|
||||
raise AssertionError('Expected an exception to be raised')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_manage_output_tags_success(opc):
|
||||
opc.write_data = AsyncMock(return_value=0.1)
|
||||
def test_manage_output_tags_success(opc):
|
||||
opc.write_data = MagicMock(return_value=0.1)
|
||||
|
||||
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||
config = {
|
||||
@@ -291,7 +277,7 @@ async def test_manage_output_tags_success(opc):
|
||||
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||
}
|
||||
|
||||
output_data, opc_metrics = await opc.manage_output_tags(
|
||||
output_data, opc_metrics = opc.manage_output_tags(
|
||||
server_id='server1',
|
||||
config=config,
|
||||
data=data,
|
||||
@@ -322,16 +308,15 @@ async def test_manage_output_tags_success(opc):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@mark.parametrize('side_effect', [[0.1, None], [None, 0.2]])
|
||||
async def test_manage_output_tags_failed(opc, side_effect):
|
||||
opc.write_data = AsyncMock(side_effect=side_effect)
|
||||
def test_manage_output_tags_failed(opc, side_effect):
|
||||
opc.write_data = MagicMock(side_effect=side_effect)
|
||||
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||
config = {
|
||||
'prediction_tags': {'tag1': {'data_type': 'float'}},
|
||||
'confidence_tags': {'tag2': {'data_type': 'float'}},
|
||||
}
|
||||
output_data, opc_metrics = await opc.manage_output_tags(
|
||||
output_data, opc_metrics = opc.manage_output_tags(
|
||||
server_id='server1',
|
||||
config=config,
|
||||
data=data,
|
||||
@@ -361,14 +346,13 @@ async def test_manage_output_tags_failed(opc, side_effect):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_manage_output_tags_do_nothing(opc):
|
||||
opc.write_data = AsyncMock(return_value=0.1)
|
||||
def test_manage_output_tags_do_nothing(opc):
|
||||
opc.write_data = MagicMock(return_value=0.1)
|
||||
data = DataFrame({'prediction': [0.75], 'prediction_confidence': [0.95]})
|
||||
config = {
|
||||
'_invalid_key': {'tag1': {'data_type': 'float'}},
|
||||
}
|
||||
output_data, opc_metrics = await opc.manage_output_tags(
|
||||
output_data, opc_metrics = opc.manage_output_tags(
|
||||
server_id='server1',
|
||||
config=config,
|
||||
data=data,
|
||||
@@ -379,11 +363,9 @@ async def test_manage_output_tags_do_nothing(opc):
|
||||
opc.write_data.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.opc.DataFrame')
|
||||
async def test_write_opc_data_success(mock_dataframe, opc):
|
||||
# Arrange
|
||||
input_data = {
|
||||
def test_write_opc_data_success(mock_dataframe, opc):
|
||||
input_data: dict[str, object] = {
|
||||
**metadata,
|
||||
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||
'opc_output_config': {
|
||||
@@ -393,19 +375,19 @@ async def test_write_opc_data_success(mock_dataframe, opc):
|
||||
}
|
||||
},
|
||||
}
|
||||
opc_output_config = input_data['opc_output_config']
|
||||
assert isinstance(opc_output_config, dict)
|
||||
|
||||
# Act
|
||||
opc.manage_output_tags = AsyncMock(return_value=(True, {'tag1': 0.1, 'tag2': 0.2}))
|
||||
opc.manage_output_tags = MagicMock(return_value=(True, {'tag1': 0.1, 'tag2': 0.2}))
|
||||
|
||||
opc.process_confidence = MagicMock(return_value={'data': 'data'})
|
||||
output_data, opc_metrics = await opc.write_opc_data(input_data)
|
||||
output_data, opc_metrics = opc.write_opc_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert output_data == {'data': 'data'}
|
||||
assert opc_metrics == {'server1': {'tag1': 0.1, 'tag2': 0.2}}
|
||||
opc.manage_output_tags.assert_called_once_with(
|
||||
'server1',
|
||||
input_data['opc_output_config']['server1'],
|
||||
opc_output_config['server1'],
|
||||
mock_dataframe.return_value,
|
||||
metadata['metadata'],
|
||||
)
|
||||
@@ -416,9 +398,7 @@ async def test_write_opc_data_success(mock_dataframe, opc):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_opc_data_empty_config(opc):
|
||||
# Arrange
|
||||
def test_write_opc_data_empty_config(opc):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||
@@ -426,16 +406,13 @@ async def test_write_opc_data_empty_config(opc):
|
||||
'opc_output_config': {'server1': {'prediction_tags': {}, 'confidence_tags': {}}},
|
||||
}
|
||||
|
||||
# Act
|
||||
await opc.write_opc_data(input_data)
|
||||
opc.write_opc_data(input_data)
|
||||
|
||||
# Assert
|
||||
opc.opc_repository['server1'].write_data.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_write_opc_data_no_validate_server(opc):
|
||||
opc.validate_server = AsyncMock(return_value=False)
|
||||
def test_write_opc_data_no_validate_server(opc):
|
||||
opc.validate_server = MagicMock(return_value=False)
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {'prediction': [0.75], 'prediction_confidence': [0.95]},
|
||||
@@ -447,10 +424,8 @@ async def test_write_opc_data_no_validate_server(opc):
|
||||
},
|
||||
}
|
||||
|
||||
# Act
|
||||
await opc.write_opc_data(input_data)
|
||||
opc.write_opc_data(input_data)
|
||||
|
||||
# Assert
|
||||
opc.opc_repository['server1'].write_data.assert_not_called()
|
||||
|
||||
|
||||
@@ -462,21 +437,18 @@ async def test_write_opc_data_no_validate_server(opc):
|
||||
],
|
||||
)
|
||||
def test_process_confidence(opc, data, success, expected):
|
||||
# Act
|
||||
result = opc.process_confidence(data, success, metadata)
|
||||
|
||||
# Assert
|
||||
assert result['prediction_confidence'][0] == expected
|
||||
|
||||
|
||||
@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
|
||||
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_close(opc):
|
||||
opc.opc_repository['server1'].disconnect = AsyncMock(return_value=True)
|
||||
await opc.close()
|
||||
opc.opc_repository['server1'].disconnect.assert_called_once()
|
||||
def test_close(opc):
|
||||
disconnect_mock = MagicMock(return_value=None)
|
||||
opc.opc_repository['server1'].disconnect = disconnect_mock
|
||||
opc.close()
|
||||
disconnect_mock.assert_called_once()
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import datetime
|
||||
import os
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
from pytest import fixture, mark, raises
|
||||
from pytest import fixture, raises
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.temporal.activities.postgres import Postgres
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.temporal.activities.postgres_sync import Postgres
|
||||
|
||||
from laborious.activities.storage import Storage
|
||||
|
||||
@@ -17,6 +18,15 @@ def _passthrough_from_dict():
|
||||
yield
|
||||
|
||||
|
||||
@fixture(autouse=True)
|
||||
def _patch_monitoring_shutdown():
|
||||
"""
|
||||
Avoid running real async SientiaMonitoring.shutdown when Storage.close runs inside tests.
|
||||
"""
|
||||
with patch.object(SientiaMonitoring, 'shutdown') as mock_shutdown:
|
||||
yield mock_shutdown
|
||||
|
||||
|
||||
metadata = {
|
||||
'metadata': {
|
||||
'model_id': 'test_model_id',
|
||||
@@ -42,7 +52,7 @@ def storage(mock_minio_repository):
|
||||
minio_repository=mock_minio_repository.return_value,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
@@ -50,7 +60,7 @@ def storage(mock_minio_repository):
|
||||
def test___init___not_hasattr(mock_minio_repository):
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
metrics_controller = MagicMock()
|
||||
minio_repo = mock_minio_repository.return_value
|
||||
storage = Storage(
|
||||
host='localhost',
|
||||
@@ -77,7 +87,7 @@ def test___init___none_minio_repository(mock_minio_repository, storage):
|
||||
storage.minio_repository = None
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
metrics_controller = AsyncMock()
|
||||
metrics_controller = MagicMock()
|
||||
storage.__init__(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
@@ -111,54 +121,55 @@ def test___init___done_repository(mock_minio_repository, storage):
|
||||
minio_repository=mock_minio_repository.return_value,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
mock_minio_repository.assert_not_called()
|
||||
assert storage.minio_repository is not None
|
||||
|
||||
|
||||
def test_close(storage):
|
||||
def test_close(storage, _patch_monitoring_shutdown):
|
||||
storage.minio_repository = MagicMock()
|
||||
|
||||
storage.close()
|
||||
|
||||
assert storage.minio_repository is None
|
||||
_patch_monitoring_shutdown.assert_called_once_with(storage)
|
||||
|
||||
|
||||
def test___del__(storage):
|
||||
storage.close = MagicMock()
|
||||
def test_close_when_minio_repository_already_none(storage, _patch_monitoring_shutdown):
|
||||
"""Closing without an initialized MinIO repository skips MinIO teardown."""
|
||||
storage.minio_repository = None
|
||||
|
||||
storage.__del__()
|
||||
storage.close()
|
||||
|
||||
storage.close.assert_called_once()
|
||||
assert storage.minio_repository is None
|
||||
_patch_monitoring_shutdown.assert_called_once_with(storage)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_load_query_with_minio_offload_no_rows(storage):
|
||||
storage.load_custom_query = AsyncMock(return_value=None)
|
||||
def test_load_query_with_minio_offload_no_rows(storage):
|
||||
storage.load_custom_query = MagicMock(return_value=None)
|
||||
storage_result = {'success': False}
|
||||
with patch(
|
||||
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
return_value=storage_result,
|
||||
) as mock_from_dataframe:
|
||||
result = await storage.load_query_with_minio_offload(
|
||||
result = storage.load_query_with_minio_offload(
|
||||
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
|
||||
)
|
||||
assert result == storage_result
|
||||
mock_from_dataframe.assert_awaited_once()
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_load_query_with_minio_offload_inline(storage):
|
||||
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
|
||||
def test_load_query_with_minio_offload_inline(storage):
|
||||
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
|
||||
storage_result = {'success': True, 'data': {'a': [1]}, 'object_key': None}
|
||||
with patch(
|
||||
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
return_value=storage_result,
|
||||
) as mock_from_dataframe:
|
||||
result = await storage.load_query_with_minio_offload(
|
||||
result = storage.load_query_with_minio_offload(
|
||||
{
|
||||
**metadata,
|
||||
'query': 'SELECT 1',
|
||||
@@ -167,44 +178,42 @@ async def test_load_query_with_minio_offload_inline(storage):
|
||||
}
|
||||
)
|
||||
assert result == storage_result
|
||||
mock_from_dataframe.assert_awaited_once()
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_load_query_with_minio_offload_minio(storage):
|
||||
storage.load_custom_query = AsyncMock(return_value=[{'a': 1}])
|
||||
def test_load_query_with_minio_offload_minio(storage):
|
||||
storage.load_custom_query = MagicMock(return_value=[{'a': 1}])
|
||||
storage_result = {'success': True, 'data': None, 'object_key': 'object-key'}
|
||||
with patch(
|
||||
'laborious.activities.storage.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
new_callable=MagicMock,
|
||||
return_value=storage_result,
|
||||
) as mock_from_dataframe:
|
||||
result = await storage.load_query_with_minio_offload(
|
||||
result = storage.load_query_with_minio_offload(
|
||||
{**metadata, 'query': 'SELECT 1', 'model_name': 'm', 'key_prefix': 'predictions/s'}
|
||||
)
|
||||
|
||||
assert result == storage_result
|
||||
mock_from_dataframe.assert_awaited_once()
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch.dict(os.environ, {'SIENTIA_MINIO_RETENTION_HOURS': '1'})
|
||||
@patch('laborious.activities.storage.now')
|
||||
async def test_cleanup_minio_objects_expired(mock_now, storage):
|
||||
def test_cleanup_minio_objects_expired(mock_now, storage):
|
||||
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
|
||||
storage.minio_repository.list_objects = AsyncMock(
|
||||
storage.minio_repository.list_objects = MagicMock(
|
||||
return_value=[
|
||||
'sientia/streamlit-connectors/training_datasets/m/m-initial-2024-12-01_00-00-00.parquet',
|
||||
'sientia/streamlit-connectors/training_datasets/m/m-initial-2025-01-10_12-00-00.parquet',
|
||||
]
|
||||
)
|
||||
storage.minio_repository.delete_file = AsyncMock()
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.delete_file = MagicMock()
|
||||
storage.send_notification = MagicMock()
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||
|
||||
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
|
||||
assert result['deleted_count'] == 1
|
||||
assert result['failed_count'] == 0
|
||||
@@ -224,72 +233,65 @@ async def test_cleanup_minio_objects_expired(mock_now, storage):
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_load_query_with_minio_offload_minio_not_initialized(storage):
|
||||
def test_load_query_with_minio_offload_minio_not_initialized(storage):
|
||||
storage.minio_repository = None
|
||||
|
||||
with raises(ValueError, match='Minio repository not initialized'):
|
||||
await storage.load_query_with_minio_offload(
|
||||
{**metadata, 'query': 'SELECT 1', 'model_name': 'm'}
|
||||
)
|
||||
storage.load_query_with_minio_offload({**metadata, 'query': 'SELECT 1', 'model_name': 'm'})
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_export_payload_to_postgres(storage):
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=MagicMock())
|
||||
storage.export_data_to_postgres = AsyncMock(return_value={'success': True})
|
||||
def test_export_payload_to_postgres(storage):
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=MagicMock())
|
||||
storage.export_data_to_postgres = MagicMock(return_value={'success': True})
|
||||
|
||||
result = await storage.export_payload_to_postgres(
|
||||
result = storage.export_payload_to_postgres(
|
||||
{**metadata, 'data': payload, 'schema': 'public', 'table': 't'}
|
||||
)
|
||||
|
||||
payload.retrieve.assert_awaited_once_with(storage.minio_repository, metadata['metadata'])
|
||||
storage.export_data_to_postgres.assert_awaited_once()
|
||||
payload.retrieve.assert_called_once_with(storage.minio_repository, metadata['metadata'])
|
||||
storage.export_data_to_postgres.assert_called_once()
|
||||
assert result == {'success': True}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_cleanup_minio_objects_expired_minio_not_initialized(storage):
|
||||
def test_cleanup_minio_objects_expired_minio_not_initialized(storage):
|
||||
storage.minio_repository = None
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.cleanup_prefix.return_value = 'test'
|
||||
with raises(ValueError, match='Minio repository not initialized'):
|
||||
await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.storage.now')
|
||||
async def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage):
|
||||
def test_cleanup_minio_objects_expired_unparseable_key(mock_now, storage):
|
||||
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
|
||||
storage.minio_repository.list_objects = AsyncMock(
|
||||
storage.minio_repository.list_objects = MagicMock(
|
||||
return_value=['some/random/key-without-timestamp.parquet']
|
||||
)
|
||||
storage.minio_repository.delete_file = AsyncMock()
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.delete_file = MagicMock()
|
||||
storage.send_notification = MagicMock()
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.cleanup_prefix.return_value = 'test'
|
||||
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
|
||||
assert result['deleted_count'] == 0
|
||||
assert result['failed_count'] == 0
|
||||
storage.minio_repository.delete_file.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.storage.now')
|
||||
async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
|
||||
def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
|
||||
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
|
||||
old_key = 'training_datasets/m/m-initial-2024-12-01_00-00-00.parquet'
|
||||
storage.minio_repository.list_objects = AsyncMock(return_value=[old_key])
|
||||
storage.minio_repository.delete_file = AsyncMock(side_effect=Exception('delete error'))
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.list_objects = MagicMock(return_value=[old_key])
|
||||
storage.minio_repository.delete_file = MagicMock(side_effect=Exception('delete error'))
|
||||
storage.send_notification = MagicMock()
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
|
||||
assert result['deleted_count'] == 0
|
||||
assert result['failed_count'] == 1
|
||||
@@ -298,21 +300,20 @@ async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
|
||||
assert result['failed'][old_key]['message'] == 'delete error'
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.storage.now')
|
||||
async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage):
|
||||
def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storage):
|
||||
mock_now.return_value = datetime.datetime(2025, 1, 10, 12, 0, 0)
|
||||
storage.minio_repository.list_objects = AsyncMock(side_effect=Exception('list error'))
|
||||
storage.send_notification_async = AsyncMock()
|
||||
storage.minio_repository.list_objects = MagicMock(side_effect=Exception('list error'))
|
||||
storage.send_notification = MagicMock()
|
||||
storage.error = MagicMock()
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
|
||||
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
result = storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
|
||||
|
||||
assert result['deleted_count'] == 0
|
||||
assert result['failed_count'] == 0
|
||||
storage.send_notification_async.assert_called_once_with(
|
||||
storage.send_notification.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='ERROR_CLEANUP_MINIO_OBJECTS_EXPIRED',
|
||||
message='Error cleaning up MinIO objects: list error',
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pandas import DataFrame
|
||||
|
||||
from laborious.utils.models.minio_dataframe_payload import (
|
||||
@@ -54,17 +53,15 @@ def test_has_data_true_when_object_key_set():
|
||||
assert payload.has_data() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_inline_dict_as_dataframe():
|
||||
def test_retrieve_inline_dict_as_dataframe():
|
||||
payload = MinioDataFramePayload(last_timestamp='t', data={'a': [1, 2]})
|
||||
minio = AsyncMock()
|
||||
out = await payload.retrieve(minio, {'metadata': {}})
|
||||
minio = MagicMock()
|
||||
out = payload.retrieve(minio, {'metadata': {}})
|
||||
assert list(out.columns) == ['a']
|
||||
minio.download_file.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_downloads_parquet_when_offloaded():
|
||||
def test_retrieve_downloads_parquet_when_offloaded():
|
||||
source = DataFrame({'a': [1, 2]})
|
||||
buf = BytesIO()
|
||||
source.to_parquet(buf, engine='pyarrow', index=True)
|
||||
@@ -76,12 +73,12 @@ async def test_retrieve_downloads_parquet_when_offloaded():
|
||||
object_key='training_datasets/m/f.parquet',
|
||||
object_prefix='training_datasets/m',
|
||||
)
|
||||
minio = AsyncMock()
|
||||
minio.download_file = AsyncMock(return_value=file_bytes)
|
||||
minio = MagicMock()
|
||||
minio.download_file = MagicMock(return_value=file_bytes)
|
||||
|
||||
out = await payload.retrieve(minio, {'metadata': {}})
|
||||
out = payload.retrieve(minio, {'metadata': {}})
|
||||
|
||||
minio.download_file.assert_awaited_once_with(
|
||||
minio.download_file.assert_called_once_with(
|
||||
object_name='training_datasets/m/f.parquet',
|
||||
metadata={'metadata': {}},
|
||||
)
|
||||
@@ -113,21 +110,19 @@ def test_parse_object_timestamp_bad_datetime():
|
||||
assert MinioDataFramePayload.parse_object_timestamp(key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_empty_when_no_data():
|
||||
def test_retrieve_empty_when_no_data():
|
||||
payload = MinioDataFramePayload(last_timestamp='t', data=None, object_key=None)
|
||||
minio = AsyncMock()
|
||||
out = await payload.retrieve(minio, {})
|
||||
minio = MagicMock()
|
||||
out = payload.retrieve(minio, {})
|
||||
assert out.empty
|
||||
minio.download_file.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.now')
|
||||
async def test_from_dataframe_none(mock_now):
|
||||
def test_from_dataframe_none(mock_now):
|
||||
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
|
||||
minio = AsyncMock()
|
||||
result = await MinioDataFramePayload.from_dataframe(
|
||||
minio = MagicMock()
|
||||
result = MinioDataFramePayload.from_dataframe(
|
||||
dataframe=None,
|
||||
minio_repo=minio,
|
||||
model_name='m',
|
||||
@@ -139,15 +134,14 @@ async def test_from_dataframe_none(mock_now):
|
||||
assert result.object_key is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.now')
|
||||
async def test_from_dataframe_empty(mock_now):
|
||||
def test_from_dataframe_empty(mock_now):
|
||||
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
|
||||
minio = AsyncMock()
|
||||
minio = MagicMock()
|
||||
mock_df = MagicMock()
|
||||
mock_df.__bool__ = MagicMock(return_value=True)
|
||||
mock_df.empty = True
|
||||
result = await MinioDataFramePayload.from_dataframe(
|
||||
result = MinioDataFramePayload.from_dataframe(
|
||||
dataframe=mock_df,
|
||||
minio_repo=minio,
|
||||
model_name='m',
|
||||
@@ -174,12 +168,11 @@ def _mock_dataframe(data_dict, timestamp_values=None):
|
||||
return mock_df
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
|
||||
async def test_from_dataframe_inline():
|
||||
minio = AsyncMock()
|
||||
def test_from_dataframe_inline():
|
||||
minio = MagicMock()
|
||||
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
|
||||
result = await MinioDataFramePayload.from_dataframe(
|
||||
result = MinioDataFramePayload.from_dataframe(
|
||||
dataframe=df,
|
||||
minio_repo=minio,
|
||||
model_name='m',
|
||||
@@ -190,12 +183,11 @@ async def test_from_dataframe_inline():
|
||||
assert result.last_timestamp == '2024-01-01'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
|
||||
async def test_from_dataframe_inline_uses_provided_last_timestamp():
|
||||
minio = AsyncMock()
|
||||
def test_from_dataframe_inline_uses_provided_last_timestamp():
|
||||
minio = MagicMock()
|
||||
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
|
||||
result = await MinioDataFramePayload.from_dataframe(
|
||||
result = MinioDataFramePayload.from_dataframe(
|
||||
dataframe=df,
|
||||
minio_repo=minio,
|
||||
model_name='m',
|
||||
@@ -205,17 +197,16 @@ async def test_from_dataframe_inline_uses_provided_last_timestamp():
|
||||
assert result.last_timestamp == '2024-01-02'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.now')
|
||||
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0)
|
||||
async def test_from_dataframe_offloaded(mock_now):
|
||||
def test_from_dataframe_offloaded(mock_now):
|
||||
mock_now.return_value = datetime(2024, 1, 1, 0, 0, 0)
|
||||
minio = AsyncMock()
|
||||
minio.upload_file = AsyncMock(return_value={'minio_object_name': 'full/key.parquet'})
|
||||
minio = MagicMock()
|
||||
minio.upload_file = MagicMock(return_value={'minio_object_name': 'full/key.parquet'})
|
||||
minio.bucket = 'test-bucket'
|
||||
|
||||
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
|
||||
result = await MinioDataFramePayload.from_dataframe(
|
||||
result = MinioDataFramePayload.from_dataframe(
|
||||
dataframe=df,
|
||||
minio_repo=minio,
|
||||
model_name='m',
|
||||
@@ -226,7 +217,7 @@ async def test_from_dataframe_offloaded(mock_now):
|
||||
assert result.object_key == 'full/key.parquet'
|
||||
assert result.bucket == 'test-bucket'
|
||||
assert result.uri == 's3://test-bucket/full/key.parquet'
|
||||
minio.upload_file.assert_awaited_once()
|
||||
minio.upload_file.assert_called_once()
|
||||
|
||||
|
||||
def test_from_dict_inline():
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, Mock, patch
|
||||
from unittest.mock import ANY, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from asyncua.crypto.security_policies import SecurityPolicyBasic256
|
||||
from opcua.crypto import security_policies
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
from laborious.utils.repository.opc_repository import OpcRepository
|
||||
@@ -27,19 +27,18 @@ 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(),
|
||||
metrics_controller=MagicMock(),
|
||||
)
|
||||
repository.disconnection_interval = 0.1
|
||||
repository.send_notification = MagicMock()
|
||||
repository.send_notification_async = AsyncMock()
|
||||
repository.emit_metric = AsyncMock()
|
||||
repository.emit_metric_sync = MagicMock()
|
||||
return repository
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_client():
|
||||
with patch('laborious.utils.repository.opc_repository.Client') as mock:
|
||||
client_instance = AsyncMock()
|
||||
client_instance = MagicMock()
|
||||
mock.return_value = client_instance
|
||||
yield client_instance
|
||||
|
||||
@@ -68,58 +67,48 @@ def test_init(opc_repository):
|
||||
assert opc_repository.error_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_security(opc_repository, mock_client):
|
||||
def test_set_security(opc_repository, mock_client):
|
||||
opc_repository.client = mock_client
|
||||
await opc_repository.set_security()
|
||||
opc_repository.set_security()
|
||||
|
||||
mock_client.application_uri = 'urn:test:server'
|
||||
mock_client.set_security.assert_called_once_with(
|
||||
SecurityPolicyBasic256,
|
||||
certificate='/path/to/cert.pem',
|
||||
private_key='/path/to/key.pem',
|
||||
server_certificate='/path/to/server_cert.pem',
|
||||
security_policies.SecurityPolicyBasic256,
|
||||
'/path/to/cert.pem',
|
||||
'/path/to/key.pem',
|
||||
'/path/to/server_cert.pem',
|
||||
)
|
||||
assert mock_client.secure_channel_timeout == 10000000
|
||||
assert mock_client.session_timeout == 10000000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_security_missing_certificates(opc_repository):
|
||||
def test_set_security_missing_certificates(opc_repository):
|
||||
opc_repository.cert_path = None
|
||||
opc_repository.private_key_path = None
|
||||
|
||||
try:
|
||||
await opc_repository.set_security()
|
||||
except ValueError as e:
|
||||
assert str(e) == 'Certificate and private key paths must be provided for secure connection.'
|
||||
with pytest.raises(ValueError, match='Certificate and private key paths'):
|
||||
opc_repository.set_security()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_security_missing_client(opc_repository):
|
||||
def test_set_security_missing_client(opc_repository):
|
||||
opc_repository.client = None
|
||||
try:
|
||||
await opc_repository.set_security()
|
||||
except ValueError as e:
|
||||
assert str(e) == 'Client must be initialized before setting security'
|
||||
with pytest.raises(ValueError, match='Client must be initialized'):
|
||||
opc_repository.set_security()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_with_security(opc_repository, mock_client):
|
||||
opc_repository.try_connect = AsyncMock(return_value=(True, {}))
|
||||
result = await opc_repository.connect()
|
||||
def test_connect_with_security(opc_repository, mock_client):
|
||||
opc_repository.try_connect = MagicMock(return_value=(True, {}))
|
||||
result = opc_repository.connect()
|
||||
|
||||
opc_repository.try_connect.assert_called_once()
|
||||
assert opc_repository.client == mock_client
|
||||
assert result == (True, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_without_security(opc_repository, mock_client):
|
||||
def test_connect_without_security(opc_repository, mock_client):
|
||||
opc_repository.cert_path = None
|
||||
opc_repository.try_connect = AsyncMock(return_value=(True, {}))
|
||||
opc_repository.set_security = AsyncMock()
|
||||
result = await opc_repository.connect()
|
||||
opc_repository.try_connect = MagicMock(return_value=(True, {}))
|
||||
opc_repository.set_security = MagicMock()
|
||||
result = opc_repository.connect()
|
||||
|
||||
opc_repository.try_connect.assert_called_once()
|
||||
opc_repository.set_security.assert_not_called()
|
||||
@@ -127,25 +116,23 @@ async def test_connect_without_security(opc_repository, mock_client):
|
||||
assert result == (True, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_try_connect_success(opc_repository):
|
||||
def test_try_connect_success(opc_repository):
|
||||
opc_repository.last_reconnection_time = None
|
||||
opc_repository.client = AsyncMock()
|
||||
result = await opc_repository.try_connect()
|
||||
opc_repository.client = MagicMock()
|
||||
result = opc_repository.try_connect()
|
||||
|
||||
opc_repository.client.connect.assert_called_once()
|
||||
assert opc_repository.last_reconnection_time is not None
|
||||
assert result == (True, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_try_connect_fail(opc_repository):
|
||||
def test_try_connect_fail(opc_repository):
|
||||
opc_repository.last_reconnection_time = None
|
||||
opc_repository.disconnect = AsyncMock()
|
||||
opc_repository.disconnect = MagicMock()
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.client.connect.side_effect = Exception('Test error')
|
||||
|
||||
is_connected, error_data = await opc_repository.try_connect()
|
||||
is_connected, error_data = opc_repository.try_connect()
|
||||
|
||||
opc_repository.disconnect.assert_called_once()
|
||||
opc_repository.client.connect.assert_called_once()
|
||||
@@ -157,10 +144,9 @@ async def test_try_connect_fail(opc_repository):
|
||||
assert error_data['attachment_content'] is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_try_connect_no_client(opc_repository):
|
||||
def test_try_connect_no_client(opc_repository):
|
||||
opc_repository.client = None
|
||||
result = await opc_repository.try_connect()
|
||||
result = opc_repository.try_connect()
|
||||
assert result == (
|
||||
False,
|
||||
{
|
||||
@@ -172,21 +158,24 @@ async def test_try_connect_no_client(opc_repository):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnection_fallback_success(opc_repository, mock_client):
|
||||
def test_session_alive_returns_false_when_client_none(opc_repository):
|
||||
opc_repository.client = None
|
||||
assert opc_repository._session_alive() is False
|
||||
|
||||
|
||||
def test_disconnection_fallback_success(opc_repository, mock_client):
|
||||
opc_repository.client = mock_client
|
||||
mock_client.disconnect.return_value = True
|
||||
result = await opc_repository.disconnection_fallback()
|
||||
mock_client.disconnect.return_value = None
|
||||
result = opc_repository.disconnection_fallback()
|
||||
|
||||
mock_client.disconnect.assert_called_once()
|
||||
assert result == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnection_fallback_fail(opc_repository, mock_client):
|
||||
def test_disconnection_fallback_fail(opc_repository, mock_client):
|
||||
opc_repository.client = mock_client
|
||||
mock_client.disconnect.side_effect = Exception('Test error')
|
||||
result = await opc_repository.disconnection_fallback()
|
||||
result = opc_repository.disconnection_fallback()
|
||||
assert result == [
|
||||
{'attempt': 1, 'error': 'Test error', 'traceback': ANY},
|
||||
{'attempt': 2, 'error': 'Test error', 'traceback': ANY},
|
||||
@@ -197,32 +186,29 @@ async def test_disconnection_fallback_fail(opc_repository, mock_client):
|
||||
assert mock_client.disconnect.call_count == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect(opc_repository, mock_client):
|
||||
def test_disconnect(opc_repository, mock_client):
|
||||
opc_repository.client = mock_client
|
||||
opc_repository.disconnection_fallback = AsyncMock(return_value=[])
|
||||
await opc_repository.disconnect()
|
||||
opc_repository.disconnection_fallback = MagicMock(return_value=[])
|
||||
opc_repository.disconnect()
|
||||
|
||||
opc_repository.disconnection_fallback.assert_called_once()
|
||||
assert opc_repository.client is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_no_client(opc_repository):
|
||||
def test_disconnect_no_client(opc_repository):
|
||||
opc_repository.client = None
|
||||
assert await opc_repository.disconnect() is None
|
||||
assert opc_repository.disconnect() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_error(opc_repository, mock_client):
|
||||
def test_disconnect_error(opc_repository, mock_client):
|
||||
opc_repository.client = mock_client
|
||||
opc_repository.disconnection_fallback = AsyncMock(
|
||||
opc_repository.disconnection_fallback = MagicMock(
|
||||
return_value=[{'attempt': 1, 'error': 'Test error', 'traceback': 'text'}]
|
||||
)
|
||||
await opc_repository.disconnect()
|
||||
opc_repository.disconnect()
|
||||
|
||||
opc_repository.disconnection_fallback.assert_called_once()
|
||||
opc_repository.send_notification_async.assert_called_once_with(
|
||||
opc_repository.send_notification.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.',
|
||||
@@ -235,63 +221,39 @@ async def test_disconnect_error(opc_repository, mock_client):
|
||||
assert opc_repository.client is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_connection_none_client(opc_repository):
|
||||
def test_validate_connection_none_client(opc_repository):
|
||||
opc_repository.client = None
|
||||
opc_repository.connect = AsyncMock(return_value=(True, {}))
|
||||
response = await opc_repository.validate_connection()
|
||||
opc_repository.connect = MagicMock(return_value=(True, {}))
|
||||
response = opc_repository.validate_connection()
|
||||
assert response == (True, {})
|
||||
opc_repository.connect.assert_called_once()
|
||||
|
||||
|
||||
# @pytest.mark.asyncio
|
||||
# async def test_validate_connection_error_count_disconnect_error(opc_repository):
|
||||
# opc_repository.error_count = 6
|
||||
# opc_repository.client = AsyncMock()
|
||||
# opc_repository.disconnect = AsyncMock(side_effect=Exception('Test error'))
|
||||
# opc_repository.connect = AsyncMock(return_value=(True, {}))
|
||||
def test_validate_connection_disconnect_raises(opc_repository):
|
||||
"""Outer except path when reconnect cleanup fails mid-validation."""
|
||||
|
||||
# response = await opc_repository.validate_connection()
|
||||
# assert response == opc_repository.connect.return_value
|
||||
# opc_repository.disconnect.assert_called_once()
|
||||
# opc_repository.connect.assert_called_once()
|
||||
# opc_repository.logger.custom_error.assert_has_calls(
|
||||
# [
|
||||
# call('Failed to disconnect from OPC server: Test error', ANY),
|
||||
# ]
|
||||
# )
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository._session_alive = MagicMock(return_value=False)
|
||||
opc_repository.last_reconnection_time = datetime(2020, 1, 1, 0, 0, 0)
|
||||
opc_repository.disconnect = MagicMock(side_effect=RuntimeError('disconnect failed'))
|
||||
|
||||
response = opc_repository.validate_connection()
|
||||
|
||||
assert response[0] is False
|
||||
assert response[1]['notification_id'] == f'OPC_CONNECTION_CHECK_ERROR_{opc_repository.id}'
|
||||
assert 'disconnect failed' in response[1]['message']
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_connection_error_validate_connection_error(opc_repository):
|
||||
opc_repository.client = MagicMock(uaclient=Exception('Test error'))
|
||||
opc_repository.error_count = 0
|
||||
|
||||
response = await opc_repository.validate_connection()
|
||||
|
||||
assert response == (
|
||||
False,
|
||||
{
|
||||
'notification_id': f'OPC_CONNECTION_CHECK_ERROR_{opc_repository.id}',
|
||||
'message': "Failed to validate connection to OPC server: 'Exception' object has no attribute 'protocol'",
|
||||
'block': 'opc_repository',
|
||||
'level': NotificationLevel.ERROR,
|
||||
'attachment_content': ANY,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.opc_repository.datetime')
|
||||
async def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, opc_repository):
|
||||
def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, opc_repository):
|
||||
_mock_datetime.now = MagicMock(return_value=datetime(2025, 1, 1, 0, 0, 0))
|
||||
opc_repository.error_count = 0
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.client.uaclient.protocol = None
|
||||
opc_repository.client.get_root_node.side_effect = RuntimeError('down')
|
||||
opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0)
|
||||
opc_repository.connect = MagicMock(return_value=(True, {}))
|
||||
|
||||
response = await opc_repository.validate_connection()
|
||||
response = opc_repository.validate_connection()
|
||||
opc_repository.connect.assert_not_called()
|
||||
assert response == (
|
||||
False,
|
||||
@@ -304,55 +266,51 @@ async def test_validate_connection_lost_not_time_to_reconnect(_mock_datetime, op
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.opc_repository.datetime')
|
||||
async def test_validate_connection_lost_time_to_reconnect(mock_datetime, opc_repository):
|
||||
def test_validate_connection_lost_time_to_reconnect(mock_datetime, opc_repository):
|
||||
mock_datetime.now = MagicMock(return_value=datetime(2025, 1, 1, 1, 0, 0))
|
||||
opc_repository.error_count = 0
|
||||
opc_repository.client = AsyncMock()
|
||||
opc_repository.client.uaclient.protocol = None
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.client.get_root_node.side_effect = RuntimeError('down')
|
||||
opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0)
|
||||
opc_repository.connect = AsyncMock(return_value=(True, {}))
|
||||
opc_repository.connect = MagicMock(return_value=(True, {}))
|
||||
|
||||
response = await opc_repository.validate_connection()
|
||||
response = opc_repository.validate_connection()
|
||||
opc_repository.connect.assert_called_once()
|
||||
assert response == opc_repository.connect.return_value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_connection_success(opc_repository):
|
||||
def test_validate_connection_success(opc_repository):
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.error_count = 0
|
||||
opc_repository.client.uaclient.protocol = MagicMock()
|
||||
opc_repository.client.uaclient.protocol.state = 'open'
|
||||
opc_repository.client.get_root_node.return_value = MagicMock()
|
||||
|
||||
output = await opc_repository.validate_connection()
|
||||
output = opc_repository.validate_connection()
|
||||
assert output == (True, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data_validate_connection_do_nothing(opc_repository):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
||||
opc_repository.client = AsyncMock(get_node=MagicMock())
|
||||
mock_node = AsyncMock()
|
||||
def test_write_data_validate_connection_do_nothing(opc_repository):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||
opc_repository.client = MagicMock()
|
||||
mock_node = MagicMock()
|
||||
opc_repository.client.get_node.return_value = mock_node
|
||||
|
||||
result = await opc_repository.write_data(
|
||||
result = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
opc_repository.validate_connection.assert_called_once()
|
||||
opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||
mock_node.set_value.assert_called_once()
|
||||
assert result == (True, {'response_time': ANY})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data_validate_connection_failed(opc_repository):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(False, {}))
|
||||
opc_repository.client = AsyncMock()
|
||||
def test_write_data_validate_connection_failed(opc_repository):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(False, {}))
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.error_count = 0
|
||||
|
||||
result = await opc_repository.write_data(
|
||||
result = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
@@ -361,14 +319,13 @@ async def test_write_data_validate_connection_failed(opc_repository):
|
||||
assert result == (False, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data_get_node_failed(opc_repository):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
||||
opc_repository.client = AsyncMock()
|
||||
def test_write_data_get_node_failed(opc_repository):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||
opc_repository.client = MagicMock()
|
||||
opc_repository.error_count = 0
|
||||
opc_repository.client.get_node = MagicMock(side_effect=Exception('Test error'))
|
||||
|
||||
is_success, error_data = await opc_repository.write_data(
|
||||
is_success, error_data = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
@@ -376,23 +333,15 @@ async def test_write_data_get_node_failed(opc_repository):
|
||||
opc_repository.client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||
assert is_success is False
|
||||
assert error_data['notification_id'] == f'OPC_WRITE_GET_NODE_ERROR_{opc_repository.id}'
|
||||
assert (
|
||||
error_data['message']
|
||||
== "Failed to get node from OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||
)
|
||||
assert error_data['block'] == 'opc_repository'
|
||||
assert error_data['level'] == NotificationLevel.ERROR
|
||||
assert error_data['attachment_content'] is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data_invalid_data_type(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
||||
def test_write_data_invalid_data_type(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||
opc_repository.client = mock_client
|
||||
mock_node = AsyncMock()
|
||||
mock_node = MagicMock()
|
||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||
|
||||
is_success, error_data = await opc_repository.write_data(
|
||||
is_success, error_data = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'invalid_type', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
@@ -401,53 +350,37 @@ async def test_write_data_invalid_data_type(opc_repository, mock_client):
|
||||
|
||||
assert is_success is False
|
||||
assert error_data['notification_id'] == f'OPC_WRITE_DATA_TYPE_ERROR_{opc_repository.id}'
|
||||
assert (
|
||||
error_data['message']
|
||||
== "Unsupported data type: invalid_type | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||
)
|
||||
assert error_data['block'] == 'opc_repository'
|
||||
assert error_data['level'] == NotificationLevel.ERROR
|
||||
assert error_data.get('attachment_content') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
||||
def test_write_data(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||
opc_repository.client = mock_client
|
||||
mock_node = AsyncMock()
|
||||
mock_node = MagicMock()
|
||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||
|
||||
result = await opc_repository.write_data(
|
||||
result = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||
mock_node.write_value.assert_called_once()
|
||||
mock_node.set_value.assert_called_once()
|
||||
assert result == (True, {'response_time': ANY})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_data_write_value_failed(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = AsyncMock(return_value=(True, {}))
|
||||
def test_write_data_write_value_failed(opc_repository, mock_client):
|
||||
opc_repository.validate_connection = MagicMock(return_value=(True, {}))
|
||||
opc_repository.client = mock_client
|
||||
mock_node = AsyncMock()
|
||||
mock_node = MagicMock()
|
||||
opc_repository.error_count = 0
|
||||
mock_client.get_node = MagicMock(return_value=mock_node)
|
||||
mock_node.write_value.side_effect = Exception('Test error')
|
||||
mock_node.set_value.side_effect = Exception('Test error')
|
||||
|
||||
is_success, error_data = await opc_repository.write_data(
|
||||
is_success, error_data = opc_repository.write_data(
|
||||
'ns=2;s=TestNode', 42.0, 'float', opc_repository.logger, metadata['metadata']
|
||||
)
|
||||
|
||||
opc_repository.validate_connection.assert_called_once()
|
||||
mock_client.get_node.assert_called_once_with('ns=2;s=TestNode')
|
||||
mock_node.write_value.assert_called_once()
|
||||
mock_node.set_value.assert_called_once()
|
||||
assert is_success is False
|
||||
assert error_data['notification_id'] == f'OPC_WRITE_DATA_ERROR_{opc_repository.id}'
|
||||
assert (
|
||||
error_data['message']
|
||||
== "Failed to write data to OPC server: Test error | metadata: {'model_id': 'test_model', 'model_name': 'test_model', 'workflow_name': 'test_workflow', 'schema_name': 'test_schedule'}"
|
||||
)
|
||||
assert error_data['block'] == 'opc_repository'
|
||||
assert error_data['level'] == NotificationLevel.ERROR
|
||||
assert error_data['attachment_content'] is not None
|
||||
|
||||
@@ -7,8 +7,8 @@ from laborious.worker import worker
|
||||
|
||||
def _build_fake_activities():
|
||||
inst = MagicMock()
|
||||
inst.init_opc = AsyncMock()
|
||||
inst.shutdown = AsyncMock()
|
||||
inst.init_opc = MagicMock()
|
||||
inst.shutdown = MagicMock()
|
||||
inst.load_query_with_minio_offload = MagicMock()
|
||||
inst.retrain_model = MagicMock()
|
||||
inst.update_production_model = MagicMock()
|
||||
@@ -176,7 +176,7 @@ async def test_main_success_exit_zero(monkeypatch):
|
||||
|
||||
notif = m_notif_cls.return_value
|
||||
notif.shutdown.assert_called_once()
|
||||
fake_activities.shutdown.assert_awaited_once()
|
||||
fake_activities.shutdown.assert_called_once()
|
||||
assert m_prepare.call_count == 4
|
||||
prepare_calls = m_prepare.call_args_list
|
||||
assert prepare_calls[0].kwargs['runtime'] == 'single'
|
||||
|
||||
Reference in New Issue
Block a user