Spaces:
Running
Running
Download tests/test_request_validation.py from astronolan/AION-Search: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/spaces/astronolan/AION-Search/resolve/main/tests/test_request_validation.py
- Command line
-
hf download hf://spaces/astronolan/AION-Search/tests/test_request_validation.py
-
curl -L -o test_request_validation.py https://huggingface.co/spaces/astronolan/AION-Search/resolve/main/tests/test_request_validation.py
10.6 kB
| """Abuse regressions with real Dash request dispatch and mocked upstream work.""" | |
| import base64 | |
| import copy | |
| import json | |
| import unittest | |
| from unittest.mock import Mock, patch | |
| import httpx | |
| import numpy as np | |
| import requests | |
| from openai import OpenAI | |
| from test_search_limits import DashFixture, SearchFixture, IMAGE | |
| from src.callbacks import prepare_search_data | |
| from src.config import ZILLIZ_PRIMARY_KEY | |
| import src.services as services | |
| from src.services import EmbeddingService, ImageProcessingService, SearchService, ZillizService | |
| from src.search_limits import SearchLimitError | |
| from src.upstream import OPENAI_MAX_RETRIES, OPENAI_TIMEOUT, ZILLIZ_TIMEOUT, UpstreamTimeoutError | |
| from src.url_state import decode_search_state | |
| from src.validation import MAX_REQUEST_BYTES, MAX_ENCODED_STATE_CHARS, validate_result_state | |
| def result_store(count=300): | |
| data = {key: [value] * count for key, value in [ | |
| (ZILLIZ_PRIMARY_KEY, 'test-galaxy'), ('ra', 180.0), ('dec', 0.0), | |
| ('r_mag', 17.0), ('distance', 0.9), | |
| ('cutout_url', 'https://alasky.cds.unistra.fr/hips-image-services/hips2fits?ra=180&dec=0'), | |
| ]} | |
| data.update(query='galaxy', loaded_count=min(60, count)) | |
| return data | |
| class NumericValidationTest(SearchFixture): | |
| def test_magnitudes_rejected_before_embedding(self): | |
| for low, high in [('13 OR r_mag >= 0', 20), (True, 20), (float('nan'), 20), | |
| (13, float('inf')), (12, 20), (13, 21), (20, 13), (10**400, 20)]: | |
| with self.subTest(low=low, high=high), self.assertRaises(SearchLimitError): | |
| self.service.search_text('galaxy', rmag_min=low, rmag_max=high) | |
| self.assert_no_upstream() | |
| def test_all_advanced_inputs_validated_before_first_embedding(self): | |
| cases = [dict(text_weights=[]), dict(text_weights=[float('nan')]), | |
| dict(text_weights=[11]), dict(text_weights=[True]), | |
| dict(image_weights=[]), dict(image_weights=[float('inf')]), | |
| dict(image_queries=[dict(IMAGE, ra=-1)]), | |
| dict(image_queries=[dict(IMAGE, ra=361)]), | |
| dict(image_queries=[dict(IMAGE, dec=91)]), | |
| dict(image_queries=[dict(IMAGE, ra='180 OR true')]), | |
| dict(image_queries=[dict(IMAGE, fov=0)]), | |
| dict(image_queries=[dict(IMAGE, fov=float('inf'))]), | |
| dict(rmag_min='13 OR true'), dict(top_k=301)] | |
| for change in cases: | |
| kwargs = dict(text_queries=['galaxy'], text_weights=[1], image_queries=[IMAGE], image_weights=[1]) | |
| kwargs.update(change) | |
| with self.subTest(change=change), self.assertRaises(SearchLimitError): | |
| self.service.search_advanced(**kwargs) | |
| self.assert_no_upstream() | |
| def test_vector_operation_lengths_and_values(self): | |
| for operations in ([], ['+', '-'], ['invalid']): | |
| with self.subTest(operations=operations), self.assertRaises(SearchLimitError): | |
| self.service.search_vector(['galaxy'], operations) | |
| self.assert_no_upstream() | |
| def test_valid_numeric_boundaries_still_search(self): | |
| self.service.search_advanced(['galaxy'], [-10], | |
| [dict(IMAGE, ra=360, dec=-90)], [10], rmag_min=13, rmag_max=20) | |
| self.zilliz.search.assert_called_once() | |
| def test_invalid_share_states_reject_without_restoring(self): | |
| for change in [dict(tw=[]), dict(tw=[float('nan')]), dict(rmin='13 OR true'), | |
| dict(iq=[[180, 91, .025]], iw=[1]), dict(iq=[[180, 0]], iw=[1])]: | |
| state = dict(tq=['galaxy'], tw=[1]) | |
| state.update(change) | |
| encoded = base64.urlsafe_b64encode(json.dumps(state).encode()).decode() | |
| decoded = decode_search_state(encoded) | |
| self.assertIn('error', decoded) | |
| self.assertEqual(decoded['text_queries'], []) | |
| def test_share_size_rejected_before_base64_decode(self): | |
| with patch('src.url_state.base64.urlsafe_b64decode') as decode: | |
| self.assertIn('error', decode_search_state('a' * (MAX_ENCODED_STATE_CHARS + 1))) | |
| decode.assert_not_called() | |
| class CallbackValidationTest(DashFixture): | |
| def test_bad_numeric_values_and_parallel_arrays_rejected(self): | |
| base = ['galaxy', 'galaxy', [13, 20], {'display': 'block'}, '+', 'text', None, None, | |
| ['text'], ['blue'], [None], [None], ['+']] | |
| for index, value in [(2, ['13 OR true', 20]), (4, 'nan'), (4, '11'), | |
| (5, 'invalid'), (9, []), (12, ['+','+'])]: | |
| states = copy.deepcopy(base) | |
| states[index] = value | |
| result = self.post_callback('perform_search', [1,None,1,None,[]], states, | |
| 'search-button-advanced.n_clicks') | |
| self.assertIsNone(result['search-data']['data']) | |
| self.assertEqual(result['search-results']['children']['props']['color'], 'warning') | |
| self.assert_no_upstream() | |
| def test_load_more_rejects_forged_state_before_rendering(self): | |
| cases = [result_store(301)] | |
| for key, value in [('loaded_count', 1000000), ('loaded_count', -1), ('loaded_count', True), | |
| ('dec', []), ('ra', [float('nan')] * 300)]: | |
| data = result_store() | |
| data[key] = value | |
| cases.append(data) | |
| for data in cases: | |
| with patch('src.callbacks.build_galaxy_card') as card: | |
| result = self.post_callback('load_more_galaxies', [1], [data], 'load-more-button.n_clicks') | |
| card.assert_not_called() | |
| self.assertIsNone(result['search-data']['data']) | |
| def test_csv_rejects_forged_state_before_dataframe(self): | |
| for data in [result_store(301), dict(result_store(), distance=[]), dict(result_store(), loaded_count=999)]: | |
| with patch('src.callbacks.pd.DataFrame') as frame: | |
| self.post_callback('download_csv', [1,None], [data], 'download-button.n_clicks', expected_status=204) | |
| frame.assert_not_called() | |
| def test_normal_pagination_csv_and_modal_work(self): | |
| data = result_store() | |
| for loaded in (180, 300): | |
| response = self.post_callback('load_more_galaxies', [1], [data], 'load-more-button.n_clicks') | |
| data = response['search-data']['data'] | |
| self.assertEqual(data['loaded_count'], loaded) | |
| response = self.post_callback('download_csv', [1,None], [data], 'download-button.n_clicks') | |
| self.assertEqual(len(response['download-csv']['data']['content'].splitlines()), 301) | |
| response = self.post_callback('expand_galaxy_from_url', [data], ['test-galaxy'], 'search-data.data') | |
| self.assertTrue(response['galaxy-modal']['is_open']) | |
| def test_small_result_sets_have_valid_pagination(self): | |
| data = prepare_search_data(self.zilliz.search.return_value, 'galaxy') | |
| self.assertEqual(data['loaded_count'], 1) | |
| self.assertEqual(validate_result_state(data), 1) | |
| def test_large_legitimate_callback_fits_body_limit(self): | |
| data = result_store() | |
| data['query'] = 'π' * 10_000 | |
| size = len(json.dumps(data).encode()) | |
| self.assertLess(size, MAX_REQUEST_BYTES // 4) | |
| response = self.post_callback('load_more_galaxies', [1], [data], 'load-more-button.n_clicks') | |
| self.assertEqual(response['search-data']['data']['loaded_count'], 180) | |
| def test_body_limit_rejects_before_dispatch(self): | |
| with patch('src.callbacks.build_galaxy_card') as card: | |
| response = self.client.post('/_dash-update-component', data=b'x' * (MAX_REQUEST_BYTES + 1), | |
| content_type='application/json') | |
| self.assertEqual(response.status_code, 413) | |
| card.assert_not_called() | |
| self.assert_no_upstream() | |
| def test_timeout_returns_friendly_retry_message(self): | |
| self.embedding.encode_text_query.side_effect = UpstreamTimeoutError() | |
| response = self.search() | |
| self.assertIn('Please try again', str(response['search-results'])) | |
| self.assertIsNone(response['search-data']['data']) | |
| class UpstreamTimeoutTest(unittest.TestCase): | |
| def test_all_zilliz_requests_have_timeouts_and_do_not_retry(self): | |
| zilliz = ZillizService() | |
| dimension = getattr(services, 'ZILLIZ_VECTOR_DIM', 1024) | |
| with patch('src.services.ZILLIZ_ENDPOINT', 'https://example.org/search'): | |
| for action in (lambda: ImageProcessingService().encode_image(180, 0), | |
| lambda: zilliz.search(np.ones(dimension))): | |
| with patch('src.services.requests.post', side_effect=requests.ReadTimeout) as post: | |
| with self.assertRaises(UpstreamTimeoutError): | |
| action() | |
| self.assertEqual(post.call_count, 1) | |
| self.assertEqual(post.call_args.kwargs['timeout'], ZILLIZ_TIMEOUT) | |
| with patch('src.services.requests.post', side_effect=requests.ConnectTimeout) as post: | |
| self.assertEqual(zilliz.get_collection_count(), 0) | |
| self.assertEqual(post.call_args.kwargs['timeout'], ZILLIZ_TIMEOUT) | |
| def test_openai_client_has_bounded_timeout_and_retries(self): | |
| with patch('src.services.OPENAI_API_KEY', 'test-key'): | |
| client = EmbeddingService(Mock())._get_openai_client() | |
| self.addCleanup(client.close) | |
| self.assertEqual(client.max_retries, OPENAI_MAX_RETRIES) | |
| self.assertEqual(client.timeout.connect, 5) | |
| self.assertEqual(client.timeout.read, 60) | |
| def test_openai_timeout_stops_search_after_one_retry(self): | |
| calls = [] | |
| def stalled(request): | |
| calls.append(request.url.path) | |
| raise httpx.ReadTimeout('simulated stall', request=request) | |
| client = OpenAI(api_key='test-key', timeout=OPENAI_TIMEOUT, max_retries=OPENAI_MAX_RETRIES, | |
| http_client=httpx.Client(transport=httpx.MockTransport(stalled))) | |
| self.addCleanup(client.close) | |
| embedding = EmbeddingService(Mock()) | |
| embedding.openai_client = client | |
| zilliz = Mock() | |
| with patch('openai._base_client.time.sleep'): | |
| with self.assertRaises(UpstreamTimeoutError): | |
| SearchService(embedding, zilliz).search_text('galaxy') | |
| self.assertEqual(len(calls), 2) | |
| self.assertTrue(all(path.endswith('/moderations') for path in calls)) | |
| zilliz.assert_not_called() | |
| self.assertEqual(zilliz.mock_calls, []) | |
| if __name__ == '__main__': | |
| unittest.main() | |