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:
vitor-aignosi
2026-05-07 17:02:25 -03:00
parent aaf647efdf
commit e6018af23f
51 changed files with 4408 additions and 2660 deletions

View File

@@ -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')

View File

@@ -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()

View File

@@ -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,

View File

@@ -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)

View File

@@ -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

View File

@@ -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'

View File

@@ -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()

View File

@@ -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',

View File

@@ -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():

View File

@@ -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

View File

@@ -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'