"""Unit testcases for EarningsTransfer logic.""" from datetime import datetime from unittest.mock import MagicMock, call, patch import pytest from abacus_common_logic.constants.constants import SYSTEM_TIMEZONE from freezegun import freeze_time from marshmallow import ValidationError from abacus_contract.tests.utils.factories import ContractFactory from royalties.constants.constants import ( EARNINGS_TRANSFER_SORT_OPTIONS, SORT_ORDER_OPTIONS, ) from royalties.logic import earnings_transfer as logic from royalties.logic.earnings_transfer import validate_update_entries from royalties.models.earnings_transfer import EarningsTransfer from royalties.schemas.earnings_transfer import EarningsTransferDetailSchema from royalties.tests.utils.factories import EarningsTransferFactory @patch('royalties.logic.earnings_transfer.EarningsTransfer') def test_get_earnings_transfer(mock_earnings_transfer_model): """Test to get earnings transfer.""" mock_earnings_transfer = EarningsTransferFactory.create() mock_earnings_transfer_model.get_earnings_transfers.return_value = ( [mock_earnings_transfer], 1, ) mock_request_params = { 'limit': 10, 'offset': 0, 'sort_by': EARNINGS_TRANSFER_SORT_OPTIONS.EARNINGS_TRANSFER_ID, 'sort_order': SORT_ORDER_OPTIONS.DESC, } result = logic.get_earnings_transfers(mock_request_params) assert result.status == 200 assert result.message['total_count'] == 1 mock_earnings_transfer_model.get_earnings_transfers.assert_called_once_with( **mock_request_params ) @patch('royalties.logic.earnings_transfer.EarningsTransfer') @patch('royalties.logic.earnings_transfer.EarningsTransferFilterSchema.load') def test_get_earnings_transfers_validation_error( mock_load, mock_earnings_transfer_model ): """Test handling of invalid query parameters.""" mock_load.side_effect = ValidationError('Invalid limit value') params = {'limit': 'not_a_number'} result = logic.get_earnings_transfers(params) assert result.status == 400 assert result.errors == { 'code': 'error', 'message': 'Invalid limit value', } mock_earnings_transfer_model.get_earnings_transfers.assert_not_called() @patch('royalties.logic.earnings_transfer.EarningsTransfer') def test_get_earnings_transfers_by_contract_id(mock_earnings_transfer_model): """Test to get earnings transfers by contract id.""" mock_earnings_transfer = EarningsTransferFactory.create() mock_earnings_transfer_model.get_earnings_transfers_by_contract_id.return_value = ( [mock_earnings_transfer], 1, ) mock_request_params = { 'limit': 10, 'offset': 0, 'sort_by': EARNINGS_TRANSFER_SORT_OPTIONS.EARNINGS_TRANSFER_ID, 'sort_order': SORT_ORDER_OPTIONS.DESC, } contract_id = mock_earnings_transfer.from_contract_id result = logic.get_earnings_transfers_by_contract_id( contract_id, None, mock_request_params ) assert result.status == 200 assert result.message['total_count'] == 1 mock_request_params.update({'contract_id': contract_id, 'transfer_type': None}) mock_earnings_transfer_model.get_earnings_transfers_by_contract_id.assert_called_once_with( **mock_request_params ) @patch('royalties.logic.earnings_transfer.EarningsTransfer') @patch('royalties.logic.earnings_transfer.EarningsTransferFilterSchema.load') def test_get_contract_earnings_transfers_invalid_query_param( mock_load, mock_earnings_transfer_model ): """Test handling of invalid query parameters.""" mock_earnings_transfer = EarningsTransferFactory.create() contract_id = mock_earnings_transfer.from_contract_id mock_load.side_effect = ValidationError('Invalid limit value') mock_request_params = {'limit': 'not_a_number'} result = logic.get_earnings_transfers_by_contract_id( contract_id, None, mock_request_params ) assert result.status == 400 assert result.errors == { 'code': 'error', 'message': 'Invalid limit value', } mock_earnings_transfer_model.get_earnings_transfers_by_contract_id.assert_not_called() @patch('royalties.logic.earnings_transfer.EarningsTransfer') def test_get_contract_earnings_transfers_invalid_transfer_type( mock_earnings_transfer_model, ): """Test handling of invalid transfer_type.""" mock_earnings_transfer = EarningsTransferFactory.create() contract_id = mock_earnings_transfer.from_contract_id mock_request_params = { 'limit': 10, 'offset': 0, 'sort_by': EARNINGS_TRANSFER_SORT_OPTIONS.EARNINGS_TRANSFER_ID, 'sort_order': SORT_ORDER_OPTIONS.DESC, } result = logic.get_earnings_transfers_by_contract_id( contract_id, 'test', mock_request_params ) assert result.status == 400 assert result.errors == { 'code': 'error', 'message': 'Invalid transfer_type. Must be one of: reclass, override, or transfer.', } mock_earnings_transfer_model.get_earnings_transfers_by_contract_id.assert_not_called() @patch('royalties.logic.earnings_transfer.EarningsTransfer') def test_get_earnings_transfer_by_id(mock_earnings_transfer_model): """Test getting an earnings transfer by ID.""" mock_earnings_transfer = EarningsTransferFactory.create() mock_earnings_transfer_model.get_by_id_or_error.return_value = ( mock_earnings_transfer ) transfer_id = mock_earnings_transfer.earnings_transfer_id result = logic.get_earnings_transfer_by_id(transfer_id) mock_earnings_transfer_model.get_by_id_or_error.assert_called_once_with( transfer_id, 404 ) assert result.status == 200 assert result.message == EarningsTransferDetailSchema().dump(mock_earnings_transfer) def test_validate_update_entries_when_valid(): """Test that no errors are raised when updating entries.""" existing_transfer_1 = EarningsTransferFactory.create() validate_update_entries( {1: existing_transfer_1}, [{'comment': 'new_comment', 'earnings_transfer_id': 1}], [1], ) def test_validate_update_entries_when_one_invalid(): """Test validate update entries when one invalid.""" existing_transfer_1 = EarningsTransferFactory.create() with pytest.raises(ValidationError) as exc_info: validate_update_entries( {1: existing_transfer_1}, [ { 'to_contract_id': existing_transfer_1.from_contract.contract_id, 'earnings_transfer_id': 1, } ], [1], ) assert ( str(exc_info.value) == 'Source and destination contracts must be different for ID [1].' ) def test_validate_update_entries_when_many_invalid(): """Test validate update entries when many invalid.""" existing_transfer_1 = EarningsTransferFactory.create() existing_transfer_2 = EarningsTransferFactory.create() id_1 = existing_transfer_1.earnings_transfer_id id_2 = existing_transfer_2.earnings_transfer_id with pytest.raises(ValidationError) as exc_info: validate_update_entries( {id_1: existing_transfer_1, id_2: existing_transfer_2}, [ { 'to_contract_id': existing_transfer_1.from_contract.contract_id, 'earnings_transfer_id': id_1, }, { 'to_contract_id': existing_transfer_2.from_contract.contract_id, 'earnings_transfer_id': id_2, }, ], [id_1, id_2], ) assert ( str(exc_info.value) == f'Source and destination contracts must be different for ID [{id_1}, {id_2}].' ) def test_validate_update_entries_when_many_invalid_transfer_amount(): """Test validate update entries when many with invalid transfer_amount.""" existing_transfer_1 = EarningsTransferFactory.create() existing_transfer_2 = EarningsTransferFactory.create() id_1 = existing_transfer_1.earnings_transfer_id id_2 = existing_transfer_2.earnings_transfer_id with pytest.raises(ValidationError) as exc_info: validate_update_entries( {id_1: existing_transfer_1, id_2: existing_transfer_2}, [ {'transfer_amount': 108, 'earnings_transfer_id': id_1}, {'transfer_amount': 108, 'earnings_transfer_id': id_2}, ], [id_1, id_2], ) assert ( str(exc_info.value) == "[{'transfer_amount': ['Transfer amount must be between 0 and 100 when rate type is percent.']}, " "{'transfer_amount': ['Transfer amount must be between 0 and 100 when rate type is percent.']}]" ) @patch('royalties.logic.earnings_transfer.Contract') def test_validate_all_contracts_exists(mock_contract_model): """Test that all contracts exist.""" mock_contracts = ContractFactory.create_batch(2) contract_ids = set([mock_contract.contract_id for mock_contract in mock_contracts]) mock_contract_model.query.filter.return_value.all.return_value = mock_contracts result = logic.validate_contracts(contract_ids) result is True @patch('royalties.logic.earnings_transfer.Contract') def test_validate_contracts_raise_validation_error(mock_contract_model): """Test that validation error is raised if contract doesn't exist.""" mock_contracts = ContractFactory.create_batch(2) contract_ids = set([mock_contract.contract_id for mock_contract in mock_contracts]) mock_contract_model.query.filter.return_value.all.return_value = mock_contracts contract_ids.add(99999) contract_ids.add(9010101) with pytest.raises(ValidationError) as exc_info: logic.validate_contracts(contract_ids) assert ( str(exc_info.value) == 'The following contract IDs do not exist: [9010101, 99999]' ) @pytest.fixture @freeze_time(datetime(2026, 5, 27, 0, 0, 0, tzinfo=SYSTEM_TIMEZONE)) @patch('royalties.logic.earnings_transfer.db') @patch('royalties.logic.earnings_transfer.validate_from_contract') @patch('royalties.logic.earnings_transfer.validate_contracts') @patch('royalties.logic.earnings_transfer.validate_post_request_contains_unique_data') @freeze_time(datetime(2026, 5, 5, 0, 0, 0, tzinfo=SYSTEM_TIMEZONE)) def test_bulk_create_earnings_transfers_success( mock_duplicate_validation, mock_validate_contracts, mock_validate_from_contract, mock_db, ): """Test should create multiple EarningsTransfers and commit all at once.""" mock_from_contract = ContractFactory.create() mock_to_contract_1 = ContractFactory.create() mock_to_contract_2 = ContractFactory.create() mock_input = [ { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract_1.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'transfer_source': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', }, { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract_2.contract_id, 'transfer_type': 'override', 'rate_type': 'flat_rate', 'transfer_amount': '1890.89', 'transfer_source': 'close_balance', 'negative': True, 'active': True, 'comment': 'Test comment', }, ] mock_duplicate_validation.return_value = True result = logic.bulk_create_earnings_transfers(mock_input) mock_duplicate_validation.assert_called_once_with(mock_input) mock_validate_contracts.assert_called_once_with( { mock_from_contract.contract_id, mock_to_contract_1.contract_id, mock_to_contract_2.contract_id, } ) mock_validate_from_contract.assert_called_once() mock_db.session.add_all.assert_called_once() mock_db.session.commit.assert_called_once() assert result.status == 201 assert result.message == [ { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract_1.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'input': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', 'earnings_transfer_id': None, 'created_at': '2026-05-05', }, { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract_2.contract_id, 'transfer_type': 'override', 'rate_type': 'flat_rate', 'transfer_amount': '1890.89', 'input': 'close_balance', 'negative': True, 'active': True, 'comment': 'Test comment', 'earnings_transfer_id': None, 'created_at': '2026-05-05', }, ] @patch('royalties.logic.earnings_transfer.db') @patch('royalties.logic.earnings_transfer.validate_contracts') @patch('royalties.logic.earnings_transfer.validate_post_request_contains_unique_data') def test_bulk_create_earnings_transfers_validation_error( mock_duplicate_validation, mock_validate_contracts, mock_db, ): """Test that error is returned if contract doesn't exist.""" mock_from_contract = ContractFactory.create() mock_to_contract = ContractFactory.create() mock_input = [ { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'transfer_source': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', } ] mock_duplicate_validation.return_value = True mock_validate_contracts.side_effect = ValidationError("Contract doesn't exist") result = logic.bulk_create_earnings_transfers(mock_input) mock_duplicate_validation.assert_called_once_with(mock_input) mock_validate_contracts.assert_called_once_with( {mock_from_contract.contract_id, mock_to_contract.contract_id} ) assert result.status == 400 assert result.errors == {'code': 'error', 'message': "Contract doesn't exist"} mock_db.session.add_all.assert_not_called() mock_db.session.commit.assert_not_called() @patch('royalties.logic.earnings_transfer.db') @patch('royalties.logic.earnings_transfer.validate_contracts') @patch('royalties.logic.earnings_transfer.validate_post_request_contains_unique_data') def test_bulk_create_earnings_transfers_duplicate_record( mock_duplicate_validation, mock_validate_contracts, mock_db, ): """Test that error is returned if post request contains duplicate records.""" mock_from_contract = ContractFactory.create() mock_to_contract = ContractFactory.create() mock_input = [ { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'transfer_source': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', }, { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'transfer_source': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', }, ] mock_duplicate_validation.side_effect = ValidationError( 'POST request contain duplicate records.' ) result = logic.bulk_create_earnings_transfers(mock_input) mock_duplicate_validation.assert_called_once_with(mock_input) mock_validate_contracts.assert_not_called() assert result.status == 400 assert result.errors == { 'code': 'error', 'message': 'POST request contain duplicate records.', } mock_db.session.add_all.assert_not_called() mock_db.session.commit.assert_not_called() @patch('royalties.logic.earnings_transfer.db') @patch('royalties.logic.earnings_transfer.validate_from_contract') @patch('royalties.logic.earnings_transfer.validate_contracts') @patch('royalties.logic.earnings_transfer.validate_post_request_contains_unique_data') def test_bulk_create_earnings_transfers_commit_fails( mock_duplicate_validation, mock_validate_contracts, mock_validate_from_contract, mock_db, ): """Test that rollback triggers if commit fails.""" mock_from_contract = ContractFactory.create() mock_to_contract = ContractFactory.create() mock_input = [ { 'from_contract_id': mock_from_contract.contract_id, 'to_contract_id': mock_to_contract.contract_id, 'transfer_type': 'override', 'rate_type': 'percent', 'transfer_amount': '1', 'transfer_source': 'gross_revenue', 'negative': True, 'active': True, 'comment': 'Test comment', } ] mock_duplicate_validation.return_value = True mock_validate_contracts.return_value = True mock_db.session.commit.side_effect = Exception('DB commit failed') with pytest.raises(Exception, match='DB commit failed'): logic.bulk_create_earnings_transfers(mock_input) mock_validate_contracts.assert_called_once_with( {mock_from_contract.contract_id, mock_to_contract.contract_id} ) mock_duplicate_validation.assert_called_once_with(mock_input) mock_db.session.rollback.assert_called_once() mock_db.session.commit.assert_called_once() def test_validate_from_contract(): """Test validating entries based on the `from_contract_id` field.""" from_contract = ContractFactory.create() to_contract_1 = ContractFactory.create() to_contract_2 = ContractFactory.create() transfer_1 = EarningsTransfer( from_contract_id=from_contract.contract_id, to_contract_id=to_contract_1.contract_id, transfer_type='override', rate_type='percent', transfer_amount='60', input='close_balance', ) transfer_2 = EarningsTransfer( from_contract_id=from_contract.contract_id, to_contract_id=to_contract_2.contract_id, transfer_type='override', rate_type='percent', transfer_amount='60', input='gross_revenue', ) result = logic.validate_from_contract([transfer_1, transfer_2]) assert result is None transfer_2.to_contract_id = to_contract_1.contract_id with pytest.raises(ValidationError) as exc_info: logic.validate_from_contract([transfer_1, transfer_2]) assert ( str(exc_info.value) == 'Contract 1 cannot have more than one transfer of earnings configured for Contract 2' ) transfer_2.to_contract_id = to_contract_2.contract_id transfer_2.input = 'close_balance' with pytest.raises(ValidationError) as exc_info: logic.validate_from_contract([transfer_1, transfer_2]) assert ( str(exc_info.value) == 'Contract 1 cannot transfer more than 100% of its earnings' )