"""Unit tests for Base Workflow class.""" import datetime from unittest.mock import call, MagicMock, patch import boto3 from garcon import activity import pytest from feed_ingestion.flows import base @patch('garcon.activity.create') def test_flow_base_init(mock_activity, monkeypatch): """Test FlowBase init sets domain, name, version, activity properties.""" monkeypatch.setattr( base, 'get_domain', MagicMock(return_value='domain')) monkeypatch.setattr( base, 'generate_feed_name', MagicMock(return_value='feed_name')) flow = base.FlowBase('feed_name', '1.0') assert flow.domain == 'domain' assert flow.name == 'feed_name' assert flow.feed_name == 'feed_name' assert flow.version == '1.0' assert mock_activity.call_count == 1 @patch('feed_ingestion.flows.base.get_domain') @patch('garcon.activity.create') def test_flow_base_workflow_id(mock_domain, mock_create, monkeypatch): """Test FlowBase workflow_id method.""" monkeypatch.setattr( base, 'generate_feed_name', MagicMock(return_value='feed_name')) flow = base.FlowBase('feed_name', '1.0') # test for passed context date dt = '2015-10-01' context = {'context_date': dt} expected_result = 'feed_name-{date}'.format(date=dt) assert flow.workflow_id(context) == expected_result # test for missing context date (defaults to today) expected_result = 'feed_name-{date}'.format( date=datetime.datetime.today().strftime('%Y-%m-%d')) assert flow.workflow_id(initial_context={}) == expected_result @patch.object(base, 'sentry_util') @patch('feed_ingestion.flows.base.get_domain') @patch('feed_ingestion.flows.base.generate_feed_name') @patch('garcon.activity.create') @patch('feed_ingestion.flows.base.logger') def test_flow_base_exception( mock_logging, mock_create, mock_feed, mock_domain, mock_sentry_util): """Test FlowBase default on_exception logs to actor.logger.""" flow = base.FlowBase('feed_name', '1.0') exception = Exception('Some error') # Non Activity calls FlowBase logger actor_mock = MagicMock() flow.on_exception(actor_mock, exception) mock_logging.error.assert_called_with(exception, exc_info=True) # Assert Activity error calls activity's own logger client = boto3.client('swf', 'us-east-1') activity_obj = activity.Activity(client) on_exception_mock = MagicMock() activity_obj.logger.error = on_exception_mock flow.on_exception(activity_obj, 'exception') on_exception_mock.assert_called_with('exception', exc_info=True) # Sentry called if DSN is an os var with patch.dict('os.environ', {'SENTRY_DSN': 'https://sentryblah.com'}): mock_sentry_util.reset_mock() flow.on_exception(MagicMock(), exception) assert mock_sentry_util.send_error_or_warning.call_args_list == [ call(exception) ] mock_sentry_util.reset_mock() # Sentry called and warning was sent exception_1 = Exception( 'The activity failures has exceeded its retry limit.') flow.on_exception(MagicMock(), exception_1) assert mock_sentry_util.send_error_or_warning.call_args_list == [ call(exception_1) ] @patch('feed_ingestion.flows.base.get_domain') @patch('feed_ingestion.flows.base.generate_feed_name') @patch('garcon.activity.create') def test_flow_base_decider(mock_domain, mock_feed, mock_create): """Test FlowBase decider throws not implemented error.""" flow = base.FlowBase('feed_name', '1.0') with pytest.raises(NotImplementedError): flow.decider('schedule') class TestGetDomain(object): """Test get_domain.""" def test_dev_domain(self, monkeypatch): """Set swf domain for dev.""" monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'dev'})) assert base.get_domain() == 'dev' def test_prod_domain(self, monkeypatch): """Set swf domain for prod.""" monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'prod'})) assert base.get_domain() == 'prod_swf_feed_ingestion' def test_qa_domain(self, monkeypatch): """Set swf domain for qa.""" monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'qa'})) assert base.get_domain() == 'qa_swf_feed_ingestion' def test_test_domain(self, monkeypatch): """Set swf domain for test.""" monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'test'})) assert base.get_domain() == 'test' def test_shadow_domain(self, monkeypatch): """Set swf domain for shadow.""" monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'shadow'})) assert base.get_domain() == 'shadow_swf_feed_ingestion' def test_custom_dev_domain(self, monkeypatch): """Set swf domain to env domain in the dev domain.""" monkeypatch.setenv('SWF_DOMAIN', 'dev_custom') monkeypatch.setattr( base.conf, 'getconf', MagicMock(return_value={'env': 'dev'})) assert base.get_domain() == 'dev_custom' def test_generate_feed_name(): """Test generating feed name.""" # valid feedname feed_name = 'feed_name' expected_result = 'feed_name_feed_ingestion' actual_result = base.generate_feed_name(feed_name) assert actual_result == expected_result # invalid feedname feed_name = 'feed_Name' with pytest.raises(AssertionError): base.generate_feed_name(feed_name) # invalid feedname feed_name = 'feed name' with pytest.raises(AssertionError): base.generate_feed_name(feed_name) def test_base_create_kwargs_collect_activity_for(): """Test create_kwargs_collect_activity_for method of FlowBase class.""" feed_name = 'test_feed' flow_version = '1.0' flow = base.FlowBase(feed_name, flow_version) # as supplied to Activity.run by garcon activity_input = { 'key1': 'value1', 'namespace2.key2': 'value2', 'namespace3.key3': 'value3' } activity_with_kwargs = 'some_activity_accepting_kwargs' kwargs_collect_activity = flow.create_kwargs_collect_activity_for( activity_with_kwargs, mapping={ 'mapped_key2': 'namespace2.key2', 'mapped_key3': 'namespace3.key3', 'mapped_key4': 'key4' # not present in input context } ) output_context = {} def store_result(result): output_context['result'] = result expected_context = { 'some_activity_accepting_kwargs_kwargs.kwargs': { 'mapped_key2': 'value2', 'mapped_key3': 'value3', 'mapped_key4': None, } } with patch.multiple( activity.Activity, poll_for_activity=lambda self, x: MagicMock( context=activity_input, client=boto3.client('swf', 'us-east-1'), activity_id='activityId', complete=store_result)): kwargs_collect_activity.run() assert output_context['result'] == expected_context def test_get_aws_config(monkeypatch): """Test aws config overrides work.""" monkeypatch.setenv('AWS_ACCESS_KEY_ID', 'AWS_ACCESS_KEY_ID') monkeypatch.setenv('AWS_SECRET_ACCESS_KEY', 'AWS_SECRET_ACCESS_KEY') monkeypatch.setenv('AWS_SESSION_TOKEN', '') monkeypatch.setenv( 'AWS_WORKFLOW_ACCESS_KEY_ID', 'AWS_WORKFLOW_ACCESS_KEY_ID') monkeypatch.setenv( 'AWS_WORKFLOW_SECRET_ACCESS_KEY', 'AWS_WORKFLOW_SECRET_ACCESS_KEY') monkeypatch.setenv('Environment', 'dev') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_WORKFLOW_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_WORKFLOW_SECRET_ACCESS_KEY' assert aws_config['access_token'] == '' monkeypatch.setenv('Environment', 'qa') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == '' monkeypatch.setenv('Environment', 'prod') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == '' # ensure prod does not break if workflow env variables are not set monkeypatch.setenv('Environment', 'prod') monkeypatch.delenv('AWS_WORKFLOW_ACCESS_KEY_ID') monkeypatch.delenv('AWS_WORKFLOW_SECRET_ACCESS_KEY') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == '' def test_get_aws_config_token_based_creds(monkeypatch): """Test aws config overrides work.""" monkeypatch.setenv('AWS_ACCESS_KEY_ID', 'AWS_ACCESS_KEY_ID') monkeypatch.setenv('AWS_SECRET_ACCESS_KEY', 'AWS_SECRET_ACCESS_KEY') monkeypatch.setenv('AWS_SESSION_TOKEN', 'AWS_SESSION_TOKEN') monkeypatch.setenv( 'AWS_WORKFLOW_ACCESS_KEY_ID', 'AWS_WORKFLOW_ACCESS_KEY_ID') monkeypatch.setenv( 'AWS_WORKFLOW_SECRET_ACCESS_KEY', 'AWS_WORKFLOW_SECRET_ACCESS_KEY') monkeypatch.setenv( 'AWS_WORKFLOW_SESSION_TOKEN', 'AWS_WORKFLOW_SESSION_TOKEN') monkeypatch.setenv('Environment', 'dev') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_WORKFLOW_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_WORKFLOW_SECRET_ACCESS_KEY' assert aws_config['access_token'] == 'AWS_WORKFLOW_SESSION_TOKEN' monkeypatch.setenv('Environment', 'qa') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == 'AWS_SESSION_TOKEN' monkeypatch.setenv('Environment', 'prod') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == 'AWS_SESSION_TOKEN' # ensure prod does not break if workflow env variables are not set monkeypatch.setenv('Environment', 'prod') monkeypatch.delenv('AWS_WORKFLOW_ACCESS_KEY_ID') monkeypatch.delenv('AWS_WORKFLOW_SECRET_ACCESS_KEY') monkeypatch.delenv('AWS_WORKFLOW_SESSION_TOKEN') aws_config = base.get_aws_config() assert aws_config['access_key'] == 'AWS_ACCESS_KEY_ID' assert aws_config['access_secret'] == 'AWS_SECRET_ACCESS_KEY' assert aws_config['access_token'] == 'AWS_SESSION_TOKEN' class FakeActivity: """Minimal stand-in for garcon's Activity, just enough for execute().""" class logger: """Fake logger with a no-op debug().""" @staticmethod def debug(msg): """No-op debug log.""" @staticmethod def heartbeat(): pass class TestQueryTagRunnerMixin: """Test QueryTagRunnerMixin.execute query tag propagation.""" def test_execute_sets_query_tag_per_call(self): """Test each execute() call sees its own context's query tag. Regression test: a runner instance is built once per activity and its execute() is called repeatedly for the life of the worker process, so tasks must be re-tagged from the original, unwrapped task on every call rather than wrapping the previous call's already-wrapped task. """ observed_tags = [] def task_fn(context, activity=None): from feed_ingestion.util.query_tag import get_query_tag observed_tags.append(get_query_tag()) return {} runner = base.SyncRunner(task_fn) runner.execute(FakeActivity, { 'execution.workflow_id': 'wf-A', 'execution.run_id': 'run-A', 'backfill': 'True'}) runner.execute(FakeActivity, { 'execution.workflow_id': 'wf-B', 'execution.run_id': 'run-B', 'backfill': 'True'}) assert observed_tags == [ ('{"workflow_id": "wf-A", "run_id": "run-A", ' '"env": "dev", "backfill": true}'), ('{"workflow_id": "wf-B", "run_id": "run-B", ' '"env": "dev", "backfill": true}'), ] def test_execute_does_not_grow_wrapper_nesting(self): """Test tasks are not re-wrapped on top of each other every call.""" def task_fn(context, activity=None): return {} runner = base.SyncRunner(task_fn) activity_input = { 'execution.workflow_id': 'wf-A', 'execution.run_id': 'run-A', 'backfill': 'True', } runner.execute(FakeActivity, activity_input) runner.execute(FakeActivity, activity_input) runner.execute(FakeActivity, activity_input) assert runner._original_tasks == (task_fn,) assert len(runner.tasks) == 1 def test_backfill_flag(self): observed_tags = [] def task_fn(context, activity=None): from feed_ingestion.util.query_tag import get_query_tag observed_tags.append(get_query_tag()) return {} runner = base.SyncRunner(task_fn) runner.execute(FakeActivity, { 'execution.workflow_id': 'wf-A', 'execution.run_id': 'run-A', 'backfill': 'True', 'other': 'value' }) assert observed_tags == [ '{"workflow_id": "wf-A", "run_id": "run-A", ' '"env": "dev", "backfill": true}' ] def test_without_backfill_flag(self): observed_tags = [] def task_fn(context, activity=None): from feed_ingestion.util.query_tag import get_query_tag query_tag = get_query_tag() if query_tag: observed_tags.append(query_tag) return {} runner = base.SyncRunner(task_fn) runner.execute(FakeActivity, { 'execution.workflow_id': 'wf-A', 'execution.run_id': 'run-A', 'other': 'value' }) assert observed_tags == []