"""Tests for src/connectors/snowflake/query.py.""" from __future__ import annotations from unittest.mock import MagicMock import pytest from snowflake.connector.errors import OperationalError as SnowflakeOperationalError from src.connectors.snowflake.connection import SnowflakeErrorCode from src.connectors.snowflake.query import ( _SQL_TEMPLATE, ROYALTY_ACCOUNTING_DATABASE, fetch_contract_data, get_balance_for_input, input_to_column, ) from src.errors import TransientError from src.types import ContractData class TestInputToColumn: """Test input_to_column mapping.""" def test_default(self) -> None: assert input_to_column(None) == 'CLOSING_BALANCE' def test_empty_string(self) -> None: assert input_to_column('') == 'CLOSING_BALANCE' def test_closing_balance(self) -> None: assert input_to_column('closing_balance') == 'CLOSING_BALANCE' def test_gross_revenue(self) -> None: assert input_to_column('gross_revenue') == 'GROSS_REVENUE' def test_unknown_falls_back_to_closing_balance(self) -> None: assert input_to_column('unknown_value') == 'CLOSING_BALANCE' def test_net_revenue(self) -> None: assert input_to_column('net_revenue') == 'NET_REVENUE' class TestGetBalanceForInput: """Test get_balance_for_input routing.""" def test_closing_balance(self) -> None: data = ContractData(closing_balance=100.0) assert get_balance_for_input(data, None) == 100.0 def test_gross_revenue(self) -> None: data = ContractData(gross_revenue=200.0) assert get_balance_for_input(data, 'gross_revenue') == 200.0 def test_net_revenue(self) -> None: data = ContractData(net_revenue=300.0) assert get_balance_for_input(data, 'net_revenue') == 300.0 def test_unavailable_returns_none(self) -> None: data = ContractData() assert get_balance_for_input(data, 'gross_revenue') is None def _make_cursor(rows: list[tuple], col_names: list[str] | None = None) -> MagicMock: """Build a mock Snowflake cursor returning given rows.""" if col_names is None: col_names = [ 'CONTRACT_ID', 'ACCOUNT_ID', 'ACCOUNT_NAME', 'CONTRACT_NAME', 'ACCOUNT_PAYEE_CURRENCY', 'CLOSING_BALANCE', 'GROSS_REVENUE', 'NET_REVENUE', 'PRIOR_CLOSING_BALANCE', ] cursor = MagicMock() cursor.description = [(name,) for name in col_names] cursor.fetchall.return_value = rows return cursor class TestFetchContractData: """Test fetch_contract_data.""" def test_returns_contract_data(self) -> None: cursor = _make_cursor( [(100, 1, 'Acme', 'Deal A', 'USD', 500.0, 1000.0, 800.0, 450.0)] ) mock_conn = MagicMock() mock_conn.cursor.return_value = cursor result = fetch_contract_data(mock_conn, [100]) assert 100 in result cd = result[100] assert cd.account_name == 'Acme' assert cd.closing_balance == 500.0 assert cd.prior_closing_balance == 450.0 def test_empty_ids_returns_empty(self) -> None: mock_conn = MagicMock() result = fetch_contract_data(mock_conn, []) assert result == {} mock_conn.cursor.assert_not_called() def test_missing_contract_not_in_result(self) -> None: cursor = _make_cursor([]) mock_conn = MagicMock() mock_conn.cursor.return_value = cursor result = fetch_contract_data(mock_conn, [999]) assert 999 not in result def test_invalid_contract_id_zero_raises(self) -> None: with pytest.raises(ValueError, match='Invalid contract IDs'): fetch_contract_data(MagicMock(), [0]) def test_invalid_contract_id_negative_raises(self) -> None: with pytest.raises(ValueError, match='Invalid contract IDs'): fetch_contract_data(MagicMock(), [-5]) def test_mixed_valid_and_invalid_raises(self) -> None: with pytest.raises(ValueError, match='Invalid contract IDs'): fetch_contract_data(MagicMock(), [100, -1, 200]) # ═══════════════════════════════════════════════════════════════════════════════ # Transient error classification # ═══════════════════════════════════════════════════════════════════════════════ def _conn_with_execute_error(err: Exception) -> MagicMock: """Build a mock connection whose cursor.execute() raises err.""" cursor = MagicMock() cursor.execute.side_effect = err conn = MagicMock() conn.cursor.return_value = cursor return conn class TestFetchContractDataTransientErrors: """Test that transient Snowflake errors are classified correctly.""" def test_fetch_query_timeout_raises_transient(self) -> None: err = SnowflakeOperationalError( msg='query timeout', errno=SnowflakeErrorCode.QUERY_TIMEOUT ) with pytest.raises(TransientError): fetch_contract_data(_conn_with_execute_error(err), [100]) def test_fetch_warehouse_suspended_raises_transient(self) -> None: err = SnowflakeOperationalError( msg='warehouse suspended', errno=SnowflakeErrorCode.WAREHOUSE_SUSPENDED ) with pytest.raises(TransientError): fetch_contract_data(_conn_with_execute_error(err), [100]) def test_fetch_connection_error_raises_transient(self) -> None: err = SnowflakeOperationalError( msg='connection error', errno=SnowflakeErrorCode.CONNECTION_ERROR ) with pytest.raises(TransientError): fetch_contract_data(_conn_with_execute_error(err), [100]) def test_fetch_internal_service_error_raises_transient(self) -> None: err = SnowflakeOperationalError( msg='internal service error', errno=SnowflakeErrorCode.INTERNAL_SERVICE_ERROR, ) with pytest.raises(TransientError): fetch_contract_data(_conn_with_execute_error(err), [100]) def test_fetch_connection_failed_message_raises_transient(self) -> None: err = SnowflakeOperationalError( msg='Failed to establish a connection to Snowflake', errno=99999 ) with pytest.raises(TransientError): fetch_contract_data(_conn_with_execute_error(err), [100]) def test_fetch_unknown_operational_error_reraises(self) -> None: err = SnowflakeOperationalError(msg='syntax error in SQL', errno=99999) with pytest.raises(SnowflakeOperationalError): fetch_contract_data(_conn_with_execute_error(err), [100]) # ═══════════════════════════════════════════════════════════════════════════════ # ACC-10361: current statement period selection via current_sp CTE # ═══════════════════════════════════════════════════════════════════════════════ class TestSqlTemplateCurrentStatementPeriod: """Verify the generated SQL uses the current_sp CTE for period selection. These tests assert the SQL structure so a regression (e.g. reverting to ORDER BY STATEMENT_PERIOD_ID DESC without the CTE) is caught immediately. """ def test_current_sp_cte_present(self) -> None: assert 'current_sp AS (' in _SQL_TEMPLATE def test_current_sp_queries_statement_period_table(self) -> None: assert f'{ROYALTY_ACCOUNTING_DATABASE}.statement_period' in _SQL_TEMPLATE def test_current_sp_filters_by_status(self) -> None: assert "statement_period_status = 'current'" in _SQL_TEMPLATE def test_ranking_references_current_sp(self) -> None: assert 'SELECT statement_period_id FROM current_sp' in _SQL_TEMPLATE def test_ranking_does_not_use_view_status_column(self) -> None: assert 'STATEMENT_PERIOD_STATUS' not in _SQL_TEMPLATE def test_execute_passes_generated_sql(self) -> None: cursor = _make_cursor( [(100, 1, 'Acme', 'Deal A', 'USD', 500.0, 1000.0, 800.0, 450.0)] ) mock_conn = MagicMock() mock_conn.cursor.return_value = cursor fetch_contract_data(mock_conn, [100]) executed_sql: str = cursor.execute.call_args[0][0] assert 'current_sp AS (' in executed_sql assert f'{ROYALTY_ACCOUNTING_DATABASE}.statement_period' in executed_sql assert 'SELECT statement_period_id FROM current_sp' in executed_sql