AION-Search / tests /test_request_validation.py
astronolan's picture
Validate callback state and bound upstream waits
980f681
Raw History Blame Contribute Delete
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()