"""Unit tests for base_flow.py.""" import datetime import os from unittest.mock import call from unittest.mock import Mock from unittest.mock import patch from freezegun import freeze_time import pytest from yt_conflict_elasticsearch.flows import base_flow @pytest.mark.parametrize('env', [None, 'dev']) def test_get_domain_dev(monkeypatch, env): """Test get_domain for dev env.""" # Mocking os_mock = Mock() os_mock.getenv.side_effect = [env, 'dev'] monkeypatch.setattr( 'yt_conflict_elasticsearch.flows.base_flow.os', os_mock) # Test function call domain = base_flow.get_domain() # Checks assert domain == 'dev' os_mock.getenv.assert_has_calls([ call('Environment'), call('SWF_DOMAIN', 'dev')]) @pytest.mark.parametrize('env', ['qa', 'prod']) def test_get_domain_not_dev(monkeypatch, env): """Test get_domain for not dev env.""" # Mocking os_mock = Mock() os_mock.getenv.side_effect = [env, 'dev'] monkeypatch.setattr( 'yt_conflict_elasticsearch.flows.base_flow.os', os_mock) # Test function call domain = base_flow.get_domain() # Checks assert domain == '{}_{}'.format(env, 'yt_conflict_elasticsearch') assert len(os_mock.getenv.mock_calls) == 1 @patch.dict(os.environ, {'AWS_REGION': 'us-east-1'}) def test_base_flow(monkeypatch, flow_mock): """Test BaseFlow class __init__ method.""" # Mocking # Test class method flow = base_flow.BaseFlow( flow_name=flow_mock.flow_name, version=flow_mock.flow_version, swf_domain=flow_mock.flow_domain) # Checks assert flow.domain == flow_mock.flow_domain assert flow.name == flow_mock.flow_name assert flow.version == flow_mock.flow_version assert flow.sentry_client is None @patch.dict(os.environ, {'AWS_REGION': 'us-east-1'}) def test_base_flow_default_domain(monkeypatch, flow_mock): """Test BaseFlow class __init__ method with default domain.""" # Mocking get_domain_mock = Mock() domain_mock = get_domain_mock.return_value monkeypatch.setattr( 'yt_conflict_elasticsearch.flows.base_flow.get_domain', get_domain_mock) # Test class method flow = base_flow.BaseFlow( flow_name=flow_mock.flow_name, version=flow_mock.flow_version) # Checks assert flow.domain == domain_mock assert flow.name == flow_mock.flow_name assert flow.version == flow_mock.flow_version assert flow.sentry_client is None @patch.dict(os.environ, {'AWS_REGION': 'us-east-1', 'SENTRY_DSN': 'dsn'}) def test_base_flow_sentry_client(monkeypatch, flow_mock): """Test BaseFlow class __init__ method, check sentry client.""" # Mocking sentry_mock = Mock() sentry_client = sentry_mock.init.return_value monkeypatch.setattr( 'yt_conflict_elasticsearch.flows.base_flow.sentry_sdk', sentry_mock) # Test class method flow = base_flow.BaseFlow( flow_name=flow_mock.flow_name, version=flow_mock.flow_version, swf_domain=flow_mock.flow_domain) # Checks assert flow.domain == flow_mock.flow_domain assert flow.name == flow_mock.flow_name assert flow.version == flow_mock.flow_version assert flow.sentry_client == sentry_client @patch.dict(os.environ, {'AWS_REGION': 'us-east-1'}) @freeze_time('2017-10-01') def test_base_flow_workflow_id(monkeypatch, flow_mock): """Test BaseFlow calls workflow_id method.""" # Mocking date = datetime.datetime.today().strftime('%Y-%m-%d') expected_workflow_id = '{flow_name}-{date}'.format( flow_name=flow_mock.flow_name, date=date) # Test class method flow = base_flow.BaseFlow( flow_name=flow_mock.flow_name, version=flow_mock.flow_version, swf_domain=flow_mock.flow_domain) workflow_id = flow.workflow_id({}) # Checks assert workflow_id == expected_workflow_id @patch.dict(os.environ, {'AWS_REGION': 'us-east-1'}) def test_base_flow_workflow_id_with_context_date(monkeypatch, flow_mock): """Test BaseFlow calls workflow_id method passed context_date.""" # Mocking expected_workflow_id = '{flow_name}-{date}'.format( flow_name=flow_mock.flow_name, date='2017-10-01') # Test class method flow = base_flow.BaseFlow( flow_name=flow_mock.flow_name, version=flow_mock.flow_version, swf_domain=flow_mock.flow_domain) workflow_id = flow.workflow_id({'context_date': '2017-10-01'}) # Checks assert workflow_id == expected_workflow_id