"""Test run_extract_sales task.""" from datetime import datetime from datetime import timezone from unittest.mock import call from unittest.mock import MagicMock from unittest.mock import patch from freezegun import freeze_time import pytest from tasks.sales_ingest.run_extract_sales import _get_ecs_credentials from tasks.sales_ingest.run_extract_sales import run_extract_sales_task @patch('tasks.sales_ingest.run_extract_sales._ecs_credentials', None) @patch('tasks.sales_ingest.run_extract_sales.time') @patch('tasks.sales_ingest.run_extract_sales.config') @patch('tasks.sales_ingest.run_extract_sales.get_security_group') @patch('tasks.sales_ingest.run_extract_sales.get_subnet') @patch('tasks.sales_ingest.run_extract_sales.get_vpc') @patch('tasks.sales_ingest.run_extract_sales.get_ec2_client') @patch('tasks.sales_ingest.run_extract_sales.get_ecs_client') @patch('tasks.sales_ingest.run_extract_sales.get_credentials_for_assumed_role') @patch('tasks.sales_ingest.run_extract_sales.get_sts_client') @patch('tasks.sales_ingest.run_extract_sales.get_abacus_event') def test_run_extract_sales_task( mock_get_abacus_event, mock_get_sts_client, mock_get_credentials, mock_get_ecs_client, mock_get_ec2_client, mock_get_vpc, mock_get_subnet, mock_get_security_group, mock_config, mock_time ): """Test running the extract sales task.""" mock_dag_run = MagicMock() mock_task_instance = MagicMock() mock_task_instance.xcom_pull.return_value = 'distro' mock_get_abacus_event.return_value.target_id = 123 mock_sts_client = MagicMock() mock_get_sts_client.return_value = mock_sts_client mock_get_credentials.return_value = { 'AccessKeyId': 'test_access_key_id', 'SecretAccessKey': 'test_secret_access_key', 'SessionToken': 'test_session_token', 'Expiration': datetime(3000, 1, 1) } mock_ecs_client = MagicMock() mock_get_ecs_client.return_value = mock_ecs_client mock_ec2_client = MagicMock() mock_get_ec2_client.return_value = mock_ec2_client mock_get_vpc.return_value = 'VPC_ID' mock_get_subnet.return_value = 'SUBNET_ID' mock_get_security_group.return_value = 'SECURITY_GROUP_ID' mock_ecs_client.run_task.return_value = {'tasks': [{'taskArn': 'TASK_ARN'}]} # The `describe_tasks` call should be made twice. # The first time the task is running. # The second one the task has successfully terminated. mock_ecs_client.describe_tasks.side_effect = [ {'tasks': [{'lastStatus': 'RUNNING'}]}, {'tasks': [{ 'lastStatus': 'STOPPED', 'containers': [{'exitCode': 0}], 'stopCode': 'STOP_CODE', 'stoppedReason': 'STOPPED_REASON' }]} ] mock_config.EXTRACT_SALES_DAG_EXECUTION_ROLE = 'TEST_ROLE' mock_config.EXTRACT_SALES_ECS_SERVICE_NAME = 'TEST_SERVICE' mock_time.sleep.return_value = None run_extract_sales_task(mock_dag_run, task_instance=mock_task_instance) mock_get_abacus_event.assert_called_once_with( mock_dag_run, task_instance=mock_task_instance ) mock_task_instance.xcom_pull.assert_called_once_with( task_ids='determine_sales_ingest_type', key='sales_ingest_type' ) mock_get_credentials.assert_called_once_with( mock_sts_client, role_arn='TEST_ROLE', session_name='run_extract_sales_session' ) mock_get_vpc.assert_called_once_with(mock_ec2_client) mock_get_subnet.assert_called_once_with(mock_ec2_client, 'VPC_ID') mock_get_security_group.assert_called_once_with( mock_ec2_client, 'VPC_ID', 'TEST_SERVICE' ) mock_ecs_client.run_task.assert_called_once_with( cluster='TEST_SERVICE', taskDefinition='TEST_SERVICE', launchType='FARGATE', networkConfiguration={ 'awsvpcConfiguration': { 'subnets': ['SUBNET_ID'], 'securityGroups': ['SECURITY_GROUP_ID'], 'assignPublicIp': 'DISABLED' } }, overrides={ 'containerOverrides': [ { 'name': 'ecs-abacus-extract-sales', 'environment': [ { 'name': 'SALES_TYPE', 'value': 'distro' }, { 'name': 'BATCH_ID', 'value': '123' } ] } ] } ) mock_ecs_client.describe_tasks.assert_has_calls([ call(cluster='TEST_SERVICE', tasks=['TASK_ARN']), call(cluster='TEST_SERVICE', tasks=['TASK_ARN']) ]) @patch('tasks.sales_ingest.run_extract_sales._ecs_credentials', None) @patch('tasks.sales_ingest.run_extract_sales.time') @patch('tasks.sales_ingest.run_extract_sales.config') @patch('tasks.sales_ingest.run_extract_sales.get_security_group') @patch('tasks.sales_ingest.run_extract_sales.get_subnet') @patch('tasks.sales_ingest.run_extract_sales.get_vpc') @patch('tasks.sales_ingest.run_extract_sales.get_ec2_client') @patch('tasks.sales_ingest.run_extract_sales.get_ecs_client') @patch('tasks.sales_ingest.run_extract_sales.get_credentials_for_assumed_role') @patch('tasks.sales_ingest.run_extract_sales.get_sts_client') @patch('tasks.sales_ingest.run_extract_sales.get_abacus_event') def test_run_extract_sales_task_error( mock_get_abacus_event, mock_get_sts_client, mock_get_credentials, mock_get_ecs_client, mock_get_ec2_client, mock_get_vpc, mock_get_subnet, mock_get_security_group, mock_config, mock_time ): """Test running the extract sales task when an error occurrs.""" mock_dag_run = MagicMock() mock_task_instance = MagicMock() mock_task_instance.xcom_pull.return_value = 'distro' mock_get_abacus_event.return_value.target_id = 123 mock_sts_client = MagicMock() mock_get_sts_client.return_value = mock_sts_client mock_get_credentials.return_value = { 'AccessKeyId': 'test_access_key_id', 'SecretAccessKey': 'test_secret_access_key', 'SessionToken': 'test_session_token', 'Expiration': datetime(3000, 1, 1) } mock_ecs_client = MagicMock() mock_get_ecs_client.return_value = mock_ecs_client mock_ec2_client = MagicMock() mock_get_ec2_client.return_value = mock_ec2_client mock_get_vpc.return_value = 'VPC_ID' mock_get_subnet.return_value = 'SUBNET_ID' mock_get_security_group.return_value = 'SECURITY_GROUP_ID' mock_ecs_client.run_task.return_value = {'tasks': [{'taskArn': 'TASK_ARN'}]} # The `describe_tasks` call should be made twice. # The first time the task is running. # The second one the task has successfully terminated. mock_ecs_client.describe_tasks.side_effect = [ {'tasks': [{'lastStatus': 'RUNNING'}]}, {'tasks': [{ 'lastStatus': 'STOPPED', 'containers': [{'exitCode': 1}], 'stopCode': 'STOP_CODE', 'stoppedReason': 'STOPPED_REASON' }]} ] mock_config.EXTRACT_SALES_DAG_EXECUTION_ROLE = 'TEST_ROLE' mock_config.EXTRACT_SALES_ECS_SERVICE_NAME = 'TEST_SERVICE' mock_time.sleep.return_value = None with pytest.raises(Exception): run_extract_sales_task(mock_dag_run, task_instance=mock_task_instance) mock_get_abacus_event.assert_called_once_with( mock_dag_run, task_instance=mock_task_instance ) mock_task_instance.xcom_pull.assert_called_once_with( task_ids='determine_sales_ingest_type', key='sales_ingest_type' ) mock_get_credentials.assert_called_once_with( mock_sts_client, role_arn='TEST_ROLE', session_name='run_extract_sales_session' ) mock_get_vpc.assert_called_once_with(mock_ec2_client) mock_get_subnet.assert_called_once_with(mock_ec2_client, 'VPC_ID') mock_get_security_group.assert_called_once_with( mock_ec2_client, 'VPC_ID', 'TEST_SERVICE' ) mock_ecs_client.run_task.assert_called_once_with( cluster='TEST_SERVICE', taskDefinition='TEST_SERVICE', launchType='FARGATE', networkConfiguration={ 'awsvpcConfiguration': { 'subnets': ['SUBNET_ID'], 'securityGroups': ['SECURITY_GROUP_ID'], 'assignPublicIp': 'DISABLED' } }, overrides={ 'containerOverrides': [ { 'name': 'ecs-abacus-extract-sales', 'environment': [ { 'name': 'SALES_TYPE', 'value': 'distro' }, { 'name': 'BATCH_ID', 'value': '123' } ] } ] } ) mock_ecs_client.describe_tasks.assert_has_calls([ call(cluster='TEST_SERVICE', tasks=['TASK_ARN']), call(cluster='TEST_SERVICE', tasks=['TASK_ARN']) ]) @freeze_time('2025-01-01 10:00:00Z') @patch( 'tasks.sales_ingest.run_extract_sales._ecs_credentials', { 'AccessKeyId': 'test_access_key_id', 'SecretAccessKey': 'test_secret_access_key', 'SessionToken': 'test_session_token', 'Expiration': datetime(2025, 1, 1, 11, 0, 0, tzinfo=timezone.utc) } ) @patch('tasks.sales_ingest.run_extract_sales.get_credentials_for_assumed_role') def test_get_ecs_credentials_valid(mock_get_credentials): """Test getting the ECS credentials when they are valid.""" result = _get_ecs_credentials() assert result['AccessKeyId'] == 'test_access_key_id' mock_get_credentials.assert_not_called() @freeze_time('2025-01-01 10:59:10Z') @patch( 'tasks.sales_ingest.run_extract_sales._ecs_credentials', { 'AccessKeyId': 'test_access_key_id', 'SecretAccessKey': 'test_secret_access_key', 'SessionToken': 'test_session_token', 'Expiration': datetime(2025, 1, 1, 11, 0, 0, tzinfo=timezone.utc) } ) @patch('tasks.sales_ingest.run_extract_sales.get_credentials_for_assumed_role') def test_get_ecs_credentials_expired(mock_get_credentials): """Test getting the ECS credentials when they have expired.""" mock_get_credentials.return_value = { 'AccessKeyId': 'test_access_key_id_refreshed', 'SecretAccessKey': 'test_secret_access_key', 'SessionToken': 'test_session_token', 'Expiration': datetime(3000, 1, 1) } result = _get_ecs_credentials() assert result['AccessKeyId'] == 'test_access_key_id_refreshed' mock_get_credentials.assert_called_once()