"""Snapshot contract processor unit tests.""" from unittest.mock import MagicMock, patch import pandas from snapshot_contract import constants from snapshot_contract.processor import SnapshotContractProcessor @patch('snapshot_contract.processor.flatten_contract_data') def test_combine_contract_data( mock_flatten_contract_data, mock_contract, mock_contract_term, mock_isrcs, mock_event, mock_flat_contract_data ): """Test main process method.""" accounting_run_id = mock_event.get('target_id') mock_flatten_contract_data.return_value = mock_flat_contract_data processor = SnapshotContractProcessor(mock_event) contract_df = pandas.DataFrame([mock_contract]) contract_term_df = pandas.DataFrame([mock_contract_term]) isrc_df = pandas.DataFrame(mock_isrcs) res = processor._combine_contract_data(contract_df, contract_term_df, isrc_df) assert res == mock_flat_contract_data mock_flatten_contract_data.assert_called_once_with( accounting_run_id, contract_df, contract_term_df, isrc_df ) @patch('snapshot_contract.processor.app_logger') @patch('snapshot_contract.processor.awswrangler') def test_get_contracts(mock_awswrangler, mock_logger, mock_contract, mock_event): """Test reading contract csv from s3.""" contract_df = pandas.DataFrame([mock_contract]) mock_awswrangler.s3.read_csv.return_value = contract_df processor = SnapshotContractProcessor(mock_event) res = processor._get_contracts() assert not res.empty assert type(res) == type(contract_df) assert mock_logger.info.call_count == 2 mock_awswrangler.s3.read_csv.mock_called_once() @patch('snapshot_contract.processor.app_logger') @patch('snapshot_contract.processor.awswrangler') def test_get_contract_terms( mock_awswrangler, mock_logger, mock_contract_term, mock_event ): """Test reading contract_term csv from s3.""" contract_term_df = pandas.DataFrame([mock_contract_term]) mock_awswrangler.s3.read_csv.return_value = contract_term_df processor = SnapshotContractProcessor(mock_event) res = processor._get_contract_terms() assert not res.empty assert type(res) == type(contract_term_df) assert mock_logger.info.call_count == 2 mock_awswrangler.s3.read_csv.mock_called_once() @patch('snapshot_contract.processor.app_logger') @patch('snapshot_contract.processor.awswrangler') def test_get_isrcs( mock_awswrangler, mock_logger, mock_isrcs, mock_event ): """Test reading isrc parquet from s3.""" isrc_df = pandas.DataFrame(mock_isrcs) mock_awswrangler.s3.read_parquet.return_value = isrc_df processor = SnapshotContractProcessor(mock_event) res = processor._get_isrcs() assert not res.empty assert type(res) == type(isrc_df) assert mock_logger.info.call_count == 2 mock_awswrangler.s3.read_csv.mock_called_once() @patch('snapshot_contract.processor.app_logger') def test_log_init(mock_logger, mock_event): """Test log init is called when processor is initialized.""" SnapshotContractProcessor(mock_event) target_id = mock_event['target_id'] mock_logger.info.assert_called_once_with(constants.PROCESSOR_MSG.format(target_id)) @patch('snapshot_contract.processor.app_logger') def test_log_success(mock_logger, mock_event): """Test processor logs a success message.""" processor = SnapshotContractProcessor(mock_event) processor._log_success() assert mock_logger.info.call_count == 2 mock_logger.info.assert_called_with(constants.SUCCESS_MSG) @patch('snapshot_contract.processor.get_accounting_run') @patch('snapshot_contract.processor.helpers') def test_set_s3_path(mock_helpers, mock_get_run, mock_event): """Test setting s3 path based on event's accounting_run.""" acct_run_id = mock_event.get('target_id') mock_acct_run = { 'accounting_period_id': 12, 'accounting_period_name': 'Acct Period Name', 'accounting_run_id': acct_run_id, 'run_controller_name': 'RC Name' } s3_path = f's3://test-bucket/12-acct-period-name/{acct_run_id}-rc-name/snapshots' mock_get_run.return_value = mock_acct_run mock_helpers.build_snapshot_path.return_value = s3_path processor = SnapshotContractProcessor(mock_event) processor._set_s3_path() assert processor.s3_path == s3_path mock_get_run.assert_called_once_with(acct_run_id) mock_helpers.build_snapshot_path.assert_called_once_with(mock_acct_run) @patch('snapshot_contract.processor.app_logger') @patch('snapshot_contract.processor.awswrangler') @patch('snapshot_contract.processor.helpers') def test_write_to_s3( mock_helpers, mock_awswrangler, mock_logger, mock_event, mock_flat_contract_data ): """Test processor writes combined contract data to s3 as parquet file.""" flat_df = pandas.DataFrame(mock_flat_contract_data) mock_awswrangler.s3.to_parquet.return_value = {'paths': ['s3_path']} mock_helpers.apply_contract_schema.return_value = flat_df processor = SnapshotContractProcessor(mock_event) processor._write_to_s3(mock_flat_contract_data) assert mock_logger.info.call_count == 2 mock_helpers.apply_contract_schema.assert_called_once() mock_awswrangler.s3.to_parquet.assert_called_once() def test_process( mock_contract, mock_contract_term, mock_event, mock_isrcs, mock_flat_contract_data): """Test SnapshotContractProcessor's main processor method.""" contract_df = pandas.DataFrame([mock_contract]) contract_term_df = pandas.DataFrame([mock_contract_term]) isrc_df = pandas.DataFrame(mock_isrcs) SnapshotContractProcessor._combine_contract_data = \ MagicMock(return_value=mock_flat_contract_data) SnapshotContractProcessor._get_contracts = MagicMock(return_value=contract_df) SnapshotContractProcessor._get_contract_terms = \ MagicMock(return_value=contract_term_df) SnapshotContractProcessor._get_isrcs = MagicMock(return_value=isrc_df) SnapshotContractProcessor._log_init = MagicMock() SnapshotContractProcessor._log_success = MagicMock() SnapshotContractProcessor._set_s3_path = MagicMock() SnapshotContractProcessor._write_to_s3 = MagicMock() processor = SnapshotContractProcessor(mock_event) processor.process() assert processor.event == mock_event SnapshotContractProcessor._log_init.assert_called_once() SnapshotContractProcessor._set_s3_path.assert_called_once() SnapshotContractProcessor._get_contracts.assert_called_once() SnapshotContractProcessor._get_contract_terms.assert_called_once() SnapshotContractProcessor._get_isrcs.assert_called_once() SnapshotContractProcessor._combine_contract_data.assert_called_once_with( contract_df, contract_term_df, isrc_df ) SnapshotContractProcessor._write_to_s3.assert_called_once_with( mock_flat_contract_data ) SnapshotContractProcessor._log_success.assert_called_once()