"""Unit tests for Repository.""" from decimal import Decimal from unittest.mock import MagicMock, Mock import pymysql import pytest from src.connectors.repository import Repository from src.errors import TransientError from src.schemas import ( ContractAdjustmentCount, LedgerAdjustmentApplied, StatementPeriodPaymentEntity, ) _UNSET = object() def _make_repo(*, fetchone=_UNSET, fetchall=_UNSET, lastrowid=_UNSET): """Create a Repository with a pre-configured mock cursor. Returns: Tuple of (repo, mock_cursor, mock_conn) for assertions. """ mock_cursor = MagicMock() if fetchone is not _UNSET: mock_cursor.fetchone.return_value = fetchone if fetchall is not _UNSET: mock_cursor.fetchall.return_value = fetchall if lastrowid is not _UNSET: mock_cursor.lastrowid = lastrowid mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor return Repository(mock_conn), mock_cursor, mock_conn def _make_transient_repo(): """Create a Repository whose cursor raises a transient OperationalError.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2003, "Can't connect to MySQL server" ) return Repository(mock_conn) class TestRepositoryInit: """Tests for Repository initialization.""" def test_init_stores_connection(self): """Test repository stores database connection.""" mock_conn = Mock() repo = Repository(mock_conn) assert repo.conn is mock_conn class TestRepositoryGetStatementPeriodPaymentEntity: """Tests for Repository.get_statement_period_payment_entity method.""" def test_get_sppe_success(self): """Test successful SPPE lookup.""" repo, _, _ = _make_repo( fetchone={ 'statement_period_payment_entity_id': 1, 'statement_period_id': 456, 'reference_payment_entity_id': 10, } ) result = repo.get_statement_period_payment_entity(1) assert isinstance(result, StatementPeriodPaymentEntity) assert result.statement_period_payment_entity_id == 1 assert result.statement_period_id == 456 assert result.reference_payment_entity_id == 10 def test_get_sppe_not_found(self): """Test SPPE lookup returns None when not found.""" repo, _, _ = _make_repo(fetchone=None) result = repo.get_statement_period_payment_entity(999) assert result is None def test_get_sppe_sql_structure(self): """Test SQL query structure for get_statement_period_payment_entity.""" repo, cursor, _ = _make_repo(fetchone=None) repo.get_statement_period_payment_entity(1) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'statement_period_payment_entity' in sql assert 'statement_period_payment_entity_id' in sql assert 'statement_period_id' in sql assert 'reference_payment_entity_id' in sql assert params == (1,) def test_get_sppe_handles_transient_errors(self): """Test SPPE lookup converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError) as exc_info: repo.get_statement_period_payment_entity(1) assert 'Database connection failed' in str(exc_info.value) class TestRepositoryGetAdjustmentsForPaymentEntity: """Tests for Repository.get_adjustments_for_payment_entity method.""" def test_get_adjustments_success(self): """Test successful adjustment retrieval.""" repo, _, _ = _make_repo( fetchall=[ { 'ledger_adjustment_applied_id': 1, 'account_id': 100, 'contract_id': 200, 'adjustment_amount': Decimal('50.00'), 'adjustment_currency_code': 'USD', 'adjustment_amount_payee_currency': Decimal('50.00'), 'adjustment_payee_currency_code': 'USD', 'account_payee_id': 300, }, { 'ledger_adjustment_applied_id': 2, 'account_id': 101, 'contract_id': 201, 'adjustment_amount': Decimal('75.50'), 'adjustment_currency_code': 'EUR', 'adjustment_amount_payee_currency': Decimal('80.00'), 'adjustment_payee_currency_code': 'USD', 'account_payee_id': 301, }, ] ) result = repo.get_adjustments_for_payment_entity(456, 10, 10_000) assert len(result) == 2 assert isinstance(result[0], LedgerAdjustmentApplied) assert result[0].ledger_adjustment_applied_id == 1 assert result[0].adjustment_amount == Decimal('50.00') assert result[1].ledger_adjustment_applied_id == 2 assert result[1].adjustment_amount == Decimal('75.50') def test_get_adjustments_empty(self): """Test get_adjustments returns empty list when no records found.""" repo, _, _ = _make_repo(fetchall=[]) result = repo.get_adjustments_for_payment_entity(456, 10, 10_000) assert result == [] def test_get_adjustments_sql_structure(self): """Test SQL query structure for get_adjustments.""" repo, cursor, _ = _make_repo(fetchall=[]) repo.get_adjustments_for_payment_entity(456, 10, 10_000) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'ledger_adjustment_applied AS laa' in sql assert 'account_payment_term AS apt' in sql assert 'account_payee AS ap' in sql assert 'apt.payment_entity_id = %s' in sql assert 'laa.statement_period_id = %s' in sql assert 'laa.apply_to_flowthrough_payment = 1' in sql assert 'NOT EXISTS' in sql assert 'payment_allocation_ledger_adjustment pala' in sql assert 'LIMIT' in sql assert params == (10, 456, 10_000) def test_get_adjustments_handles_transient_errors(self): """Test get_adjustments converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError) as exc_info: repo.get_adjustments_for_payment_entity(456, 10, 10_000) assert 'Database connection failed' in str(exc_info.value) class TestRepositoryCreatePaymentAllocation: """Tests for Repository.create_payment_allocation method.""" def test_create_payment_allocation_success(self): """Test successful payment allocation creation.""" repo, cursor, _ = _make_repo(lastrowid=789) result = repo.create_payment_allocation( contract_id=200, payee_type='account_payee', payee_id=300, statement_period_id=456, payment_allocation_type='flowthrough', amount_to_payment=Decimal('100.00'), amount_to_ledger=Decimal('95.00'), currency_code='USD', description='Flowthrough payment allocation', created_by='lambda-abacus-payment-allocation', ) assert result == 789 # Verify SQL cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'INSERT INTO payment_allocation' in sql assert 'contract_id' in sql assert 'payee_type' in sql assert 'payee_id' in sql assert 'payment_status_modified' in sql assert 'ledger_status_modified' in sql assert 'description' in sql assert params[0] == 200 # contract_id assert params[1] == 'account_payee' # payee_type assert params[2] == 300 # payee_id assert params[3] == 456 # statement_period_id assert params[4] == 'flowthrough' # payment_allocation_type assert params[5] == Decimal('100.00') # amount_to_payment assert params[6] == 'init' # payment_status assert params[7] == Decimal('95.00') # amount_to_ledger assert params[8] == 'init' # ledger_status assert params[9] == 'USD' # currency_code assert params[10] == 'Flowthrough payment allocation' # description def test_create_payment_allocation_does_not_commit(self): """Test create_payment_allocation does not commit.""" repo, _, conn = _make_repo(lastrowid=789) repo.create_payment_allocation( contract_id=200, payee_type='account_payee', payee_id=300, statement_period_id=456, payment_allocation_type='flowthrough', amount_to_payment=Decimal('100.00'), amount_to_ledger=Decimal('95.00'), currency_code='USD', description='Flowthrough payment allocation', created_by='test', ) conn.commit.assert_not_called() def test_create_payment_allocation_handles_transient_errors(self): """Test create_payment_allocation converts connection errors.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2006, 'MySQL server has gone away' ) repo = Repository(mock_conn) with pytest.raises(TransientError): repo.create_payment_allocation( contract_id=200, payee_type='account_payee', payee_id=300, statement_period_id=456, payment_allocation_type='flowthrough', amount_to_payment=Decimal('100.00'), amount_to_ledger=Decimal('95.00'), currency_code='USD', description='Flowthrough payment allocation', created_by='test', ) class TestRepositoryCreatePaymentAllocationLedgerAdjustment: """Tests for Repository.create_payment_allocation_ledger_adjustment method.""" def test_create_link_success(self): """Test successful link creation.""" repo, cursor, _ = _make_repo(lastrowid=100) result = repo.create_payment_allocation_ledger_adjustment( payment_allocation_id=789, ledger_adjustment_applied_id=1, created_by='test', ) assert result == 100 # Verify SQL cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'INSERT INTO payment_allocation_ledger_adjustment' in sql assert 'payment_allocation_id' in sql assert 'ledger_adjustment_applied_id' in sql assert params == (789, 1, 'test', 'test') def test_create_link_handles_transient_errors(self): """Test link creation converts connection errors.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2013, 'Lost connection to MySQL server during query' ) repo = Repository(mock_conn) with pytest.raises(TransientError): repo.create_payment_allocation_ledger_adjustment( payment_allocation_id=789, ledger_adjustment_applied_id=1, created_by='test', ) class TestRepositoryGetContractAdjustmentCounts: """Tests for Repository.get_contract_adjustment_counts method.""" def test_get_contract_counts_success(self): """Test successful contract adjustment count retrieval.""" repo, _, _ = _make_repo( fetchall=[ {'contract_id': 200, 'adjustment_count': 5}, {'contract_id': 201, 'adjustment_count': 12}, ] ) result = repo.get_contract_adjustment_counts(456, 10, 250_000) assert len(result) == 2 assert isinstance(result[0], ContractAdjustmentCount) assert result[0].contract_id == 200 assert result[0].adjustment_count == 5 assert result[1].contract_id == 201 assert result[1].adjustment_count == 12 def test_get_contract_counts_empty(self): """Test returns empty list when no contracts found.""" repo, _, _ = _make_repo(fetchall=[]) result = repo.get_contract_adjustment_counts(456, 10, 250_000) assert result == [] def test_get_contract_counts_sql_structure(self): """Test SQL query structure for get_contract_adjustment_counts.""" repo, cursor, _ = _make_repo(fetchall=[]) repo.get_contract_adjustment_counts(456, 10, 250_000) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'GROUP BY laa.contract_id' in sql assert 'COUNT(*)' in sql assert 'apt.payment_entity_id = %s' in sql assert 'laa.statement_period_id = %s' in sql assert 'NOT EXISTS' in sql assert 'LIMIT' in sql assert params == (10, 456, 250_000) def test_get_contract_counts_handles_transient_errors(self): """Test converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError): repo.get_contract_adjustment_counts(456, 10, 250_000) class TestRepositoryGetAdjustmentsForContracts: """Tests for Repository.get_adjustments_for_contracts method.""" def test_get_adjustments_single_contract(self): """Test retrieval for a single contract.""" repo, _, _ = _make_repo( fetchall=[ { 'ledger_adjustment_applied_id': 1, 'account_id': 100, 'contract_id': 200, 'adjustment_amount': Decimal('50.00'), 'adjustment_currency_code': 'USD', 'adjustment_amount_payee_currency': Decimal('50.00'), 'adjustment_payee_currency_code': 'USD', 'account_payee_id': 300, }, ] ) result = repo.get_adjustments_for_contracts(456, 10, [200]) assert len(result) == 1 assert isinstance(result[0], LedgerAdjustmentApplied) assert result[0].contract_id == 200 def test_get_adjustments_multiple_contracts(self): """Test retrieval for multiple contracts.""" repo, _, _ = _make_repo( fetchall=[ { 'ledger_adjustment_applied_id': 1, 'account_id': 100, 'contract_id': 200, 'adjustment_amount': Decimal('50.00'), 'adjustment_currency_code': 'USD', 'adjustment_amount_payee_currency': Decimal('50.00'), 'adjustment_payee_currency_code': 'USD', 'account_payee_id': 300, }, { 'ledger_adjustment_applied_id': 2, 'account_id': 101, 'contract_id': 201, 'adjustment_amount': Decimal('75.00'), 'adjustment_currency_code': 'EUR', 'adjustment_amount_payee_currency': Decimal('80.00'), 'adjustment_payee_currency_code': 'USD', 'account_payee_id': 301, }, ] ) result = repo.get_adjustments_for_contracts(456, 10, [200, 201]) assert len(result) == 2 assert result[0].contract_id == 200 assert result[1].contract_id == 201 def test_get_adjustments_sql_has_in_clause(self): """Test SQL contains IN clause with parameterized placeholders.""" repo, cursor, _ = _make_repo(fetchall=[]) repo.get_adjustments_for_contracts(456, 10, [200, 201, 202]) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'contract_id IN (%s, %s, %s)' in sql assert 'NOT EXISTS' in sql assert params == (10, 456, 200, 201, 202) def test_get_adjustments_with_limit(self): """Test optional LIMIT is appended when provided.""" repo, cursor, _ = _make_repo(fetchall=[]) repo.get_adjustments_for_contracts(456, 10, [200], limit=5000) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'LIMIT' in sql assert params == (10, 456, 200, 5000) def test_get_adjustments_without_limit(self): """Test no LIMIT clause when limit is None.""" repo, cursor, _ = _make_repo(fetchall=[]) repo.get_adjustments_for_contracts(456, 10, [200]) cursor.execute.assert_called_once() sql = cursor.execute.call_args[0][0] params = cursor.execute.call_args[0][1] assert 'LIMIT' not in sql assert params == (10, 456, 200) def test_get_adjustments_handles_transient_errors(self): """Test converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError): repo.get_adjustments_for_contracts(456, 10, [200]) class TestRepositoryGetMaxAllowedPacket: """Tests for Repository.get_max_allowed_packet method.""" def test_get_max_allowed_packet_success(self): """Test successful retrieval of max_allowed_packet.""" repo, cursor, _ = _make_repo( fetchone={ 'Variable_name': 'max_allowed_packet', 'Value': '16777216', } ) result = repo.get_max_allowed_packet() assert result == 16_777_216 cursor.execute.assert_called_once_with( "SHOW VARIABLES LIKE 'max_allowed_packet'" ) def test_get_max_allowed_packet_returns_none(self): """Test returns None when variable not found.""" repo, _, _ = _make_repo(fetchone=None) result = repo.get_max_allowed_packet() assert result is None def test_get_max_allowed_packet_handles_transient_errors(self): """Test converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError): repo.get_max_allowed_packet() class TestRepositoryGetStateStatus: """Tests for Repository.get_state_status method.""" def test_get_state_status_returns_status(self): """Test returns action_status when row is found.""" repo, cursor, _ = _make_repo(fetchone={'action_status': 'complete'}) result = repo.get_state_status( 'statement_period_payment_entity', 123, 'close_balance' ) assert result == 'complete' cursor.execute.assert_called_once() args = cursor.execute.call_args assert 'abacus_state' in args[0][0] assert args[0][1] == ('statement_period_payment_entity', 123, 'close_balance') def test_get_state_status_returns_none_when_not_found(self): """Test returns None when no matching row exists.""" repo, _, _ = _make_repo(fetchone=None) result = repo.get_state_status( 'statement_period_payment_entity', 999, 'close_balance' ) assert result is None def test_get_state_status_handles_transient_errors(self): """Test converts connection errors to TransientError.""" repo = _make_transient_repo() with pytest.raises(TransientError): repo.get_state_status( 'statement_period_payment_entity', 123, 'close_balance' )