SIENTIAPDE-1712

Update dependencies and refactor input filter handling for consistency

- Updated sientia-dataops-library dependency version from 1.10.3 to 1.10.4 in requirements.txt.
- Refactored input filter handling in the Gates class to read policy and config keys in a case-insensitive manner.
- Updated test cases to ensure consistency in filter key naming conventions across various scenarios.
This commit is contained in:
vitor-aignosi
2026-03-23 14:45:21 -03:00
parent f22cc49b93
commit 503d9aa485
24 changed files with 331 additions and 234 deletions

View File

@@ -90,7 +90,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
)
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -103,7 +103,7 @@ async def test_input_gate_filter_exception(mock_input_filter_functions, gates_ac
gates_activity.send_notification_async.assert_called_once_with(
metadata=metadata['metadata'],
notification_id='INTPUT_GATE_ERROR__EMPTY_DATA',
message="Error in filter EMPTY_DATA:{'policy': 'STOP', 'config': {}}: \n Test error",
message="Error in filter EMPTY_DATA:{'POLICY': 'STOP', 'CONFIG': {}}: \n Test error",
block='input_gate',
level=NotificationLevel.ERROR,
attachment_content=ANY,
@@ -133,7 +133,7 @@ async def test_input_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -147,11 +147,45 @@ async def test_input_gate_with_filter(gates_activity):
@mark.asyncio
async def test_input_gate_with_filter_not_caught(gates_activity):
async def test_input_gate_with_filter_lowercase_keys(gates_activity):
# Arrange
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'policy': 'STOP', 'config': {}}},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
# Act
result = await 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):
# Arrange
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'Policy': 'STOP', 'Config': {}}},
'data': _minio_payload(DataFrame({'value': []})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
# Act
result = await 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):
# Arrange
input_data = {
**metadata,
'filters': {'EMPTY_DATA': {'POLICY': 'STOP', 'CONFIG': {}}},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
@@ -248,7 +282,7 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}},
'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': _minio_payload(
{'content': {'message': 'API error occurred', 'traceback': 'error trace'}},
status={'success': False, 'message': 'API error occurred'},
@@ -266,12 +300,33 @@ async def test_mlflow_response_gate_with_filter(gates_activity):
gates_activity.send_notification_async.assert_called()
@mark.asyncio
async def test_mlflow_response_gate_with_filter_capitalized_keys(gates_activity):
# Arrange
input_data = {
**metadata,
'filters': {'API_ERROR': {'Policy': 'STOP'}},
'data': _minio_payload(
{'content': {'message': 'API error occurred', 'traceback': 'error trace'}},
status={'success': False, 'message': 'API error occurred'},
),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
}
# Act
result = await 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):
# Arrange
input_data = {
**metadata,
'filters': {'API_ERROR': {'policy': 'STOP'}},
'filters': {'API_ERROR': {'POLICY': 'STOP'}},
'data': _minio_payload(
{'content': {'message': 'success'}},
status={'success': True},
@@ -364,7 +419,7 @@ async def test_mlflow_content_gate_with_filter(gates_activity):
# Arrange
input_data = {
**metadata,
'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}},
'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}},
'data': _minio_payload(DataFrame({'value': [None, None, None]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
@@ -402,7 +457,7 @@ async def test_mlflow_content_gate_with_filter_not_caught(gates_activity):
async def test_mlflow_content_gate_filter_returns_false(gates_activity):
input_data = {
**metadata,
'filters': {'NAN_VALUES': {'policy': 'STOP', 'config': {}}},
'filters': {'NAN_VALUES': {'POLICY': 'STOP', 'CONFIG': {}}},
'data': _minio_payload(DataFrame({'value': [1, 2, 3]})),
'type': 'test',
'path_priority': ['STOP', 'CONTINUE', 'REPEAT'],
@@ -823,6 +878,7 @@ async def test_write_metrics(mock_metrics, gates_activity):
core_tags = {
'pod_id': gates_activity.pod_id,
'runtime': gates_activity.runtime,
'operation_type': 'predict',
'model_name': metadata['metadata']['model_name'],
'workflow_name': metadata['metadata']['workflow_name'],
}
@@ -930,6 +986,7 @@ async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_act
tags={
'pod_id': gates_activity.pod_id,
'runtime': gates_activity.runtime,
'operation_type': 'predict',
'model_name': metadata['metadata']['model_name'],
'workflow_name': metadata['metadata']['workflow_name'],
'opc_server_id': 'server1',

View File

@@ -180,6 +180,7 @@ async def test_request_transform_failure(mock_from_dataframe, mlflow):
operation='transform',
status=transform_response,
workflow_metadata=metadata['metadata'],
last_timestamp=payload.last_timestamp,
)
assert response_data == mock_from_dataframe.return_value
@@ -250,6 +251,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml
operation='predict',
status=predict_response,
workflow_metadata=metadata['metadata'],
last_timestamp=payload.last_timestamp,
)
assert response_data == mock_from_dataframe.return_value

View File

@@ -204,9 +204,7 @@ async def test_cleanup_minio_objects_expired(mock_now, storage):
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'data': data_mock}
)
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 1
assert result['failed_count'] == 0
@@ -291,9 +289,7 @@ async def test_cleanup_minio_objects_expired_delete_fails(mock_now, storage):
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'data': data_mock}
)
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 0
assert result['failed_count'] == 1
@@ -312,9 +308,7 @@ async def test_cleanup_minio_objects_expired_list_objects_error(mock_now, storag
data_mock = MagicMock()
data_mock.cleanup_prefix.return_value = 'training_datasets/m'
result = await storage.cleanup_minio_objects_expired(
{**metadata, 'data': data_mock}
)
result = await storage.cleanup_minio_objects_expired({**metadata, 'data': data_mock})
assert result['deleted_count'] == 0
assert result['failed_count'] == 0

View File

@@ -90,22 +90,20 @@ async def test_retrieve_downloads_parquet_when_offloaded():
def test_build_object_key():
key, prefix = _build_object_key('my-model', 'initial', '2024-01-01_00-00-00')
assert key == 'training_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet'
assert prefix == 'training_datasets/my-model'
assert key == 'prediction_datasets/my-model/my-model-initial-2024-01-01_00-00-00.parquet'
assert prefix == 'prediction_datasets/my-model'
def test_build_object_key_strips_slashes():
key, prefix = _build_object_key(' /my-model/ ', 'transform', '2024-06-15_10-30-45')
assert prefix == 'training_datasets/my-model'
assert key.startswith('training_datasets/my-model/')
assert prefix == 'prediction_datasets/my-model'
assert key.startswith('prediction_datasets/my-model/')
def test_estimate_size_bytes_fallback():
df = DataFrame({'a': [1, 2]})
original_to_dict = df.to_dict
df.to_dict = lambda *a, **kw: (_ for _ in ()).throw(RuntimeError('to_dict failed'))
size = MinioDataFramePayload.estimate_size_bytes(df)
df.to_dict = original_to_dict
with patch.object(df, 'to_dict', side_effect=RuntimeError('to_dict failed')):
size = MinioDataFramePayload.estimate_size_bytes(df)
assert isinstance(size, int)
assert size > 0

View File

@@ -6,15 +6,6 @@ from laborious.activities.activities import Activities
from laborious.workflows.sub_workflows.prediction_process import PredictionProcess
@fixture(autouse=True)
def _passthrough_from_dict():
with patch(
'laborious.workflows.sub_workflows.prediction_process.MinioDataFramePayload.from_dict',
side_effect=lambda x: x,
):
yield
@fixture
def prediction_process():
return PredictionProcess()
@@ -38,7 +29,9 @@ async def test_run(workflow_mock, prediction_process):
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,
@@ -128,7 +121,7 @@ async def test_run(workflow_mock, prediction_process):
{
**metadata,
'filters': input_data['mlflow_transform_filters'],
'data': 'transformed_data',
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'type': 'transform',
'path_priority': input_data['path_priority'],
},
@@ -143,7 +136,7 @@ async def test_run(workflow_mock, prediction_process):
Activities.request_predict,
{
**metadata,
'data': 'transformed_data',
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
},
@@ -174,8 +167,8 @@ async def test_run(workflow_mock, prediction_process):
{
'metadata': metadata,
'path_flag': 'continue',
'data': 'predicted_data',
'transformed_data': 'transformed_data',
'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'},
'transformed_data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'prediction_confidence': 0.95,
'timestamp': '2024-01-01',
'model_id': 1,
@@ -199,7 +192,9 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,
@@ -251,7 +246,9 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,
@@ -336,7 +333,9 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,
@@ -421,7 +420,7 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
Activities.mlflow_content_gate,
{
'filters': input_data['mlflow_transform_filters'],
'data': 'transformed_data',
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'type': 'transform',
'path_priority': input_data['path_priority'],
**metadata,
@@ -441,7 +440,9 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
# Arrange
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,
@@ -527,7 +528,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
Activities.mlflow_content_gate,
{
'filters': input_data['mlflow_transform_filters'],
'data': 'transformed_data',
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'type': 'transform',
'path_priority': input_data['path_priority'],
**metadata,
@@ -542,7 +543,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
call(
Activities.request_predict,
{
'data': 'transformed_data',
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
'model_name': input_data['model_name'],
'model_config': input_data['model_config'],
**metadata,
@@ -784,7 +785,9 @@ async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process):
data_payload = MagicMock()
data_payload.cleanup_prefix.return_value = 'training_datasets/test'
data_payload.__getitem__ = lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
data_payload.__getitem__ = (
lambda self, key: '2024-01-01' if key == 'last_timestamp' else MagicMock()
)
input_data = {
'metadata': metadata,
'data': data_payload,

View File

@@ -45,10 +45,9 @@ async def test_run(workflow_mock: AsyncMock, drift: Drift):
reference_data = {'data': 'test_reference_data'}
drift_data = {'drift': 'test_drift_data'}
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method.return_value = drift_data
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
# Act
await drift.run(input_data)
@@ -64,7 +63,7 @@ async def test_run(workflow_mock: AsyncMock, drift: Drift):
ORDER BY timestamp ASC
"""
workflow_mock.start_local_activity_method.assert_has_calls(
workflow_mock.start_activity_method.assert_has_calls(
[
call(
Activities.load_custom_query,
@@ -140,16 +139,15 @@ async def test_run_empty_target_data(workflow_mock: AsyncMock, drift: Drift):
target_data = None
reference_data = {'data': 'test_reference_data'}
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock()
# Act
await drift.run(input_data)
# Assert - Should not call calculate_drift or export
workflow_mock.execute_local_activity_method.assert_not_called()
workflow_mock.execute_activity_method.assert_not_called()
@@ -175,9 +173,8 @@ async def test_run_empty_drift_data(workflow_mock: AsyncMock, drift: Drift):
reference_data = {'data': 'test_reference_data'}
drift_data = None
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method.return_value = drift_data
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
workflow_mock.execute_activity_method = AsyncMock()
# Act
@@ -226,10 +223,9 @@ async def test_run_default_chunk_period(workflow_mock: AsyncMock, drift: Drift):
reference_data = {'data': 'test_reference_data'}
drift_data = {'drift': 'test_drift_data'}
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method.return_value = drift_data
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.start_activity_method.side_effect = [target_data, reference_data]
workflow_mock.execute_local_activity_method = AsyncMock(return_value=drift_data)
workflow_mock.execute_activity_method = AsyncMock(return_value=None)
# Act
await drift.run(input_data)

View File

@@ -1,4 +1,4 @@
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from unittest.mock import ANY, AsyncMock, call, patch
from pytest import fixture, mark
@@ -39,8 +39,15 @@ async def test_run(workflow_mock: AsyncMock, minimal_retrain: MinimalRetrain):
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = True
storage_result = {
'last_timestamp': '2024-01-01 00:00:00+0000',
'status': {'success': True},
'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}},
'bucket': None,
'object_key': None,
'object_prefix': None,
'uri': None,
}
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
@@ -163,8 +170,15 @@ async def test_run_storage_fail(workflow_mock: AsyncMock, minimal_retrain: Minim
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = False
storage_result = {
'last_timestamp': '2024-01-01 00:00:00+0000',
'status': {'success': True},
'data': {},
'bucket': None,
'object_key': None,
'object_prefix': None,
'uri': None,
}
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[
@@ -218,8 +232,15 @@ async def test_run_fail_retrain(workflow_mock: AsyncMock, minimal_retrain: Minim
},
}
storage_result = MagicMock()
storage_result.has_data.return_value = True
storage_result = {
'last_timestamp': '2024-01-01 00:00:00+0000',
'status': {'success': True},
'data': {'timestamp': {0: '2024-01-01 00:00:00+0000'}, 'value': {0: 1.0}},
'bucket': None,
'object_key': None,
'object_prefix': None,
'uri': None,
}
workflow_mock.execute_activity_method = AsyncMock(
side_effect=[

View File

@@ -12,12 +12,10 @@ def predictions_batch() -> PredictionsBatch:
metadata = {
'metadata': {
'model_id': 'test_model_id',
'model_name': 'test_model',
'workflow_name': 'predictions_batch',
'schedule_name': 'test_schedule',
},
'model_id': 'test_model_id',
'model_name': 'test_model',
'workflow_name': 'predictions_batch',
'schedule_name': 'test_schedule',
}
@@ -49,7 +47,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
call(
Activities.load_query_with_minio_offload,
{
**metadata,
'metadata': metadata,
'query': input_data['query'],
'datetime_columns': input_data.get('datetime_columns', []),
'model_name': input_data['model_name'],
@@ -60,19 +58,39 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
]
)
prediction_input = {
'metadata': metadata,
'metadata': {'metadata': metadata},
'data': activity_return,
'schema': input_data['schema'],
'table_name': input_data['table_name'],
'transform_table_name': input_data['transform_table_name'],
'model_id': input_data['model_id'],
'model_name': input_data['model_name'],
'input_filters': input_data.get('input_filters', {'EMPTY_DATA': {'POLICY': 'STOP'}}),
'input_filters': input_data.get(
'input_filters',
{
'EMPTY_DATA': {
'POLICY': 'STOP',
'CONFIG': {},
}
},
),
'mlflow_transform_filters': input_data.get(
'mlflow_transform_filters', {'API_ERROR': {'POLICY': 'STOP'}}
'mlflow_transform_filters',
{
'API_ERROR': {
'POLICY': 'STOP',
'CONFIG': {},
}
},
),
'mlflow_predict_filters': input_data.get(
'mlflow_predict_filters', {'API_ERROR': {'POLICY': 'STOP'}}
'mlflow_predict_filters',
{
'API_ERROR': {
'POLICY': 'STOP',
'CONFIG': {},
}
},
),
'model_config': input_data.get('model_config', {}),
'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']),

View File

@@ -42,9 +42,8 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
target_data = {'data': 'test_target_data'}
simple_metrics_data = {'metrics': 'test_simple_metrics_data'}
workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data]
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None])
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
# Act
await simple_metrics.run(input_data)
@@ -66,7 +65,7 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
p."timestamp" desc;
"""
workflow_mock.execute_local_activity_method.assert_has_calls(
workflow_mock.execute_activity_method.assert_has_calls(
[
call(
Activities.load_custom_query,
@@ -80,29 +79,30 @@ async def test_run(workflow_mock: AsyncMock, simple_metrics: SimpleMetrics):
start_to_close_timeout=ANY,
),
call(
Activities.calculate_simple_metrics,
Activities.export_data_to_postgres,
{
**metadata,
'model_id': input_data['model_id'],
'target_data': target_data,
'metrics': input_data['metrics'],
'interval_minutes': input_data['interval_minutes'],
'data': simple_metrics_data,
'schema': input_data['schema'],
'table_name': input_data['target_table_name'],
'timestamp_conversion': {
'column': 'timestamp',
'format': DATETIME_FORMAT_WITH_TZ,
},
},
retry_policy=ANY,
start_to_close_timeout=ANY,
),
]
)
# Assert - Check export_data_to_postgres call
workflow_mock.execute_activity_method.assert_called_once_with(
Activities.export_data_to_postgres,
workflow_mock.execute_local_activity_method.assert_called_once_with(
Activities.calculate_simple_metrics,
{
**metadata,
'data': simple_metrics_data,
'schema': input_data['schema'],
'table_name': input_data['target_table_name'],
'timestamp_conversion': {'column': 'timestamp', 'format': DATETIME_FORMAT_WITH_TZ},
'model_id': input_data['model_id'],
'target_data': target_data,
'metrics': input_data['metrics'],
'interval_minutes': input_data['interval_minutes'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
@@ -128,15 +128,13 @@ async def test_run_empty_target_data(workflow_mock: AsyncMock, simple_metrics: S
target_data = None
workflow_mock.execute_local_activity_method.return_value = target_data
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock(return_value=target_data)
# Act
await simple_metrics.run(input_data)
# Assert - Should not call calculate_simple_metrics or export
assert workflow_mock.execute_local_activity_method.call_count == 1
workflow_mock.execute_activity_method.assert_not_called()
assert workflow_mock.execute_activity_method.call_count == 1
@mark.asyncio
@@ -159,16 +157,15 @@ async def test_run_empty_simple_metrics(workflow_mock: AsyncMock, simple_metrics
target_data = {'data': 'test_target_data'}
simple_metrics_data = None
workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data]
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock(return_value=target_data)
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
# Act
await simple_metrics.run(input_data)
# Assert - Should call calculate_simple_metrics but not export
assert workflow_mock.execute_local_activity_method.call_count == 2
workflow_mock.execute_activity_method.assert_not_called()
workflow_mock.execute_activity_method.assert_called_once()
workflow_mock.execute_local_activity_method.assert_called_once()
@mark.asyncio
@@ -191,33 +188,28 @@ async def test_run_default_metrics(workflow_mock: AsyncMock, simple_metrics: Sim
target_data = {'data': 'test_target_data'}
simple_metrics_data = {'metrics': 'test_simple_metrics_data'}
workflow_mock.execute_local_activity_method.side_effect = [target_data, simple_metrics_data]
workflow_mock.execute_activity_method = AsyncMock()
workflow_mock.execute_activity_method = AsyncMock(side_effect=[target_data, None])
workflow_mock.execute_local_activity_method = AsyncMock(return_value=simple_metrics_data)
# Act
await simple_metrics.run(input_data)
# Assert - Check calculate_simple_metrics call with default metrics
workflow_mock.execute_local_activity_method.assert_has_calls(
[
call(
Activities.load_custom_query,
ANY,
retry_policy=ANY,
start_to_close_timeout=ANY,
),
call(
Activities.calculate_simple_metrics,
{
**metadata,
'model_id': input_data['model_id'],
'target_data': target_data,
'metrics': ['rmse', 'mse', 'mae', 'r2'], # Default value
'interval_minutes': input_data['interval_minutes'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
),
]
workflow_mock.execute_activity_method.assert_any_call(
Activities.load_custom_query,
ANY,
retry_policy=ANY,
start_to_close_timeout=ANY,
)
workflow_mock.execute_local_activity_method.assert_called_once_with(
Activities.calculate_simple_metrics,
{
**metadata,
'model_id': input_data['model_id'],
'target_data': target_data,
'metrics': ['rmse', 'mse', 'mae', 'r2'], # Default value
'interval_minutes': input_data['interval_minutes'],
},
retry_policy=ANY,
start_to_close_timeout=ANY,
)