"""Test CortexSearchClient.""" import logging from unittest.mock import MagicMock, patch, PropertyMock import pytest from payment.connectors.cortex_search.client import ( CortexSearchClient, CortexSearchConfig, ) from payment.connectors.cortex_search.exceptions import ( CortexSearchAuthError, CortexSearchConfigError, CortexSearchServiceError, ) RAW_KEY = b'raw_private_key_bytes' # Snapshot the original ``service`` descriptor at import time so we can # restore it between tests. Several tests below monkey-patch # ``type(client).service = PropertyMock(...)``, which mutates the class and # leaks into subsequent tests if not reset. _ORIGINAL_SERVICE_DESCRIPTOR = CortexSearchClient.__dict__['service'] @pytest.fixture(autouse=True) def _restore_service_descriptor(): yield CortexSearchClient.service = _ORIGINAL_SERVICE_DESCRIPTOR @pytest.fixture(autouse=True) def mock_load_key(): with patch( 'payment.connectors.cortex_search.client.load_private_key', return_value=RAW_KEY, ) as m: yield m def make_config(**kwargs): defaults = dict( SNOWFLAKE_ACCOUNT='test-account', SNOWFLAKE_USER='test-user', SNOWFLAKE_DATABASE='TEST_DB', SNOWFLAKE_SCHEMA='TEST_SCHEMA', SNOWFLAKE_CORTEX_SEARCH_SERVICE_NAME='TEST_SERVICE', SNOWFLAKE_PRIVATE_KEY=RAW_KEY.decode(), RETRY_WAIT_MULTIPLIER=0, RETRY_WAIT_MIN=0, RETRY_WAIT_MAX=0, RETRY_MAX_ATTEMPTS=1, ) defaults.update(kwargs) return CortexSearchConfig(**defaults) class TestEnsurePrivateKey: def test_passes_key_to_load_private_key_when_provided(self, mock_load_key): config = make_config( SNOWFLAKE_PRIVATE_KEY='pem-key-string', SNOWFLAKE_PRIVATE_KEY_PATH=None ) client = CortexSearchClient(config, None) mock_load_key.assert_called_once_with('pem-key-string', None, None) assert client._private_key == RAW_KEY @patch('payment.connectors.cortex_search.client.load_private_key') def test_loads_key_from_path(self, mock_load): mock_load.return_value = b'loaded_key' config = make_config( SNOWFLAKE_PRIVATE_KEY=None, SNOWFLAKE_PRIVATE_KEY_PATH='/path/to/key.p8', SNOWFLAKE_KEY_PASSPHRASE='pass', ) client = CortexSearchClient(config, None) assert client._private_key == b'loaded_key' mock_load.assert_called_once_with(None, '/path/to/key.p8', 'pass') def test_raises_config_error_when_no_key_provided(self): config = make_config( SNOWFLAKE_PRIVATE_KEY=None, SNOWFLAKE_PRIVATE_KEY_PATH=None, ) with pytest.raises(CortexSearchConfigError): CortexSearchClient(config, None) @patch('payment.connectors.cortex_search.client.load_private_key') def test_raises_config_error_when_key_loading_fails(self, mock_load): mock_load.side_effect = OSError('permission denied') config = make_config( SNOWFLAKE_PRIVATE_KEY=None, SNOWFLAKE_PRIVATE_KEY_PATH='/bad/path.p8', ) with pytest.raises(CortexSearchConfigError, match='Error loading private key'): CortexSearchClient(config, None) class TestServiceProperty: @patch('payment.connectors.cortex_search.client.Root') @patch('payment.connectors.cortex_search.client.connect') def test_creates_connection_and_service(self, mock_connect, mock_Root): config = make_config() client = CortexSearchClient(config, None) mock_conn = MagicMock() mock_connect.return_value = mock_conn mock_service = MagicMock() ( mock_Root.return_value.databases.__getitem__.return_value.schemas.__getitem__.return_value.cortex_search_services.__getitem__.return_value ) = mock_service result = client.service assert result is mock_service mock_connect.assert_called_once_with( account='test-account', user='test-user', private_key=RAW_KEY, authenticator='snowflake_jwt', login_timeout=None, network_timeout=None, ) mock_Root.assert_called_once_with(mock_conn) @patch('payment.connectors.cortex_search.client.Root') @patch('payment.connectors.cortex_search.client.connect') def test_passes_timeouts_to_connect(self, mock_connect, mock_Root): config = make_config(SNOWFLAKE_LOGIN_TIMEOUT=10, SNOWFLAKE_NETWORK_TIMEOUT=30) client = CortexSearchClient(config, None) mock_connect.return_value = MagicMock() _ = client.service call_kwargs = mock_connect.call_args.kwargs assert call_kwargs['login_timeout'] == 10 assert call_kwargs['network_timeout'] == 30 @patch('payment.connectors.cortex_search.client.Root') @patch('payment.connectors.cortex_search.client.connect') def test_reuses_existing_connection(self, mock_connect, mock_Root): config = make_config() client = CortexSearchClient(config, None) mock_connect.return_value = MagicMock() _ = client.service _ = client.service mock_connect.assert_called_once() mock_Root.assert_called_once() class TestLoggerProperty: def test_returns_module_logger_when_no_factory(self): config = make_config() client = CortexSearchClient(config, None) import payment.connectors.cortex_search.client as client_module assert client.logger is client_module.logger def test_returns_factory_logger_when_provided(self): custom_logger = MagicMock(spec=logging.Logger) config = make_config() client = CortexSearchClient(config, lambda: custom_logger) assert client.logger is custom_logger def test_falls_back_to_module_logger_when_factory_returns_none(self): config = make_config() client = CortexSearchClient(config, lambda: None) import payment.connectors.cortex_search.client as client_module assert client.logger is client_module.logger class TestSearch: def _make_client_with_service(self, mock_service): config = make_config() client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) return client def test_returns_results_on_success(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-1' mock_response.results = [{'col': 'val'}] mock_service.search.return_value = mock_response client = self._make_client_with_service(mock_service) result = client._search('my query', ['col']) assert result == [{'col': 'val'}] def test_passes_query_params_to_service(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-2' mock_response.results = [] mock_service.search.return_value = mock_response client = self._make_client_with_service(mock_service) query_filter = {'@eq': {'status': 'OPEN'}} scoring = {'semantic_weight': 0.8} client._search( 'query', ['a', 'b'], query_filter=query_filter, scoring_config=scoring, limit=5, ) mock_service.search.assert_called_once_with( query='query', columns=['a', 'b'], filter=query_filter, scoring_config=scoring, limit=5, ) def test_uses_config_scoring_when_none_passed(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-3' mock_response.results = [] mock_service.search.return_value = mock_response config_scoring = {'bm25_weight': 1.0} config = make_config(SNOWFLAKE_CORTEX_SEARCH_SCORING_CONFIG=config_scoring) client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) client._search('query', ['col']) call_kwargs = mock_service.search.call_args.kwargs assert call_kwargs['scoring_config'] == config_scoring def test_passes_nested_filter_to_service(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-nested' mock_response.results = [] mock_service.search.return_value = mock_response client = self._make_client_with_service(mock_service) nested_filter = { '@and': [ {'@or': [{'@eq': {'status': 'OPEN'}}, {'@eq': {'status': 'PENDING'}}]}, {'@not': {'@eq': {'region': 'EU'}}}, ] } client._search('query', ['col'], query_filter=nested_filter) mock_service.search.assert_called_once_with( query='query', columns=['col'], filter=nested_filter, scoring_config=None, limit=20, ) def test_validate_filter_called_with_filter(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-vf' mock_response.results = [] mock_service.search.return_value = mock_response client = self._make_client_with_service(mock_service) query_filter = {'@eq': {'status': 'OPEN'}} with patch( 'payment.connectors.cortex_search.client.validate_filter' ) as mock_vf: client._search('query', ['col'], query_filter=query_filter) mock_vf.assert_called_once_with(query_filter) def test_validate_filter_called_with_none_when_no_filter(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-vf-none' mock_response.results = [] mock_service.search.return_value = mock_response client = self._make_client_with_service(mock_service) with patch( 'payment.connectors.cortex_search.client.validate_filter' ) as mock_vf: client._search('query', ['col']) mock_vf.assert_called_once_with(None) def test_validate_filter_error_propagates(self): from payment.connectors.cortex_search.exceptions import ( CortexSearchBadRequestError, ) mock_service = MagicMock() client = self._make_client_with_service(mock_service) with patch( 'payment.connectors.cortex_search.client.validate_filter', side_effect=CortexSearchBadRequestError('bad filter'), ): with pytest.raises(CortexSearchBadRequestError, match='bad filter'): client._search( 'query', ['col'], query_filter={'@gte': {'amount': 'bad'}} ) mock_service.search.assert_not_called() @pytest.mark.parametrize( 'make_exc', [ pytest.param( lambda: ( __import__( 'snowflake.connector', fromlist=['errors'] ).errors.ForbiddenError ), id='connector.ForbiddenError', ), pytest.param( lambda: ( __import__( 'snowflake.connector', fromlist=['errors'] ).errors.TokenExpiredError ), id='connector.TokenExpiredError', ), pytest.param( lambda: __import__( 'snowflake.connector.network', fromlist=['ReauthenticationRequest'] ).ReauthenticationRequest(cause=Exception('token expired')), id='connector.ReauthenticationRequest', ), pytest.param( lambda: __import__( 'snowflake.core', fromlist=['exceptions'] ).exceptions.UnauthorizedError(MagicMock()), id='core.UnauthorizedError', ), pytest.param( lambda: __import__( 'snowflake.core', fromlist=['exceptions'] ).exceptions.ForbiddenError(MagicMock()), id='core.ForbiddenError', ), ], ) def test_auth_error_calls_close_and_raises(self, make_exc): mock_service = MagicMock() mock_service.search.side_effect = make_exc() config = make_config() client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) client.close = MagicMock() with pytest.raises(CortexSearchAuthError): client._search('query', ['col']) client.close.assert_called_once() def test_service_error_raises_mapped_exception(self): from snowflake.connector import errors mock_service = MagicMock() mock_service.search.side_effect = errors.InternalServerError client = self._make_client_with_service(mock_service) with pytest.raises(CortexSearchServiceError): client._search('query', ['col']) def test_search_retries_on_retryable_error(self): mock_service = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-retry' mock_response.results = [{'id': 1}] mock_service.search.side_effect = [ CortexSearchServiceError('temporary failure'), mock_response, ] config = make_config(RETRY_MAX_ATTEMPTS=2) client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) result = client.search('query', ['id']) assert result == [{'id': 1}] assert mock_service.search.call_count == 2 def test_search_raises_after_max_attempts(self): mock_service = MagicMock() mock_service.search.side_effect = CortexSearchServiceError('always fails') config = make_config(RETRY_MAX_ATTEMPTS=2) client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) with pytest.raises(CortexSearchServiceError): client.search('query', ['col']) assert mock_service.search.call_count == 2 @pytest.mark.parametrize( 'make_exc', [ pytest.param( lambda: ( __import__( 'snowflake.connector', fromlist=['errors'] ).errors.TokenExpiredError ), id='connector.TokenExpiredError', ), pytest.param( lambda: __import__( 'snowflake.connector.network', fromlist=['ReauthenticationRequest'] ).ReauthenticationRequest(cause=Exception('token expired')), id='connector.ReauthenticationRequest', ), pytest.param( lambda: __import__( 'snowflake.core', fromlist=['exceptions'] ).exceptions.UnauthorizedError(MagicMock()), id='core.UnauthorizedError', ), ], ) @patch('payment.connectors.cortex_search.client.Root') @patch('payment.connectors.cortex_search.client.connect') def test_search_reconnects_and_retries_on_token_expired_error( self, mock_connect, mock_Root, make_exc ): """Token-expired error must close, reconnect, and retry. Exercises the real ``service`` property so the test catches a regression where ``close()`` is removed or the lazy reconnect path breaks. We assert ``connect`` is called twice (once per attempt). """ mock_service_expired = MagicMock() mock_service_expired.search.side_effect = make_exc() mock_service_fresh = MagicMock() mock_response = MagicMock() mock_response.request_id = 'req-retry-after-token-expiry' mock_response.results = [{'id': 1}] mock_service_fresh.search.return_value = mock_response def make_root(service): root = MagicMock() ( root.databases.__getitem__.return_value.schemas.__getitem__.return_value.cortex_search_services.__getitem__.return_value ) = service return root mock_Root.side_effect = [ make_root(mock_service_expired), make_root(mock_service_fresh), ] mock_connect.side_effect = [MagicMock(), MagicMock()] config = make_config(RETRY_MAX_ATTEMPTS=2) client = CortexSearchClient(config, None) result = client.search('query', ['id']) assert result == [{'id': 1}] assert mock_connect.call_count == 2 assert mock_Root.call_count == 2 mock_service_expired.search.assert_called_once() mock_service_fresh.search.assert_called_once() @pytest.mark.parametrize( 'make_exc', [ pytest.param( lambda: ( __import__( 'snowflake.connector', fromlist=['errors'] ).errors.ForbiddenError ), id='connector.ForbiddenError', ), pytest.param( lambda: __import__( 'snowflake.core', fromlist=['exceptions'] ).exceptions.ForbiddenError(MagicMock()), id='core.ForbiddenError', ), ], ) def test_search_does_not_retry_on_forbidden_error(self, make_exc): """Forbidden (403) errors are permission failures and must not retry.""" mock_service = MagicMock() mock_service.search.side_effect = make_exc() # Budget of 3 attempts proves the absence of retries is intentional, # not just a side-effect of RETRY_MAX_ATTEMPTS=1. config = make_config(RETRY_MAX_ATTEMPTS=3) client = CortexSearchClient(config, None) type(client).service = PropertyMock(return_value=mock_service) client.close = MagicMock() with pytest.raises(CortexSearchAuthError): client.search('query', ['col']) assert mock_service.search.call_count == 1 client.close.assert_called_once() class TestClose: def test_closes_connection_and_resets_state(self): config = make_config() client = CortexSearchClient(config, None) mock_conn = MagicMock() client._connection = mock_conn client._service = MagicMock() client.close() mock_conn.close.assert_called_once() assert client._connection is None assert client._service is None def test_close_is_noop_when_not_connected(self): config = make_config() client = CortexSearchClient(config, None) assert client._connection is None client.close() # should not raise def test_close_logs_error_if_connection_close_fails(self): config = make_config() client = CortexSearchClient(config, None) mock_conn = MagicMock() mock_conn.close.side_effect = Exception('close failed') client._connection = mock_conn mock_logger = MagicMock() with patch.object( type(client), 'logger', new_callable=PropertyMock, return_value=mock_logger ): client.close() mock_logger.error.assert_called_once() assert client._connection is None