"""Tests for src/connectors/snowflake/connection.py.""" from __future__ import annotations import pytest from pydantic import ValidationError from snowflake.connector.errors import OperationalError as SnowflakeOperationalError from src.connectors.snowflake.connection import ( TRANSIENT_ERROR_CODES, SnowflakeConfig, SnowflakeErrorCode, build_snowflake_config, handle_snowflake_errors, ) from src.errors import TransientError class TestSnowflakeConfig: """Tests for the SnowflakeConfig Pydantic model.""" def _minimal(self, **overrides: object) -> SnowflakeConfig: base: dict = { 'account': 'xy12345', 'user': 'svc_user', 'warehouse': 'compute_wh', 'database': 'my_db', 'schema': 'my_schema', } base.update(overrides) return SnowflakeConfig(**base) def test_construction(self) -> None: cfg = self._minimal() assert cfg.account == 'xy12345' assert cfg.user == 'svc_user' assert cfg.warehouse == 'compute_wh' assert cfg.database == 'my_db' def test_schema_alias_populates_schema_name(self) -> None: cfg = self._minimal(schema='prod_schema') assert cfg.schema_name == 'prod_schema' def test_optional_fields_default_to_none(self) -> None: cfg = self._minimal() assert cfg.host is None assert cfg.role is None assert cfg.private_key is None assert cfg.private_key_path is None assert cfg.private_key_passphrase is None def test_default_connection_timeout(self) -> None: cfg = self._minimal() assert cfg.connection_timeout == 10 def test_default_autocommit_is_false(self) -> None: cfg = self._minimal() assert cfg.autocommit is False def test_frozen_model_raises_on_mutation(self) -> None: cfg = self._minimal() with pytest.raises((ValidationError, TypeError)): cfg.account = 'other' # type: ignore[misc] def test_missing_required_field_raises(self) -> None: with pytest.raises(ValidationError): SnowflakeConfig( user='u', warehouse='w', database='d', schema='s' ) # missing account class TestBuildSnowflakeConfig: """Tests for build_snowflake_config().""" @pytest.fixture(autouse=True) def _snowflake_env(self, monkeypatch: pytest.MonkeyPatch) -> None: build_snowflake_config.cache_clear() monkeypatch.setenv('ENVIRONMENT', 'dev') monkeypatch.setenv('SNOWFLAKE_ACCOUNT', 'default_account') monkeypatch.setenv('SNOWFLAKE_USER', 'default_user') monkeypatch.setenv('SNOWFLAKE_WAREHOUSE', 'default_wh') monkeypatch.setenv('SNOWFLAKE_DATABASE', 'default_db') monkeypatch.setenv('SNOWFLAKE_SCHEMA', 'default_schema') def test_reads_account_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_ACCOUNT', 'my_account') cfg = build_snowflake_config() assert cfg.account == 'my_account' def test_reads_user_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_USER', 'svc_user') cfg = build_snowflake_config() assert cfg.user == 'svc_user' def test_missing_database_raises(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv('SNOWFLAKE_DATABASE', raising=False) with pytest.raises(KeyError): build_snowflake_config() def test_database_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_DATABASE', 'custom_db') cfg = build_snowflake_config() assert cfg.database == 'custom_db' def test_schema_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_SCHEMA', 'custom_schema') cfg = build_snowflake_config() assert cfg.schema_name == 'custom_schema' def test_private_key_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_PRIVATE_KEY', 'fake_pem_content') cfg = build_snowflake_config() assert cfg.private_key == 'fake_pem_content' def test_role_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv('SNOWFLAKE_ROLE', 'my_role') cfg = build_snowflake_config() assert cfg.role == 'my_role' def test_returns_snowflake_config_instance(self) -> None: cfg = build_snowflake_config() assert isinstance(cfg, SnowflakeConfig) class TestSnowflakeErrorCode: """Tests for SnowflakeErrorCode and TRANSIENT_ERROR_CODES.""" def test_all_transient_codes_present(self) -> None: for code in SnowflakeErrorCode: assert code in TRANSIENT_ERROR_CODES def test_known_values(self) -> None: assert SnowflakeErrorCode.QUERY_TIMEOUT == 57002 assert SnowflakeErrorCode.WAREHOUSE_SUSPENDED == 604 assert SnowflakeErrorCode.CONNECTION_ERROR == 252001 assert SnowflakeErrorCode.STATEMENT_CANCELLED == 607 assert SnowflakeErrorCode.INTERNAL_SERVICE_ERROR == 300001 class TestHandleSnowflakeErrors: """Tests for the handle_snowflake_errors decorator.""" def test_successful_call_returns_value(self) -> None: @handle_snowflake_errors def func() -> int: return 42 assert func() == 42 def test_transient_errno_raises_transient_error(self) -> None: @handle_snowflake_errors def func() -> None: raise SnowflakeOperationalError( msg='query timeout', errno=SnowflakeErrorCode.QUERY_TIMEOUT ) with pytest.raises(TransientError): func() def test_warehouse_suspended_raises_transient(self) -> None: @handle_snowflake_errors def func() -> None: raise SnowflakeOperationalError( msg='warehouse suspended', errno=SnowflakeErrorCode.WAREHOUSE_SUSPENDED ) with pytest.raises(TransientError): func() def test_connection_failed_message_raises_transient(self) -> None: @handle_snowflake_errors def func() -> None: raise SnowflakeOperationalError( msg='Failed to establish a connection to Snowflake', errno=99999 ) with pytest.raises(TransientError): func() def test_unknown_errno_reraises_operational_error(self) -> None: @handle_snowflake_errors def func() -> None: raise SnowflakeOperationalError(msg='syntax error', errno=99999) with pytest.raises(SnowflakeOperationalError): func() def test_non_operational_error_propagates_unchanged(self) -> None: @handle_snowflake_errors def func() -> None: raise ValueError('bad input') with pytest.raises(ValueError, match='bad input'): func() def test_preserves_function_name(self) -> None: @handle_snowflake_errors def my_function() -> None: pass assert my_function.__name__ == 'my_function'