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