"""Unit tests for SFN utils.""" import json import unittest from unittest import mock from unittest.mock import patch from unittest.mock import call from unittest.mock import MagicMock from lambdacommon.aws import sfn class TestStateMachine(unittest.TestCase): """Test StateMachine class.""" def setUp(self): """Set up class test.""" self.sfn_arn = 'arn:aws:states:us-east-1:1234567890:stateMachine:TestStateMachine' self.run_arn = f'{self.sfn_arn}:test_run' # noqa E231 self.state_machine = sfn.StateMachine(self.sfn_arn) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_describe(self, mock_sfn_client): """Test mock StateMachine.describe method.""" mocked_response = { 'stateMachineArn': self.sfn_arn, 'name': 'TestStateMachine', 'status': 'ACTIVE' } mock_sfn_client.describe_state_machine.return_value = mocked_response response = self.state_machine.describe() assert response == mocked_response mock_sfn_client.describe_state_machine.assert_called_once_with(stateMachineArn=self.sfn_arn) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_start(self, mock_sfn_client): """Test mock StateMachine.start method.""" run_name = 'test_run' execution_input = {'key': 'value'} mocked_response = {'executionArn': self.run_arn} mock_sfn_client.start_execution.return_value = mocked_response response = self.state_machine.start(run_name, execution_input) assert response == mocked_response['executionArn'] mock_sfn_client.start_execution.assert_called_once_with( stateMachineArn=self.sfn_arn, name=run_name, input=execution_input ) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_find_execution_by_name_prefix(self, mock_sfn_client): """Test mock StateMachine.find_execution_by_name_prefix method.""" name_prefix = 'test' mocked_response = { 'executions': [ {'name': 'test_01', 'executionArn': self.run_arn} ] } mock_sfn_client.list_executions.return_value = mocked_response response_found = self.state_machine.find_execution_by_name_prefix(name_prefix) assert response_found == mocked_response['executions'][0]['executionArn'] response_not_found = self.state_machine.find_execution_by_name_prefix('notExists') assert response_not_found is None mock_sfn_client.list_executions.assert_has_calls([ call(stateMachineArn=self.sfn_arn), call(stateMachineArn=self.sfn_arn) ]) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_describe_execution(self, mock_sfn_client): """Test mock StateMachine.describe_execution method.""" mocked_response = { 'executionArn': self.run_arn, 'stateMachineArn': self.sfn_arn, 'name': 'test_run', 'status': 'RUNNING' } mock_sfn_client.describe_execution.return_value = mocked_response response = self.state_machine.describe_execution(self.run_arn) assert response == mocked_response mock_sfn_client.describe_execution.assert_called_once_with(executionArn=self.run_arn) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_wait_execution_end(self, mock_sfn_client): """Test mock StateMachine.wait_execution_end method.""" mocked_response = { 'executionArn': self.run_arn, 'stateMachineArn': self.sfn_arn, 'name': 'test_run', 'status': 'SUCCEEDED' } mock_sfn_client.describe_execution.side_effect = [ {'executionArn': self.run_arn, 'stateMachineArn': self.sfn_arn, 'name': 'test_run', 'status': 'RUNNING'}, {'executionArn': self.run_arn, 'stateMachineArn': self.sfn_arn, 'name': 'test_run', 'status': 'SUCCEEDED'} ] response = self.state_machine.wait_execution_end(self.run_arn) assert response == mocked_response assert mock_sfn_client.describe_execution.call_count >= 2 mock_sfn_client.describe_execution.assert_called_with(executionArn=self.run_arn) @patch('lambdacommon.aws.sfn.SFN_CLIENT') def test_get_all_task_history(self, mock_sfn_client): """Test mock StateMachine.get_all_task_history methid.""" mocked_response = { 'events': [ {'id': 1, 'type': 'TaskStarted'}, {'id': 2, 'type': 'TaskEnded'} ] } mock_sfn_client.get_execution_history.side_effect = [ {'events': [{'id': 1, 'type': 'TaskStarted'}], 'nextToken': 'token'}, {'events': [{'id': 2, 'type': 'TaskEnded'}]} ] response = self.state_machine.get_all_task_history(self.run_arn, False) assert response == mocked_response['events'] mock_sfn_client.get_execution_history.assert_any_call( executionArn=self.run_arn, maxResults=120, includeExecutionData=False ) mock_sfn_client.get_execution_history.assert_called_with( executionArn=self.run_arn, maxResults=120, includeExecutionData=False, nextToken='token' ) class TestExecutedState(unittest.TestCase): """Test ExecutedState class.""" def test_init_full_event(self): """Test class init with a full event.""" state = sfn.ExecutedState({ 'id': 1, 'previousEventId': 0, 'type': 'StateStarted', 'StateStartedEventDetails': { 'name': 'Event Name', 'input': 'someinput' } }) expected_instance_repr = str({ '_id': 1, '_previous_id': 0, '_state_type': 'StateStarted', '_state_name': 'Event Name', '_details': { 'name': 'Event Name', 'input': 'someinput' } }) assert repr(state) == expected_instance_repr def test_init_event_no_details(self): """Test class init with an incomplete event.""" state = sfn.ExecutedState({ 'id': 1, 'previousEventId': 0, 'type': 'StateStarted' }) expected_instance_repr = str({ '_id': 1, '_previous_id': 0, '_state_type': 'StateStarted', }) assert getattr(state, '_state_name', None) is None assert getattr(state, '_details', None) is None assert repr(state) == expected_instance_repr class TestStateMachineExecutionTest(unittest.TestCase): """Test TestStateMachineExecutionTest class.""" @mock.patch.object(sfn.StateMachine, 'find_execution_by_name_prefix') def setUp(self, find_execution_by_name_prefix): """Set up class test.""" find_execution_by_name_prefix.return_value = None self.sfn_arn = 'arn:aws:states:us-east-1:1234567890:stateMachine:TestStateMachine' self.execution_input = json.dumps({'input': 'data'}) self.run_name = 'test_name' self.test = sfn.StateMachineExecutionTest( self.sfn_arn, self.run_name, self.execution_input ) self.mock_state_machine = MagicMock(spec=self.test._state_machine) self.test._state_machine = self.mock_state_machine def test_init(self): """Test class init.""" assert self.test._test_name == self.run_name assert self.test._execution_input == self.execution_input assert self.test._execution_name == f'{self.run_name}-{sfn.UNIQUE_IDENTIFIER}' assert self.test._execution_arn is None assert self.test._was_started_parallel is False assert self.test.was_started_parallel is False assert self.test.execution_error is None assert self.test.final_status is None assert self.test.final_output is None assert self.test._execution_history == [] def wait_execution_end(self): """Test wait_execution_end.""" execution = 'my-execution' self.test._execution_arn = execution self.mock_state_machine.wait_execution_end.return_value = { 'status': 'SUCCEEDED', 'output': 'someOutput' } self.test.wait_execution_end() assert self.test.final_status == 'SUCCEEDED' assert self.test.execution_error is None assert self.test.final_output == 'someOutput' self.mock_state_machine.wait_execution_end.assert_called_once_with(execution) def test_run(self): """Test run.""" new_execution = 'my-execution' self.mock_state_machine.start.return_value = new_execution self.test.run() self.mock_state_machine.start.assert_called_once_with( self.test._execution_name, self.test._execution_input ) assert self.test._execution_arn == new_execution def test_was_successful(self): """Test was_successful.""" self.test._completed_execution = {'status': 'SUCCEEDED'} assert self.test.was_successful is True self.test._completed_execution = {'status': 'FAILED'} assert self.test.was_successful is False def test_was_failed(self): """Test was_failed.""" self.test._completed_execution = {'status': 'FAILED'} assert self.test.was_failed is True def test_get_error(self): """Test get_error .""" self.test._completed_execution = { 'error': 'Error', 'cause': 'Error Cause' } assert self.test.execution_error.error == 'Error' assert self.test.execution_error.cause == 'Error Cause' def test_get_output(self): """Test get_output.""" self.test._completed_execution = {'output': {'key': 'value'}} assert self.test.final_output == {'key': 'value'} def test_get_executed_tasks(self): """Test get_executed_tasks.""" task_name = 'TestTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskTypeEntered', 'MyTaskEventDetails': {'name': task_name} } expected_task = sfn.ExecutedState(task_info) self.mock_state_machine.get_all_task_history.return_value = [ {'id': 2, 'previousEventId': 1, 'type': 'TaskTypeScheduled'}, task_info ] result = self.test.refresh_execution_history() # Test get_all_task_history just gets called once when get_task is called more than once result2 = self.test.refresh_execution_history() assert len(result) == 2 assert result[1] == expected_task assert result == result2 def test_get_task_found(self): """Test get_task on a found task.""" task_name = 'TestTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskTypeEntered', 'MyTaskEventDetails': {'name': 'TestTask'} } expected_task = sfn.ExecutedState(task_info) self.mock_state_machine.get_all_task_history.return_value = [ {'id': 2, 'previousEventId': 1, 'type': 'TaskTypeScheduled'}, task_info ] result = self.test.get_task(task_name) assert result == expected_task def test_get_task_not_found(self): """Test get_task on a found task.""" task_name = 'OtherTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskTypeEntered', 'MyTaskEventDetails': {'name': 'TestTask'} } self.mock_state_machine.get_all_task_history.return_value = [ task_info ] result = self.test.get_task(task_name) assert result is None self.mock_state_machine.get_all_task_history.assert_called_once_with( self.test._execution_arn) def test_get_task_with_type_found(self): """Test get_task with a type on a found task.""" task_name = 'TestTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskTypeEntered', 'MyTaskEventDetails': {'name': task_name} } expected_task = sfn.ExecutedState(task_info) self.mock_state_machine.get_all_task_history.return_value = [ {'id': 2, 'previousEventId': 1, 'type': 'TaskTypeScheduled'}, task_info ] result = self.test.get_task(task_name, sfn.StateType.ENTERED) assert result == expected_task def test_get_task_with_type_not_found(self): """Test get_task on a found task.""" task_name = 'TestTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskTypeEntered', 'MyTaskEventDetails': {'name': task_name} } self.mock_state_machine.get_all_task_history.return_value = [ {'id': 2, 'previousEventId': 1, 'type': 'TaskTypeScheduled'}, task_info ] result = self.test.get_task(task_name, sfn.StateType.STARTED) assert result is None self.mock_state_machine.get_all_task_history.assert_called_once_with( self.test._execution_arn) def test_task_happened_preloaded_history(self): """Test task_happened on a found task, with history already retrieved previously.""" task_name = 'TestTask' task_info = { 'id': 1, 'previousEventId': 0, 'type': 'TaskType', 'MyTaskEventDetails': {'name': task_name} } self.test._execution_history = [sfn.ExecutedState(task_info)] print(self.test._execution_history) result = self.test.task_was_called(task_name) assert result is True self.mock_state_machine.get_all_task_history.assert_not_called() class TestStateMachineExecutionTestForParallelExecution(unittest.TestCase): """Test StateMachineExecutionTest for a parallel execution.""" @mock.patch.object(sfn.StateMachine, 'find_execution_by_name_prefix') def setUp(self, find_execution_by_name_prefix): """Set up class test.""" self.execution_arn = 'my-execution-arn' find_execution_by_name_prefix.return_value = self.execution_arn self.sfn_arn = 'arn:aws:states:us-east-1:1234567890:stateMachine:TestStateMachine' self.execution_input = json.dumps({'input': 'data'}) self.run_name = 'test_name' self.test = sfn.StateMachineExecutionTest( self.sfn_arn, self.run_name, self.execution_input ) self.mock_state_machine = MagicMock(spec=self.test._state_machine) self.test._state_machine = self.mock_state_machine def test_init(self): """Test class init.""" assert self.test._test_name == self.run_name assert self.test._execution_input == self.execution_input assert self.test._execution_name == f'{self.run_name}-{sfn.UNIQUE_IDENTIFIER}' assert self.test._execution_arn == self.execution_arn assert self.test._was_started_parallel is True assert self.test.was_started_parallel is True assert self.test.execution_error is None assert self.test.final_status is None assert self.test.final_output is None assert self.test._execution_history == []