"""Abacus State logic tests.""" from unittest.mock import MagicMock, patch import pytest from marshmallow import ValidationError from owsresponse import response from abacus_state.constants import constants from abacus_state.logic import abacus_state as logic from abacus_state.schemas.abacus_state import AbacusStateDetailSchema from tests.utils.factories import AbacusStateFactory @patch('abacus_state.logic.abacus_state.AbacusState') def test_create_abacus_states(mock_model): """Test create_abacus_state method.""" mock_abacus_state = AbacusStateFactory.build(action_status='init') mock_model.build.return_value = [mock_abacus_state] mock_model.get_action_status_list.return_value = [] params = [ { 'action_name': mock_abacus_state.action_name, 'parent_table_id': mock_abacus_state.parent_table_id, 'parent_table_name': mock_abacus_state.parent_table_name, } ] result = logic.create_abacus_states(params) mock_model.build.assert_called_with( action_name=mock_abacus_state.action_name, parent_table_id=mock_abacus_state.parent_table_id, parent_table_name=mock_abacus_state.parent_table_name, ) assert result.status == 201 mock_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.AbacusState') def test_create_abacus_states_by_parent_table_success(mock_model, mock_validate): """Test successful creation of all abacus states given a parent_table_name.""" parent_table_name = constants.PARENT_TABLE_NAMES.ACCOUNTING_PERIOD parent_table_id = 123 contract_type = constants.CONTRACT_TYPES.DISTRIBUTION mock_validate.return_value = None mock_model.build.return_value = None mock_model.commit_changes.return_value = None res = logic.create_abacus_states_by_parent_table( parent_table_name, parent_table_id, contract_type ) assert res.status == 201 assert len(res.message) == len( constants.ACTIONS_BY_PARENT_TABLE.get(parent_table_name) ) mock_validate.assert_called_once_with(parent_table_name, parent_table_id) assert mock_model.build.call_count == len( constants.ACTIONS_BY_PARENT_TABLE.get(parent_table_name) ) mock_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.AbacusState') def test_create_abacus_state_by_parent_table_nr_accounting_period( mock_model, mock_validate ): """Test creation of NR accounting period abacus states.""" parent_table_name = constants.PARENT_TABLE_NAMES.ACCOUNTING_PERIOD parent_table_id = 123 contract_type = constants.CONTRACT_TYPES.NEIGHBOURING_RIGHTS mock_validate.return_value = None mock_model.commit_changes.return_value = None res = logic.create_abacus_states_by_parent_table( parent_table_name, parent_table_id, contract_type ) assert res.status == 201 assert len(res.message) == 4 assert len(res.message) < len(constants.ACCOUNTING_PERIOD_ACTION_NAMES) action_names = [c.kwargs['action_name'] for c in mock_model.build.call_args_list] assert ( constants.ACCOUNTING_PERIOD_ACTION_NAMES.PREP_MECHANICAL_DEDUCTIONS not in action_names ) mock_validate.assert_called_once_with(parent_table_name, parent_table_id) mock_model.build.call_count == len(res.message) mock_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.AbacusState') def test_create_abacus_states_by_parent_table_error(mock_model, mock_validate): """Test error is raised when create params are invalid.""" parent_table_name = 'accounting_period' parent_table_id = 123 mock_validate.side_effect = ValidationError('you shall not pass') res = logic.create_abacus_states_by_parent_table(parent_table_name, parent_table_id) assert res.status == 400 assert res.errors assert not res.message mock_validate.assert_called_once_with(parent_table_name, parent_table_id) mock_model.build.assert_not_called() mock_model.commit_changes.assert_not_called() @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.AbacusState') def test_create_abacus_state_by_parent_table_payment_group_payment( mock_model, mock_validate ): """Test creation of payment_group_payment states new flow.""" parent_table_name = constants.PARENT_TABLE_NAMES.PAYMENT_GROUP_PAYMENT parent_table_id = 123 mock_validate.return_value = None mock_model.commit_changes.return_value = None res = logic.create_abacus_states_by_parent_table(parent_table_name, parent_table_id) assert res.status == 201 action_names = [c.kwargs['action_name'] for c in mock_model.build.call_args_list] assert tuple(action_names) == constants.PAYMENT_GROUP_PAYMENT_ACTION_NAMES mock_validate.assert_called_once_with(parent_table_name, parent_table_id) mock_model.build.call_count == len(res.message) mock_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state.AbacusState') def test_get_action_status_list(mock_model): """Test that state action statuses are retrieved.""" mocked_result = { 'abacus_state_id': 1, 'action_name': 'contract_sync', 'action_status': 'init', 'parent_table_name': 'accounting_period', 'parent_table_id': 1, } mock_model.get_action_status_list.return_value = [mocked_result] result = logic.get_action_status_list( parent_table_name='accounting_period', parent_table_id=1 ) mock_model.get_action_status_list.assert_called_with('accounting_period', 1) assert result == [mocked_result] @patch('abacus_state.logic.abacus_state.update_account_payee_state') def test_update_abacus_state_account_payee(mock_update_account_payee_state): """Test update_abacus_state method for account_payee state.""" mock_update_account_payee_state.return_value = response.Response( message='ok', status=200 ) abacus_state = AbacusStateFactory.create( action_name=constants.ACCOUNT_PAYEE_ACTION_NAMES.PAYMENT_ELIGIBILITY, action_status=constants.ACTION_STATUSES.INIT, parent_table_name=constants.PARENT_TABLE_NAMES.ACCOUNT_PAYEE, parent_table_id=123, ) params = {'action_status': constants.ACTION_STATUSES.RUNNING} result = logic.update_abacus_state(abacus_state, **params) assert result.status == 200 mock_update_account_payee_state.assert_called_once_with(abacus_state, **params) @patch('abacus_state.logic.abacus_state.update_accounting_period_state') def test_update_abacus_state_for_acc_period(mock_accounting_period_state_logic): """Test update_abacus_state method for accounting_period actions.""" mock_accounting_period_state_logic.return_value = response.Response( message='ok', status=200 ) mock_abacus_state = AbacusStateFactory.build( action_status='init', parent_table_name='accounting_period' ) params = {'action_status': 'running'} result = logic.update_abacus_state(mock_abacus_state, **params) assert result.status == 200 mock_accounting_period_state_logic.assert_called_with(mock_abacus_state, **params) @patch('abacus_state.logic.abacus_state.update_sales_file_state') def test_update_abacus_state_for_sales_file(mock_sales_file_state_logic): """Test update_abacus_state method for sales_file actions.""" mock_sales_file_state_logic.return_value = response.Response( message='ok', status=200 ) mock_abacus_state = AbacusStateFactory.build( action_status='init', parent_table_name='sales_file' ) params = {'action_status': 'running'} result = logic.update_abacus_state(mock_abacus_state, **params) assert result.status == 200 mock_sales_file_state_logic.assert_called_with(mock_abacus_state, **params) @patch('abacus_state.logic.abacus_state.update_statement_period_state') def test_update_abacus_state_for_statement_period(mock_statement_period_state_logic): """Test update_abacus_state method for statement_period actions.""" mock_statement_period_state_logic.return_value = response.Response( message='ok', status=200 ) mock_abacus_state = AbacusStateFactory.build( action_status='init', parent_table_name='statement_period' ) params = {'action_status': 'running'} result = logic.update_abacus_state(mock_abacus_state, **params) assert result.status == 200 mock_statement_period_state_logic.assert_called_with(mock_abacus_state, **params) @patch('abacus_state.logic.abacus_state.update_statement_period_adjustment_file_state') def test_update_abacus_state_for_statement_period_adjustment_file(mock_logic): """Test update_abacus_state method for statement_period_adjustment_file actions.""" mock_logic.return_value = response.Response(message='ok', status=200) mock_abacus_state = AbacusStateFactory.build( action_status='init', parent_table_name='statement_period_adjustment_file' ) params = {'action_status': 'running'} result = logic.update_abacus_state(mock_abacus_state, **params) assert result.status == 200 mock_logic.assert_called_with(mock_abacus_state, **params) @patch('abacus_state.logic.abacus_state.update_accounting_run_state') def test_update_abacus_state_for_accounting_run(mock_logic): """Test update_abacus_state method for accounting_run actions.""" mock_logic.return_value = response.Response(message='ok', status=200) mock_abacus_state = AbacusStateFactory.build( action_status='init', parent_table_name='accounting_run' ) params = {'action_status': 'running'} result = logic.update_abacus_state(mock_abacus_state, **params) assert result.status == 200 mock_logic.assert_called_with(mock_abacus_state, **params) @patch( 'abacus_state.logic.abacus_state.bulk_update_payment_group_payment_account_states' ) def test_bulk_update_abacus_states_for_payment_account( mock_payment_account_state_logic, ): """Test bulk_update_abacus_states_by_parent_table method for send_payment action.""" parent_table_name = 'payment_group_payment_account' mock_payment_account_state_logic.return_value = response.Response( message='ok', status=200 ) AbacusStateFactory.build( action_status='init', parent_table_name=parent_table_name, parent_table_id=1 ) params = [{'action_status': 'running', 'parent_table_id': 1}] result = logic.bulk_update_abacus_states_by_parent_table(parent_table_name, params) assert result.status == 200 mock_payment_account_state_logic.assert_called_with(parent_table_name, params) def test_validate_create_params_invalid_parent_table(): """Test error is raised when an invalid parent_table_name is provided.""" parent_table_name = 'parent_table' with pytest.raises(ValidationError): logic._validate_create_params(parent_table_name, 123) @patch('abacus_state.logic.abacus_state.AbacusState') def test_validate_create_params_states_already_exist(mock_model): """Test error is raised when parent table already has abacus state records.""" abacus_state = 'existing record' mock_model.get_action_status_list.return_value = [abacus_state] with pytest.raises(ValidationError): logic._validate_create_params('accounting_period', 123) @patch('abacus_state.logic.abacus_state.AbacusState') def test_validate_create_params_valid(mock_model): """Test no error is raised when create params are valid.""" mock_model.get_action_status_list.return_value = [] res = logic._validate_create_params('accounting_period', 123) assert res is None mock_model.get_action_status_list.assert_called_once_with('accounting_period', 123) @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_dataload_states_by_ids(mock_query): """Test dataload_states_by_ids function.""" state = AbacusStateFactory() state_ids = [state.abacus_state_id] mock_query.return_value = MagicMock(all=MagicMock(return_value=[state])) res = logic.dataload_states_by_ids(state_ids) assert res.message == {'items': [{'data': AbacusStateDetailSchema().dump(state)}]} mock_query.assert_called_once_with(state_ids=state_ids) @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_no_filters(mock_query): """Test bulk_query_states with no filters.""" state = AbacusStateFactory() mock_count_query = MagicMock(count=MagicMock(return_value=1)) mock_all_query = MagicMock(all=MagicMock(return_value=[state])) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states() assert res.status == 200 assert res.message['total'] == 1 assert res.message['limit'] == 100 assert res.message['offset'] == 0 assert len(res.message['items']) == 1 @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_with_parent_table_name(mock_query): """Test bulk_query_states filtering by parent_table_name.""" state = AbacusStateFactory(parent_table_name='accounting_period') mock_count_query = MagicMock(count=MagicMock(return_value=1)) mock_all_query = MagicMock(all=MagicMock(return_value=[state])) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states(parent_table_name='accounting_period') assert res.status == 200 assert res.message['total'] == 1 mock_query.assert_any_call( parent_table_name='accounting_period', parent_table_ids=None, action_name=None, ) @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_with_parent_table_ids(mock_query): """Test bulk_query_states filtering by parent_table_ids.""" state = AbacusStateFactory(parent_table_name='accounting_period', parent_table_id=1) mock_count_query = MagicMock(count=MagicMock(return_value=1)) mock_all_query = MagicMock(all=MagicMock(return_value=[state])) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states( parent_table_name='accounting_period', parent_table_ids=[1, 2] ) assert res.status == 200 assert res.message['total'] == 1 mock_query.assert_any_call( parent_table_name='accounting_period', parent_table_ids=[1, 2], action_name=None, ) @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_with_action_name(mock_query): """Test bulk_query_states filtering by action_name.""" state = AbacusStateFactory(action_name='deliver_sales_files') mock_count_query = MagicMock(count=MagicMock(return_value=1)) mock_all_query = MagicMock(all=MagicMock(return_value=[state])) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states(action_name='deliver_sales_files') assert res.status == 200 assert res.message['total'] == 1 mock_query.assert_any_call( parent_table_name=None, parent_table_ids=None, action_name='deliver_sales_files', ) @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_limit_and_offset(mock_query): """Test bulk_query_states with limit and offset.""" states = [ AbacusStateFactory(parent_table_id=parent_id) for parent_id in range(1, 6) ] mock_count_query = MagicMock(count=MagicMock(return_value=10)) mock_all_query = MagicMock(all=MagicMock(return_value=states)) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states(limit=5, offset=3) assert res.status == 200 assert res.message['total'] == 10 assert res.message['limit'] == 5 assert res.message['offset'] == 3 mock_query.assert_any_call( parent_table_name=None, parent_table_ids=None, action_name=None, limit=5, offset=3, ) def test_bulk_query_states_parent_table_ids_without_parent_table_name(): """Test bulk_query_states raises error when parent_table_ids without parent_table_name.""" with pytest.raises(ValidationError): logic.bulk_query_states(parent_table_ids=[1, 2]) def test_bulk_query_states_invalid_parent_table_name(): """Test bulk_query_states raises error for invalid parent_table_name.""" with pytest.raises(ValidationError): logic.bulk_query_states(parent_table_name='invalid_table_name') @patch('abacus_state.models.abacus_state.AbacusState.get_filtered_query') def test_bulk_query_states_response_structure(mock_query): """Test bulk_query_states returns correct response structure.""" state = AbacusStateFactory() mock_count_query = MagicMock(count=MagicMock(return_value=1)) mock_all_query = MagicMock(all=MagicMock(return_value=[state])) mock_query.side_effect = [mock_count_query, mock_all_query] res = logic.bulk_query_states() assert res.status == 200 assert 'items' in res.message assert 'total' in res.message assert 'limit' in res.message assert 'offset' in res.message assert isinstance(res.message['items'], list) @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.StatementPeriodAdjustmentFile') @patch('abacus_state.logic.abacus_state.AbacusState') @patch('abacus_state.logic.abacus_state.is_abacus_flowthrough_automation_enabled') def test_create_abacus_states_for_auto_adjustment_file_success( mock_feature_enabled, mock_state_model, mock_file_model, mock_validate ): """Test successful creation of all abacus states for auto generated adjustments.""" parent_table_name = constants.PARENT_TABLE_NAMES.STATEMENT_PERIOD_ADJUSTMENT_FILE parent_table_id = 123 mock_file_model.get_statement_period_adjustment_file.return_value = { 'batch_type': 'auto' } mock_validate.return_value = None mock_state_model.build.return_value = None mock_state_model.commit_changes.return_value = None mock_feature_enabled.return_value = True res = logic.create_abacus_states_by_parent_table( parent_table_name, parent_table_id, None ) assert res.status == 201 assert len(res.message) == len( constants.AUTO_GENERATED_STATEMENT_PERIOD_ADJUSTMENT_FILE_ACTION_NAMES ) mock_validate.assert_called_once_with(parent_table_name, parent_table_id) assert mock_state_model.build.call_count == len( constants.AUTO_GENERATED_STATEMENT_PERIOD_ADJUSTMENT_FILE_ACTION_NAMES ) mock_state_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state._validate_create_params') @patch('abacus_state.logic.abacus_state.StatementPeriodAdjustmentFile') @patch('abacus_state.logic.abacus_state.AbacusState') @patch('abacus_state.logic.abacus_state.is_abacus_flowthrough_automation_enabled') def test_create_abacus_states_for_adjustment_file_success( mock_feature_enabled, mock_state_model, mock_file_model, mock_validate ): """Test successful creation of all abacus states for uploaded adjustments file.""" parent_table_name = constants.PARENT_TABLE_NAMES.STATEMENT_PERIOD_ADJUSTMENT_FILE parent_table_id = 123 mock_file_model.get_statement_period_adjustment_file.return_value = { 'batch_type': 'upload' } mock_validate.return_value = None mock_state_model.build.return_value = None mock_state_model.commit_changes.return_value = None mock_feature_enabled.return_value = True res = logic.create_abacus_states_by_parent_table( parent_table_name, parent_table_id, None ) assert res.status == 201 assert len(res.message) == len( constants.STATEMENT_PERIOD_ADJUSTMENT_FILE_ACTION_NAMES ) mock_validate.assert_called_once_with(parent_table_name, parent_table_id) assert mock_state_model.build.call_count == len( constants.STATEMENT_PERIOD_ADJUSTMENT_FILE_ACTION_NAMES ) mock_state_model.commit_changes.assert_called_once() @patch('abacus_state.logic.abacus_state.AbacusState') def test_reset_abacus_state_actions(mock_model): """Test to reset all abacus states for specified table and id.""" parent_table_name = 'statement_period_adjustment_file' parent_table_id = 9999 mock_abacus_states = [ AbacusStateFactory.create( parent_table_name=parent_table_name, parent_table_id=parent_table_id, action_name='upload_file', action_status=constants.ACTION_STATUSES.INIT, ) ] mock_model.reset_abacus_state_actions.return_value = mock_abacus_states result = logic.reset_abacus_state_actions(parent_table_name, parent_table_id) assert result.status == 200 mock_model.reset_abacus_state_actions.assert_called_once_with( parent_table_name, parent_table_id ) @patch('abacus_state.logic.abacus_state.g') @patch('abacus_state.logic.abacus_state.schema') @patch('abacus_state.logic.abacus_state.AbacusState') def test_dataload_states_by_target_groups_states_per_id( mock_model, mock_schema, mock_g ): """dataload_states_by_target returns one ordered entry per requested parent id.""" mock_model.get_filtered_query.return_value.all.return_value = ['s1', 's2', 's3'] mock_schema.dump.return_value = [ {'parent_table_id': 1, 'action_name': 'a'}, {'parent_table_id': 1, 'action_name': 'b'}, {'parent_table_id': 2, 'action_name': 'c'}, ] result = logic.dataload_states_by_target( 'statement_period_adjustment_file', [1, 2, 3] ) assert result.message == { 'items': [ { 'data': [ {'parent_table_id': 1, 'action_name': 'a'}, {'parent_table_id': 1, 'action_name': 'b'}, ] }, {'data': [{'parent_table_id': 2, 'action_name': 'c'}]}, {'data': None}, ] } mock_g.log.error.assert_not_called() @patch('abacus_state.logic.abacus_state.g') @patch('abacus_state.logic.abacus_state.schema') @patch('abacus_state.logic.abacus_state.AbacusState') def test_dataload_states_by_target_logs_when_records_lack_parent_table_id( mock_model, mock_schema, mock_g ): """A non-empty fetch whose records carry no parent_table_id is logged loudly.""" mock_model.get_filtered_query.return_value.all.return_value = ['state'] # Simulate a serializer field mismatch: records present, none carry the key. mock_schema.dump.return_value = [{'action_name': 'x'}] mock_g.log.error = MagicMock() logic.dataload_states_by_target('statement_period_adjustment_file', [1]) mock_g.log.error.assert_called_once()