781 lines
25 KiB
Python
781 lines
25 KiB
Python
from unittest.mock import ANY, AsyncMock, call, patch
|
|
|
|
from pytest import fixture, mark
|
|
|
|
from model_manager.activities.activities import Activities
|
|
from model_manager.workflows.sub_workflows.prediction_process import PredictionProcess
|
|
|
|
|
|
@fixture
|
|
def prediction_process():
|
|
return PredictionProcess()
|
|
|
|
|
|
metadata = {
|
|
'metadata': {
|
|
'model_id': 'test_model',
|
|
'model_name': 'test_model',
|
|
'workflow_name': 'test_workflow',
|
|
'schema_name': 'test_schedule',
|
|
},
|
|
}
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_run(workflow_mock, prediction_process):
|
|
prediction_process.path_flag_handler = AsyncMock(return_value=False)
|
|
# Arrange
|
|
input_data = {
|
|
'metadata': metadata,
|
|
'data': {'test': 'data'},
|
|
'schema': 'test_schema',
|
|
'table_name': 'test_table',
|
|
'model_id': 1,
|
|
'input_filters': {'test': 'filter'},
|
|
'mlflow_transform_filters': {'test': 'filter'},
|
|
'mlflow_predict_filters': {'test': 'filter'},
|
|
'model_name': 'test_model_name',
|
|
'model_config': {'retention': '30'},
|
|
'path_priority': ['continue', 'repeat', 'stop'],
|
|
'prediction_store_policy': 'lts:1',
|
|
}
|
|
|
|
# Mock the activity responses
|
|
workflow_mock.execute_local_activity_method.side_effect = [
|
|
'2024-01-01', # get_last_timestamp
|
|
('continue', 0.95, 'Input data with bad quality'), # input_gate
|
|
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
|
|
# mlflow_response_gate (transform)
|
|
('continue', 0.95, 'Error'),
|
|
# mlflow_content_gate (transform)
|
|
('continue', 0.95, 'Transformed data not passed the content filter'),
|
|
{'content': 'predicted_data', 'timestamp': '2024-01-01'}, # request_predict
|
|
# mlflow_response_gate (predict)
|
|
('continue', 0.95, 'Error'),
|
|
]
|
|
|
|
# Act
|
|
await prediction_process.run(input_data)
|
|
|
|
# Assert
|
|
assert workflow_mock.execute_local_activity_method.call_count == 7
|
|
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.get_last_timestamp,
|
|
{
|
|
**metadata,
|
|
'data': input_data['data'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.input_gate,
|
|
{
|
|
**metadata,
|
|
'filters': input_data['input_filters'],
|
|
'data': input_data['data'],
|
|
'path_priority': input_data['path_priority'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_transform,
|
|
{
|
|
**metadata,
|
|
'data': input_data['data'],
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
**metadata,
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_content_gate,
|
|
{
|
|
**metadata,
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': 'transformed_data',
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_predict,
|
|
{
|
|
**metadata,
|
|
'data': 'transformed_data',
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
**metadata,
|
|
'filters': input_data['mlflow_predict_filters'],
|
|
'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'},
|
|
'type': 'predict',
|
|
'path_priority': input_data['path_priority'],
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
|
|
workflow_mock.execute_child_workflow.assert_called_once_with(
|
|
'format_and_export_prediction',
|
|
{
|
|
'metadata': metadata,
|
|
'path_flag': 'continue',
|
|
'data': 'predicted_data',
|
|
'prediction_confidence': 0.95,
|
|
'timestamp': '2024-01-01',
|
|
'model_id': 1,
|
|
'model_name': 'test_model_name',
|
|
'model_config': input_data['model_config'],
|
|
'schema': input_data['schema'],
|
|
'table_name': input_data['table_name'],
|
|
'comment': 'Error',
|
|
'prediction_store_policy': input_data['prediction_store_policy'],
|
|
},
|
|
)
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
|
|
prediction_process.path_flag_handler = AsyncMock(return_value=True)
|
|
# Arrange
|
|
input_data = {
|
|
'metadata': metadata,
|
|
'data': {'test': 'data'},
|
|
'schema': 'test_schema',
|
|
'table_name': 'test_table',
|
|
'model_id': 1,
|
|
'input_filters': {'test': 'filter'},
|
|
'mlflow_transform_filters': {'test': 'filter'},
|
|
'mlflow_predict_filters': {'test': 'filter'},
|
|
'model_name': 'test_model_name',
|
|
'model_config': {'retention': '30'},
|
|
'path_priority': ['continue', 'repeat', 'stop'],
|
|
}
|
|
|
|
# Mock the activity responses
|
|
workflow_mock.execute_local_activity_method.side_effect = [
|
|
'2024-01-01', # get_last_timestamp
|
|
('stop', 0.95, 'Input data with bad quality'), # input_gate
|
|
]
|
|
|
|
# Act
|
|
await prediction_process.run(input_data)
|
|
|
|
# Assert
|
|
assert workflow_mock.execute_local_activity_method.call_count == 2
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.get_last_timestamp,
|
|
{
|
|
'data': input_data['data'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
),
|
|
call(
|
|
Activities.input_gate,
|
|
{
|
|
'filters': input_data['input_filters'],
|
|
'data': input_data['data'],
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
),
|
|
]
|
|
)
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_process):
|
|
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, True])
|
|
# Arrange
|
|
input_data = {
|
|
'metadata': metadata,
|
|
'data': {'test': 'data'},
|
|
'schema': 'test_schema',
|
|
'table_name': 'test_table',
|
|
'model_id': 1,
|
|
'input_filters': {'test': 'filter'},
|
|
'mlflow_transform_filters': {'test': 'filter'},
|
|
'mlflow_predict_filters': {'test': 'filter'},
|
|
'model_name': 'test_model_name',
|
|
'model_config': {'retention': '30'},
|
|
'path_priority': ['continue', 'repeat', 'stop'],
|
|
}
|
|
|
|
# Mock the activity responses
|
|
workflow_mock.execute_local_activity_method.side_effect = [
|
|
'2024-01-01', # get_last_timestamp
|
|
('repeat', 0.95, 'Input data with bad quality'), # input_gate
|
|
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
|
|
('continue', 0.95, 'Error'), # mlflow_response_gate (transform)
|
|
]
|
|
|
|
# Act
|
|
await prediction_process.run(input_data)
|
|
|
|
# Assert
|
|
assert workflow_mock.execute_local_activity_method.call_count == 4
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.get_last_timestamp,
|
|
{
|
|
'data': input_data['data'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.input_gate,
|
|
{
|
|
'filters': input_data['input_filters'],
|
|
'data': input_data['data'],
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_transform,
|
|
{
|
|
'data': input_data['data'],
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process):
|
|
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, True])
|
|
# Arrange
|
|
input_data = {
|
|
'metadata': metadata,
|
|
'data': {'test': 'data'},
|
|
'schema': 'test_schema',
|
|
'table_name': 'test_table',
|
|
'model_id': 1,
|
|
'input_filters': {'test': 'filter'},
|
|
'mlflow_transform_filters': {'test': 'filter'},
|
|
'mlflow_predict_filters': {'test': 'filter'},
|
|
'model_name': 'test_model_name',
|
|
'model_config': {'retention': '30'},
|
|
'path_priority': ['continue', 'repeat', 'stop'],
|
|
}
|
|
|
|
# Mock the activity responses
|
|
workflow_mock.execute_local_activity_method.side_effect = [
|
|
'2024-01-01', # get_last_timestamp
|
|
('continue', 0.95, 'Input data with bad quality'), # input_gate
|
|
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
|
|
# mlflow_response_gate (transform)
|
|
('continue', 0.95, 'Error'),
|
|
# mlflow_content_gate (transform)
|
|
('continue', 0.95, 'Transformed data not passed the content filter'),
|
|
]
|
|
|
|
# Act
|
|
await prediction_process.run(input_data)
|
|
|
|
# Assert
|
|
assert workflow_mock.execute_local_activity_method.call_count == 5
|
|
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.get_last_timestamp,
|
|
{
|
|
'data': input_data['data'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.input_gate,
|
|
{
|
|
'filters': input_data['input_filters'],
|
|
'data': input_data['data'],
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_transform,
|
|
{
|
|
'data': input_data['data'],
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_content_gate,
|
|
{
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': 'transformed_data',
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_process):
|
|
prediction_process.path_flag_handler = AsyncMock(side_effect=[False, False, False, True])
|
|
# Arrange
|
|
input_data = {
|
|
'metadata': metadata,
|
|
'data': {'test': 'data'},
|
|
'schema': 'test_schema',
|
|
'table_name': 'test_table',
|
|
'model_id': 1,
|
|
'input_filters': {'test': 'filter'},
|
|
'mlflow_transform_filters': {'test': 'filter'},
|
|
'mlflow_predict_filters': {'test': 'filter'},
|
|
'model_name': 'test_model_name',
|
|
'model_config': {'retention': '30'},
|
|
'path_priority': ['continue', 'repeat', 'stop'],
|
|
}
|
|
|
|
# Mock the activity responses
|
|
workflow_mock.execute_local_activity_method.side_effect = [
|
|
'2024-01-01', # get_last_timestamp
|
|
('continue', 0.95, 'Input data with bad quality'), # input_gate
|
|
{'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data
|
|
# mlflow_response_gate (transform)
|
|
('continue', 0.95, 'Error'),
|
|
# mlflow_content_gate (transform)
|
|
('continue', 0.95, 'Transformed data not passed the content filter'),
|
|
{'content': 'predicted_data', 'timestamp': '2024-01-01'}, # request_predict
|
|
('continue', 0.95, 'Error'), # mlflow_response_gate (predict)
|
|
]
|
|
|
|
# Act
|
|
await prediction_process.run(input_data)
|
|
|
|
# Assert
|
|
assert workflow_mock.execute_local_activity_method.call_count == 7
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.get_last_timestamp,
|
|
{
|
|
'data': input_data['data'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.input_gate,
|
|
{
|
|
'filters': input_data['input_filters'],
|
|
'data': input_data['data'],
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_transform,
|
|
{
|
|
'data': input_data['data'],
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'},
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_content_gate,
|
|
{
|
|
'filters': input_data['mlflow_transform_filters'],
|
|
'data': 'transformed_data',
|
|
'type': 'transform',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.request_predict,
|
|
{
|
|
'data': 'transformed_data',
|
|
'model_name': input_data['model_name'],
|
|
'model_config': input_data['model_config'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_local_activity_method.assert_has_calls(
|
|
[
|
|
call(
|
|
Activities.mlflow_response_gate,
|
|
{
|
|
'filters': input_data['mlflow_predict_filters'],
|
|
'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'},
|
|
'type': 'predict',
|
|
'path_priority': input_data['path_priority'],
|
|
**metadata,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
]
|
|
)
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_path_flag_handler_stop(workflow_mock, prediction_process):
|
|
# Arrange
|
|
data = {'test': 'data'}
|
|
path_flag = 'STOP'
|
|
confidence = 0.95
|
|
schema = 'test_schema'
|
|
table_name = 'test_table'
|
|
model = 'test_model'
|
|
last_timestamp = '2024-01-01'
|
|
model_name = 'test_model_name'
|
|
model_config = {'retention': '30'}
|
|
|
|
# Act
|
|
result = await prediction_process.path_flag_handler(
|
|
data,
|
|
path_flag,
|
|
{
|
|
'metadata': metadata,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'model_id': model,
|
|
'last_timestamp': last_timestamp,
|
|
'model_name': model_name,
|
|
'model_config': model_config,
|
|
},
|
|
confidence,
|
|
last_timestamp,
|
|
'',
|
|
)
|
|
|
|
# Assert
|
|
assert result is True
|
|
workflow_mock.execute_local_activity_method.assert_not_called()
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_path_flag_handler_repeat(workflow_mock, prediction_process):
|
|
# Arrange
|
|
data = {'test': 'data'}
|
|
path_flag = 'repeat'
|
|
confidence = 0.95
|
|
schema = 'test_schema'
|
|
table_name = 'test_table'
|
|
model = 'test_model'
|
|
last_timestamp = '2024-01-01'
|
|
model_name = 'test_model_name'
|
|
model_config = {'retention': '30'}
|
|
|
|
# Act
|
|
result = await prediction_process.path_flag_handler(
|
|
data,
|
|
path_flag,
|
|
{
|
|
'metadata': metadata,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'model_id': model,
|
|
'last_timestamp': last_timestamp,
|
|
'model_name': model_name,
|
|
'model_config': model_config,
|
|
},
|
|
confidence,
|
|
last_timestamp,
|
|
'',
|
|
)
|
|
|
|
# Assert
|
|
assert result is True
|
|
workflow_mock.execute_activity_method.assert_called_once_with(
|
|
Activities.repeat_last_prediction,
|
|
{
|
|
**metadata,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'model': model,
|
|
'last_timestamp': last_timestamp,
|
|
},
|
|
retry_policy=ANY,
|
|
start_to_close_timeout=ANY,
|
|
)
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_path_flag_handler_continue(workflow_mock, prediction_process):
|
|
# Arrange
|
|
data = {'test': 'data'}
|
|
path_flag = 'CONTINUE'
|
|
confidence = 0.95
|
|
schema = 'test_schema'
|
|
table_name = 'test_table'
|
|
model = 'test_model'
|
|
last_timestamp = '2024-01-01'
|
|
model_name = 'test_model_name'
|
|
model_config = {'retention': '30'}
|
|
prediction_store_policy = 'erl:1'
|
|
|
|
# Act
|
|
result = await prediction_process.path_flag_handler(
|
|
data,
|
|
path_flag,
|
|
{
|
|
'metadata': metadata,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'model_id': model,
|
|
'last_timestamp': last_timestamp,
|
|
'model_name': model_name,
|
|
'model_config': model_config,
|
|
'prediction_store_policy': prediction_store_policy,
|
|
},
|
|
confidence,
|
|
last_timestamp,
|
|
'Prediction Process',
|
|
)
|
|
|
|
# Assert
|
|
assert result is True
|
|
workflow_mock.execute_activity_method.assert_not_called()
|
|
workflow_mock.execute_child_workflow.assert_called_once_with(
|
|
'format_and_export_prediction',
|
|
{
|
|
'metadata': metadata,
|
|
'path_flag': path_flag,
|
|
'data': data,
|
|
'prediction_confidence': confidence,
|
|
'timestamp': last_timestamp,
|
|
'model_id': model,
|
|
'model_name': model_name,
|
|
'model_config': model_config,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'comment': 'Prediction Process',
|
|
'prediction_store_policy': prediction_store_policy,
|
|
},
|
|
)
|
|
|
|
|
|
@mark.asyncio
|
|
@patch('model_manager.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
|
|
async def test_path_flag_handler_unknown(workflow_mock, prediction_process):
|
|
# Arrange
|
|
data = {'test': 'data'}
|
|
path_flag = 'unknown'
|
|
confidence = 0.95
|
|
schema = 'test_schema'
|
|
table_name = 'test_table'
|
|
model = 'test_model'
|
|
last_timestamp = '2024-01-01'
|
|
model_name = 'test_model_name'
|
|
model_config = {'retention': '30'}
|
|
prediction_store_policy = 'erl:1'
|
|
# Act
|
|
result = await prediction_process.path_flag_handler(
|
|
data,
|
|
path_flag,
|
|
{
|
|
**metadata,
|
|
'schema': schema,
|
|
'table_name': table_name,
|
|
'model_id': model,
|
|
'last_timestamp': last_timestamp,
|
|
'model_name': model_name,
|
|
'model_config': model_config,
|
|
'prediction_store_policy': prediction_store_policy,
|
|
},
|
|
confidence,
|
|
last_timestamp,
|
|
'',
|
|
)
|
|
|
|
# Assert
|
|
assert result is False
|
|
workflow_mock.execute_activity_method.assert_not_called()
|
|
workflow_mock.execute_child_workflow.assert_not_called()
|