"""Unit tests for PaymentAllocationProcessor.""" import logging from decimal import Decimal from unittest.mock import Mock import pytest from src.errors import ( BalancesNotClosedError, StatementPeriodPaymentEntityNotFoundError, ) from src.processor import PaymentAllocationProcessor from src.schemas import ( ContractAdjustmentCount, LedgerAdjustmentApplied, PaymentAllocationResponse, StatementPeriodPaymentEntity, ) def _make_sppe( sppe_id: int = 1, statement_period_id: int = 456, reference_payment_entity_id: int = 10, ) -> StatementPeriodPaymentEntity: """Create a StatementPeriodPaymentEntity for testing.""" return StatementPeriodPaymentEntity( statement_period_payment_entity_id=sppe_id, statement_period_id=statement_period_id, reference_payment_entity_id=reference_payment_entity_id, ) def _make_adjustment( laa_id: int = 1, account_id: int = 100, contract_id: int = 200, amount: str = '50.00', currency: str = 'USD', payee_currency_amount: str = '50.00', payee_currency: str = 'USD', account_payee_id: int = 300, ) -> LedgerAdjustmentApplied: """Create an adjustment record for testing.""" return LedgerAdjustmentApplied( ledger_adjustment_applied_id=laa_id, account_id=account_id, contract_id=contract_id, adjustment_amount=Decimal(amount), adjustment_currency_code=currency, adjustment_amount_payee_currency=Decimal(payee_currency_amount), adjustment_payee_currency_code=payee_currency, account_payee_id=account_payee_id, ) def _make_mock_repository(**sppe_kwargs) -> Mock: """Create a mock repository with sensible defaults for process() tests. Returns a Mock with all repository methods pre-configured for the happy-path: SPPE found, balances closed, no packet clamping, no contracts to process. Callers override specific attributes as needed. """ mock = Mock() mock.get_statement_period_payment_entity.return_value = _make_sppe(**sppe_kwargs) mock.get_state_status.return_value = 'complete' mock.get_max_allowed_packet.return_value = None mock.get_contract_adjustment_counts.return_value = [] mock.create_payment_allocation.return_value = 789 mock.create_payment_allocation_ledger_adjustment.return_value = 1 mock.conn = Mock() return mock class TestPaymentAllocationProcessorInit: """Tests for PaymentAllocationProcessor initialization.""" def test_init_stores_repository(self): """Test processor stores repository reference.""" repository = Mock() processor = PaymentAllocationProcessor(repository) assert processor._repository is repository class TestPaymentAllocationProcessorProcess: """Tests for PaymentAllocationProcessor.process method.""" def test_process_success_creates_allocations(self): """Test successful allocation creation.""" mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ContractAdjustmentCount(contract_id=200, adjustment_count=2)], [], ] mock_repository.get_adjustments_for_contracts.side_effect = [ [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=200, account_payee_id=300), ], [], ] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert isinstance(result, PaymentAllocationResponse) assert result.allocations_created == 1 # 2 adjustments grouped into 1 assert result.ledger_adjustments_linked == 2 assert result.statement_period_id == 456 assert result.statement_period_payment_entity_id == 1 # Verify allocation was created with summed amounts mock_repository.create_payment_allocation.assert_called_once() call_kwargs = mock_repository.create_payment_allocation.call_args[1] assert call_kwargs['contract_id'] == 200 assert call_kwargs['payee_type'] == 'account_payee' assert call_kwargs['payee_id'] == 300 assert call_kwargs['statement_period_id'] == 456 assert call_kwargs['amount_to_payment'] == Decimal('100.00') assert call_kwargs['amount_to_ledger'] == Decimal('100.00') assert call_kwargs['currency_code'] == 'USD' assert call_kwargs['payment_allocation_type'] == 'flowthrough' # Verify correct ledger adjustment IDs were linked link_calls = ( mock_repository.create_payment_allocation_ledger_adjustment.call_args_list ) linked_laa_ids = {call.args[1] for call in link_calls} assert linked_laa_ids == {1, 2} # Verify batch commit mock_repository.conn.commit.assert_called_once() def test_process_multiple_contracts_multiple_allocations(self): """Test multiple contracts create separate allocations.""" mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ ContractAdjustmentCount(contract_id=200, adjustment_count=1), ContractAdjustmentCount(contract_id=201, adjustment_count=1), ], [], ] mock_repository.get_adjustments_for_contracts.side_effect = [ [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=201, account_payee_id=301), ], [], ] mock_repository.create_payment_allocation.side_effect = [789, 790] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 2 assert result.ledger_adjustments_linked == 2 assert mock_repository.create_payment_allocation.call_count == 2 mock_repository.conn.commit.assert_called_once() def test_process_groups_by_currency(self): """Test adjustments with different currencies create separate allocations.""" mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ContractAdjustmentCount(contract_id=200, adjustment_count=2)], [], ] mock_repository.get_adjustments_for_contracts.side_effect = [ [ _make_adjustment( laa_id=1, contract_id=200, account_payee_id=300, payee_currency='USD', ), _make_adjustment( laa_id=2, contract_id=200, account_payee_id=300, payee_currency='EUR', ), ], [], ] mock_repository.create_payment_allocation.side_effect = [789, 790] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 2 assert result.ledger_adjustments_linked == 2 def test_process_cross_currency_sums_correct_amounts(self): """Test amount_to_payment sums payee currency, amount_to_ledger sums adjustment currency.""" mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ContractAdjustmentCount(contract_id=200, adjustment_count=2)], [], ] mock_repository.get_adjustments_for_contracts.side_effect = [ [ _make_adjustment( laa_id=1, contract_id=200, account_payee_id=300, amount='100.00', currency='EUR', payee_currency_amount='110.50', payee_currency='USD', ), _make_adjustment( laa_id=2, contract_id=200, account_payee_id=300, amount='200.00', currency='EUR', payee_currency_amount='221.00', payee_currency='USD', ), ], [], ] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 1 call_kwargs = mock_repository.create_payment_allocation.call_args[1] # amount_to_payment sums adjustment_amount_payee_currency assert call_kwargs['amount_to_payment'] == Decimal('331.50') # amount_to_ledger sums adjustment_amount assert call_kwargs['amount_to_ledger'] == Decimal('300.00') # currency_code uses adjustment currency, not payee currency assert call_kwargs['currency_code'] == 'EUR' def test_process_multiple_discovery_rounds(self): """Test processing across multiple outer discovery rounds.""" mock_repository = _make_mock_repository() # Two discovery rounds, then empty mock_repository.get_contract_adjustment_counts.side_effect = [ [ContractAdjustmentCount(contract_id=200, adjustment_count=1)], [ContractAdjustmentCount(contract_id=201, adjustment_count=1)], [], ] mock_repository.get_adjustments_for_contracts.side_effect = [ [_make_adjustment(laa_id=1, contract_id=200, account_payee_id=300)], [], [_make_adjustment(laa_id=2, contract_id=201, account_payee_id=301)], [], ] mock_repository.create_payment_allocation.side_effect = [789, 790] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 2 assert result.ledger_adjustments_linked == 2 assert mock_repository.conn.commit.call_count == 2 assert mock_repository.get_contract_adjustment_counts.call_count == 3 assert mock_repository.get_adjustments_for_contracts.call_count == 4 def test_process_no_adjustments_returns_zero(self): """Test process returns zero counts when no contracts found.""" mock_repository = _make_mock_repository() processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 0 assert result.ledger_adjustments_linked == 0 mock_repository.create_payment_allocation.assert_not_called() mock_repository.get_adjustments_for_contracts.assert_not_called() def test_process_sppe_not_found(self): """Test process raises error when SPPE not found in DB.""" mock_repository = Mock() mock_repository.get_statement_period_payment_entity.return_value = None processor = PaymentAllocationProcessor(mock_repository) with pytest.raises(StatementPeriodPaymentEntityNotFoundError) as exc_info: processor.process(1) assert 'statement_period_payment_entity_id=1' in str(exc_info.value) def test_process_balances_not_closed(self): """Test process raises error when close_balance is not complete.""" mock_repository = _make_mock_repository() mock_repository.get_state_status.return_value = 'init' processor = PaymentAllocationProcessor(mock_repository) with pytest.raises(BalancesNotClosedError) as exc_info: processor.process(1) assert 'Close balance is not complete' in str(exc_info.value) def test_process_no_close_balance_state(self): """Test process raises error when close_balance state is missing.""" mock_repository = _make_mock_repository() mock_repository.get_state_status.return_value = None processor = PaymentAllocationProcessor(mock_repository) with pytest.raises(BalancesNotClosedError): processor.process(1) def test_process_derives_statement_period_id_from_db(self): """Test that statement_period_id is derived from the SPPE DB record.""" mock_repository = _make_mock_repository( sppe_id=5, statement_period_id=999, reference_payment_entity_id=20 ) processor = PaymentAllocationProcessor(mock_repository) result = processor.process(5) assert result.statement_period_id == 999 assert result.statement_period_payment_entity_id == 5 mock_repository.get_contract_adjustment_counts.assert_called_once_with( 999, 20, 250_000 ) class TestPaymentAllocationProcessorGroupAdjustments: """Tests for PaymentAllocationProcessor._group_adjustments method.""" def test_group_single_adjustment(self): """Test grouping a single adjustment.""" adjustments = [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 1 key = (200, 300, 'USD', 'USD') assert key in result assert len(result[key]) == 1 def test_group_same_contract_payee_currency(self): """Test adjustments with same key are grouped together.""" adjustments = [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=3, contract_id=200, account_payee_id=300), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 1 key = (200, 300, 'USD', 'USD') assert len(result[key]) == 3 def test_group_different_contracts(self): """Test adjustments with different contracts are grouped separately.""" adjustments = [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=201, account_payee_id=301), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 2 def test_group_different_payee_ids(self): """Test adjustments with different payee IDs are grouped separately.""" adjustments = [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=200, account_payee_id=301), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 2 def test_group_different_payee_currencies(self): """Test adjustments with different payee currencies are grouped separately.""" adjustments = [ _make_adjustment( laa_id=1, contract_id=200, account_payee_id=300, payee_currency='USD' ), _make_adjustment( laa_id=2, contract_id=200, account_payee_id=300, payee_currency='EUR' ), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 2 def test_group_different_adjustment_currencies(self): """Test adjustments with different source currencies are grouped separately.""" adjustments = [ _make_adjustment( laa_id=1, contract_id=200, account_payee_id=300, currency='USD' ), _make_adjustment( laa_id=2, contract_id=200, account_payee_id=300, currency='GBP' ), ] result = PaymentAllocationProcessor._group_adjustments(adjustments) assert len(result) == 2 class TestProcessMixedBatches: """Tests for mixed regular and oversized contracts in a single round.""" def test_process_regular_and_oversized_in_same_round(self): """Test counts accumulate across oversized and remainder+regular batches. Contract counts: (200, 1), (201, 15_000), (202, 2) with batch_size=10_000. build_batches produces: [[(201, 10000)], [(201, 5000), (202, 2), (200, 1)]]. Batch 1 processes the oversized portion of 201. Batch 2 processes the remainder of 201 packed with regular contracts 202 and 200. """ mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ ContractAdjustmentCount(contract_id=200, adjustment_count=1), ContractAdjustmentCount(contract_id=201, adjustment_count=15_000), ContractAdjustmentCount(contract_id=202, adjustment_count=2), ], [], ] # Batch 1: oversized [201] → 1 adj, then empty # Batch 2: remainder+regular [201, 202, 200] → 3 adjs, then empty mock_repository.get_adjustments_for_contracts.side_effect = [ [_make_adjustment(laa_id=4, contract_id=201, account_payee_id=301)], [], [ _make_adjustment(laa_id=1, contract_id=200, account_payee_id=300), _make_adjustment(laa_id=2, contract_id=202, account_payee_id=302), _make_adjustment(laa_id=3, contract_id=202, account_payee_id=302), ], [], ] mock_repository.create_payment_allocation.side_effect = [789, 790, 791] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) # 1 from oversized batch (201) + 2 from remainder+regular (200, 202) assert result.allocations_created == 3 assert result.ledger_adjustments_linked == 4 # 1 commit per non-empty query result assert mock_repository.conn.commit.call_count == 2 # All calls use limit=batch_size for call in mock_repository.get_adjustments_for_contracts.call_args_list: assert call.kwargs.get('limit') == 10_000 class TestProcessOversizedContract: """Tests for oversized contract sub-batching in process flow.""" def test_process_oversized_contract_sub_batches(self): """Test oversized contract is split into full portion + remainder batches. Contract (200, 15_000) with batch_size=10_000: build_batches produces: [[(200, 10000)], [(200, 5000)]]. Batch 1 (full portion): LIMIT loop processes first 10_000 adjustments. Batch 2 (remainder): LIMIT loop processes remaining 5_000 adjustments. """ mock_repository = _make_mock_repository() mock_repository.get_contract_adjustment_counts.side_effect = [ [ContractAdjustmentCount(contract_id=200, adjustment_count=15_000)], [], ] # Batch 1 (full portion): two rounds of adjustments, then empty # Batch 2 (remainder): empty (all adjustments consumed by batch 1) mock_repository.get_adjustments_for_contracts.side_effect = [ [ _make_adjustment(laa_id=i, contract_id=200, account_payee_id=300) for i in range(1, 4) ], [ _make_adjustment(laa_id=i, contract_id=200, account_payee_id=300) for i in range(4, 6) ], [], [], ] mock_repository.create_payment_allocation.side_effect = [789, 790] processor = PaymentAllocationProcessor(mock_repository) result = processor.process(1) assert result.allocations_created == 2 assert result.ledger_adjustments_linked == 5 assert mock_repository.conn.commit.call_count == 2 # All calls target contract 200 with limit=batch_size for call in mock_repository.get_adjustments_for_contracts.call_args_list: assert call.args[2] == [200] assert call.kwargs.get('limit') == 10_000 class TestGetEffectiveBatchSize: """Tests for PaymentAllocationProcessor._get_effective_batch_size.""" def test_returns_config_when_max_packet_is_none(self): """Test uses config.batch_size when max_allowed_packet is undetermined.""" mock_repository = Mock() mock_repository.get_max_allowed_packet.return_value = None processor = PaymentAllocationProcessor(mock_repository) result = processor._get_effective_batch_size() assert result == 10_000 # config default def test_no_clamping_when_packet_large_enough(self): """Test batch_size unchanged when max_allowed_packet is large.""" mock_repository = Mock() mock_repository.get_max_allowed_packet.return_value = 16_777_216 # 16MB processor = PaymentAllocationProcessor(mock_repository) result = processor._get_effective_batch_size() assert result == 10_000 # config default, no clamping def test_clamps_when_packet_too_small(self): """Test batch_size is reduced when max_allowed_packet is small.""" mock_repository = Mock() # 22_000 bytes: after 90% margin (19_800) minus 1000 base = 18_800 available # (18_800 + 1) / (20 + 1) = 895 mock_repository.get_max_allowed_packet.return_value = 22_000 processor = PaymentAllocationProcessor(mock_repository) result = processor._get_effective_batch_size() assert result < 10_000 assert result == 895 def test_clamps_logs_warning(self, caplog): """Test warning is logged when batch_size is clamped.""" mock_repository = Mock() mock_repository.get_max_allowed_packet.return_value = 22_000 processor = PaymentAllocationProcessor(mock_repository) with caplog.at_level(logging.WARNING): processor._get_effective_batch_size() assert 'batch_size clamped' in caplog.text assert 'max_allowed_packet=22000' in caplog.text def test_raises_when_packet_too_small(self): """Test raises ValueError when max_allowed_packet can't fit any entries.""" mock_repository = Mock() mock_repository.get_max_allowed_packet.return_value = 100 # Tiny processor = PaymentAllocationProcessor(mock_repository) with pytest.raises(ValueError, match='too small'): processor._get_effective_batch_size()