"""Test airflow connector.""" import base64 from unittest.mock import MagicMock, patch import pytest from werkzeug.exceptions import BadRequest from abacus_event.connectors.airflow import ( _accounting_period_close_dag_config, _accounting_period_mechanicals_dag_config, _accounting_run_calculate_nr_dag_config, _adjustment_file_generate_dag_config, _adjustment_file_import_dag_config, _adjustment_file_upload_dag_config, _apply_adjustments_dag_config, _payments_generate_dag_config, _payments_generate_export_dag_config, _payments_upload_approval_dag_config, _payoneer_payments_payout_dag_config, _sales_ingest_dag_config, ensure_dag_not_running_in_aws, handle_event_actions, trigger_dag_from_aws, ) from abacus_event.constants.constants import DAG_IDS, EVENT_NAMES from tests.utils.factories import AbacusEventFactory def test__accounting_period_close_dag_config(): """Test accounting_period_close callback.""" mock_event = AbacusEventFactory.create( event_name=EVENT_NAMES.ACCOUNTING_PERIOD_CLOSE, target_type='accounting_period' ) dag_name, config = _accounting_period_close_dag_config(mock_event) assert dag_name == DAG_IDS.ACCOUNTING_PERIOD_CLOSE assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__payments_generate_dag_config(): """Test payments_generate callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _payments_generate_dag_config(mock_event) assert dag_name == DAG_IDS.PAYMENTS_GENERATE assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__payments_generate_export_dag_config(): """Test payments_generate_export callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _payments_generate_export_dag_config(mock_event) assert dag_name == DAG_IDS.PAYMENTS_GENERATE_EXPORT assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__payments_upload_approval_dag_config(): """Test payments_upload_approval callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _payments_upload_approval_dag_config(mock_event) assert dag_name == DAG_IDS.PAYMENTS_UPLOAD_APPROVAL assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__payoneer_payments_payout_dag_config(): """Test payoneer_payments_payout callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _payoneer_payments_payout_dag_config(mock_event) assert dag_name == DAG_IDS.PAYONEER_PAYMENTS_PAYOUT assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__accounting_period_mechanicals_dag_config(): """Test accounting_period_mechanicals callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _accounting_period_mechanicals_dag_config(mock_event) assert dag_name == DAG_IDS.ACCOUNTING_PERIOD_MECHANICALS assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__apply_pending_adjustments_dag_config(): """Test apply_adjustments callback.""" mock_event = AbacusEventFactory.create( event_name='apply_pending_adjustments', target_type='statement_period_adjustment_file', ) dag_name, config = _apply_adjustments_dag_config(mock_event) assert dag_name == DAG_IDS.APPLY_PENDING_ADJUSTMENTS assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__adjustment_file_upload_dag_config(): """Test adjustment_file_upload callback.""" mock_event = AbacusEventFactory.create() dag_name, config = _adjustment_file_upload_dag_config(mock_event) assert dag_name == DAG_IDS.ADJUSTMENT_FILE_UPLOAD assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__accounting_run_calculate_nr_dag_config(): """Test that the accounting_run_calculate_nr DAG is called.""" mock_event = AbacusEventFactory.create() dag_name, config = _accounting_run_calculate_nr_dag_config(mock_event) assert dag_name == DAG_IDS.ACCOUNTING_RUN_CALCULATE_NR assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__adjustment_file_generate_dag_config(): """Test that the adjustment_file_generate DAG is called.""" mock_event = AbacusEventFactory.create() dag_name, config = _adjustment_file_generate_dag_config(mock_event) assert dag_name == DAG_IDS.ADJUSTMENT_FILE_GENERATE assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__adjustment_file_import_dag_config(): """Test that the adjustment_file_import DAG is called.""" mock_event = AbacusEventFactory.create() dag_name, config = _adjustment_file_import_dag_config(mock_event) assert dag_name == DAG_IDS.ADJUSTMENT_FILE_IMPORT assert config.get('abacus_event_id') == mock_event.abacus_event_id def test__sales_ingest_dag_config(): """Test that the sales_ingest DAG is called.""" mock_event = AbacusEventFactory.create() dag_name, config = _sales_ingest_dag_config(mock_event) assert dag_name == DAG_IDS.SALES_INGEST assert config.get('abacus_event_id') == mock_event.abacus_event_id @pytest.fixture def mock_event_handlers(): """Mock event handlers.""" mock_handlers = { EVENT_NAMES.ACCOUNTING_PERIOD_CLOSE: MagicMock(return_value=('close', {})), EVENT_NAMES.ACCOUNTING_RUN_COMMIT: MagicMock(return_value=('foo', {})), EVENT_NAMES.PAYMENTS_GENERATE: MagicMock(return_value=('bar', {})), EVENT_NAMES.SEND_PAYMENTS: MagicMock(return_value=('test', {})), EVENT_NAMES.ACCOUNTING_RUN_CALCULATE: MagicMock(return_value=('baz', {})), EVENT_NAMES.ACCOUNTING_RUN_CALCULATE_NR: MagicMock(return_value=('NR', {})), EVENT_NAMES.ADJUSTMENT_FILE_WORKSHEET_IMPORT: MagicMock( return_value=('imported', {}) ), EVENT_NAMES.SALES_INGEST_DISTRO: MagicMock( return_value=('ingest_sales_distro', {}) ), EVENT_NAMES.SALES_INGEST_NR: MagicMock(return_value=('ingest_sales_nr', {})), } with patch('abacus_event.connectors.airflow.event_handlers', mock_handlers): yield mock_handlers @patch('abacus_event.connectors.airflow.ensure_dag_not_running_in_aws') @patch('abacus_event.connectors.airflow.trigger_dag_from_aws') def test_event_handler(trigger_mock, ensure_mock, mock_event_handlers): """Test event handler behavior.""" mock_commit_event = AbacusEventFactory.create( event_name=EVENT_NAMES.ACCOUNTING_RUN_COMMIT ) handle_event_actions(mock_commit_event) assert mock_event_handlers[EVENT_NAMES.ACCOUNTING_RUN_COMMIT].call_count == 1 mock_payment_event = AbacusEventFactory.create( event_name=EVENT_NAMES.PAYMENTS_GENERATE ) handle_event_actions(mock_payment_event) assert mock_event_handlers[EVENT_NAMES.PAYMENTS_GENERATE].call_count == 1 assert ensure_mock.call_count == 2 assert trigger_mock.call_count == 2 @patch('abacus_event.connectors.airflow.ensure_dag_not_running_in_aws') @patch('abacus_event.connectors.airflow.trigger_dag_from_aws') def test_event_handler_triggers_accounting_period_close_dag( mock_trigger_dag, mock_ensure_not_running, mock_event_handlers ): """Test that dag is triggered when 'accounting_period_close' event is created.""" mock_close_accounting_period_event = AbacusEventFactory.create( event_name=EVENT_NAMES.ACCOUNTING_PERIOD_CLOSE ) handle_event_actions(mock_close_accounting_period_event) mock_event_handlers[EVENT_NAMES.ACCOUNTING_PERIOD_CLOSE].assert_called_once() mock_ensure_not_running.assert_called_once() mock_trigger_dag.assert_called_once() @patch('abacus_event.connectors.airflow.ensure_dag_not_running_in_aws') @patch('abacus_event.connectors.airflow.trigger_dag_from_aws') def test_event_handler_triggers_sales_ingest_dag_with_distro( mock_trigger_dag, mock_ensure_not_running, mock_event_handlers ): """Test that dag is triggered when an `ingest_sales_distro` event is created.""" mock_sales_ingest_distro_event = AbacusEventFactory.create( event_name=EVENT_NAMES.SALES_INGEST_DISTRO ) handle_event_actions(mock_sales_ingest_distro_event) mock_event_handlers[EVENT_NAMES.SALES_INGEST_DISTRO].assert_called_once() mock_ensure_not_running.assert_called_once() mock_trigger_dag.assert_called_once() @patch('abacus_event.connectors.airflow.ensure_dag_not_running_in_aws') @patch('abacus_event.connectors.airflow.trigger_dag_from_aws') def test_event_handler_triggers_sales_ingest_dag_with_nr( mock_trigger_dag, mock_ensure_not_running, mock_event_handlers ): """Test that dag is triggered when an `ingest_sales_nr` event is created.""" mock_sales_ingest_nr_event = AbacusEventFactory.create( event_name=EVENT_NAMES.SALES_INGEST_NR ) handle_event_actions(mock_sales_ingest_nr_event) mock_event_handlers[EVENT_NAMES.SALES_INGEST_NR].assert_called_once() mock_ensure_not_running.assert_called_once() mock_trigger_dag.assert_called_once() @patch('abacus_event.connectors.airflow._aws_airflow_request') def test_failed_airflow_trigger_raises(mock_airflow_request, mock_event_handlers): """Test a failed invocation.""" mock_airflow_request.side_effect = BadRequest('Error') with pytest.raises(BadRequest): mock_commit_event = AbacusEventFactory.create( event_name=EVENT_NAMES.ACCOUNTING_RUN_COMMIT ) handle_event_actions(mock_commit_event) @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_trigger_dag_from_aws(boto_mock, requests_mock): """Test triggering a dag from AWS Airflow.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': '', 'stdout': base64.b64encode(b'success'), } requests_mock.post.return_value = result_mock config = {'event_name': 'EventName', 'target_id': 1} dag_result = trigger_dag_from_aws('super_dag', config) assert dag_result == 'success' requests_mock.post.assert_called_once_with( 'https://mock-host-name/aws_mwaa/cli', data="""dags trigger -c '{"event_name": "EventName", "target_id": 1}' super_dag""", headers={'Authorization': 'Bearer mock-token', 'Content-Type': 'text/plain'}, ) @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_trigger_dag_from_aws_with_deprecation_warnings(boto_mock, requests_mock): """Test triggering a dag if deprecation warnings are in error output.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': base64.b64encode(b'message DeprecationWarning: description'), 'stdout': base64.b64encode(b'success'), } requests_mock.post.return_value = result_mock config = {'event_name': 'EventName', 'target_id': 1} dag_result = trigger_dag_from_aws('super_dag', config) assert dag_result == 'success' @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') @patch('abacus_event.connectors.airflow.uuid') def test_trigger_parallel_dag_from_aws(uuid_mock, boto_mock, requests_mock): """Test triggering a dag that can run in parallel from AWS Airflow.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': '', 'stdout': base64.b64encode(b'success'), } requests_mock.post.return_value = result_mock uuid_mock.uuid4.return_value = '32ff0b41-3259-4106' config = {'event_name': 'EventName', 'target_id': 1} dag_result = trigger_dag_from_aws('legacy_sync_contract', config) assert dag_result == 'success' requests_mock.post.assert_called_once_with( 'https://mock-host-name/aws_mwaa/cli', data="""dags trigger -r legacy_sync_contract_32ff0b41-3259-4106 -c '{"event_name": "EventName", "target_id": 1}' legacy_sync_contract""", headers={'Authorization': 'Bearer mock-token', 'Content-Type': 'text/plain'}, ) @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_ensure_dag_not_running_in_aws_does_not_allow_concurrency( boto_mock, requests_mock ): """Test function raises if non-parallel dag is already running in AWS Airflow.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': '', 'stdout': base64.b64encode(b'["run1", "run2"]\n'), } requests_mock.post.return_value = result_mock with pytest.raises(BadRequest) as err: ensure_dag_not_running_in_aws('awesome_dag') assert str(err.value) == "400 Bad Request: DAG 'awesome_dag' already in progress." @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_ensure_dag_not_running_in_aws(boto_mock, requests_mock): """Test function returns True if dag is not running in AWS Airflow.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = {'stderr': '', 'stdout': base64.b64encode(b'[]\n')} requests_mock.post.return_value = result_mock result = ensure_dag_not_running_in_aws('awesome_dag') assert result @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_ensure_dag_not_running_in_aws_allows_parallel_dags(boto_mock, requests_mock): """Test function returns True if parallel dag does not exceed active runs limit.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': '', 'stdout': base64.b64encode(b'["run1", "run2"]\n'), } requests_mock.post.return_value = result_mock result = ensure_dag_not_running_in_aws('legacy_sync_contract') assert result @patch('abacus_event.connectors.airflow.requests') @patch('abacus_event.connectors.airflow.boto3') def test_ensure_dag_not_running_in_aws_checks_parallel_dags_concurrency( boto_mock, requests_mock ): """Test function raises if parallel dag exceeds active runs limit.""" client_mock = MagicMock() client_mock.create_cli_token.return_value = { 'WebServerHostname': 'mock-host-name', 'CliToken': 'mock-token', } boto_mock.client.return_value = client_mock result_mock = MagicMock(status_code=200) result_mock.json.return_value = { 'stderr': '', 'stdout': base64.b64encode( b'[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]\n' ), } requests_mock.post.return_value = result_mock with pytest.raises(BadRequest) as err: ensure_dag_not_running_in_aws('legacy_sync_contract') assert ( str(err.value) == "400 Bad Request: DAG 'legacy_sync_contract' already in progress." )