"""Tests for repository module.""" from typing import Any from unittest.mock import MagicMock import pytest from pymysql.err import OperationalError from sync_contract_sap.error_handling import TransientError from sync_contract_sap.repository import Repository from sync_contract_sap.schemas import ContractSyncRow def _make_conn(rows: list[Any]) -> tuple[MagicMock, MagicMock]: """Build a mock connection whose cursor returns rows.""" mock_conn = MagicMock() mock_cursor = MagicMock() mock_conn.cursor.return_value.__enter__ = MagicMock(return_value=mock_cursor) mock_conn.cursor.return_value.__exit__ = MagicMock(return_value=False) mock_cursor.fetchall.return_value = rows return mock_conn, mock_cursor def test_get_contracts_pending_sap_sync_passes_params() -> None: """Test the query is executed with stale minutes and limit as parameters.""" mock_conn, mock_cursor = _make_conn([]) repo = Repository(mock_conn) repo.get_contracts_pending_sap_sync(limit=250, stale_sync_minutes=5) call_args = mock_cursor.execute.call_args assert call_args[0][1] == (5, 250) assert 'ORDER BY' in call_args[0][0] def test_get_contracts_pending_sap_sync_returns_rows() -> None: """Test that fetched rows are parsed into ContractSyncRow models.""" rows = [ { 'contract_id': 1, 'contract_type': 'distribution', 'account_id': 11, 'sap_created_at': None, 'abacus_state_id': None, 'action_status': None, 'state_created_at': None, 'state_last_modified': None, } ] mock_conn, _ = _make_conn(rows) repo = Repository(mock_conn) result = repo.get_contracts_pending_sap_sync(limit=10, stale_sync_minutes=5) assert result == [ContractSyncRow.model_validate(rows[0])] def test_get_contracts_pending_sap_sync_empty() -> None: """Test that an empty result set is returned as an empty list.""" mock_conn, _ = _make_conn([]) repo = Repository(mock_conn) result = repo.get_contracts_pending_sap_sync(limit=250, stale_sync_minutes=5) assert result == [] def test_get_contracts_pending_sap_sync_transient_error() -> None: """Test transient MySQL errors are classified as TransientError.""" mock_conn, mock_cursor = _make_conn([]) mock_cursor.execute.side_effect = OperationalError( 2013, 'Lost connection to MySQL server during query' ) repo = Repository(mock_conn) with pytest.raises(TransientError, match='Lost connection'): repo.get_contracts_pending_sap_sync(limit=250, stale_sync_minutes=5) def test_get_contracts_pending_sap_sync_non_transient_error_reraised() -> None: """Test non-transient MySQL errors are re-raised unchanged.""" mock_conn, mock_cursor = _make_conn([]) mock_cursor.execute.side_effect = OperationalError(1045, 'Access denied') repo = Repository(mock_conn) with pytest.raises(OperationalError, match='Access denied'): repo.get_contracts_pending_sap_sync(limit=250, stale_sync_minutes=5)