41 lines
2.2 KiB
Python
41 lines
2.2 KiB
Python
"""Search request validation and published-only retrieval regressions."""
|
|
from unittest import TestCase, mock
|
|
import test_knowledge as fixtures
|
|
from knowledge_retriever import KnowledgeRetriever
|
|
|
|
class KnowledgeSearchApiTest(TestCase):
|
|
setUp = fixtures.KnowledgeApiTest.setUp
|
|
tearDown = fixtures.KnowledgeApiTest.tearDown
|
|
|
|
def search(self, body):
|
|
return self.client.post('/api/v2/knowledge/search?account_id=0', headers=self.headers, json=body)
|
|
|
|
def test_invalid_query_never_reaches_retriever(self):
|
|
with mock.patch.object(KnowledgeRetriever, 'search') as retrieve:
|
|
for value in ['', ' ', '\t\n', '约', ' 约 ', 'a' * 2001, None, {}, 123]:
|
|
with self.subTest(value_type=type(value).__name__, length=len(value) if isinstance(value,str) else None):
|
|
response = self.search({'query': value})
|
|
self.assertEqual(response.status_code, 422)
|
|
self.assertIsInstance(response.json()['detail'], list)
|
|
self.assertEqual(self.search({}).status_code, 422)
|
|
retrieve.assert_not_called()
|
|
|
|
def test_valid_query_trims_before_validation_and_search(self):
|
|
with mock.patch.object(KnowledgeRetriever, 'search', return_value={'hits':[], 'mode':'lexical','elapsed_ms':0}) as retrieve:
|
|
for text in ['预约', '🙂🙂', 'a' * 2000]:
|
|
response = self.search({'query': ' \t'+text+'\n '})
|
|
self.assertEqual(response.status_code, 200, response.text)
|
|
self.assertEqual(retrieve.call_args.args[-1], text)
|
|
|
|
def test_no_published_knowledge_is_empty_success(self):
|
|
with mock.patch('knowledge_retriever.VectorIndex') as index:
|
|
index.return_value.configured = False
|
|
response = self.search({'query': '预约挂号'})
|
|
self.assertEqual(response.status_code, 200, response.text)
|
|
self.assertEqual(response.json()['hits'], [])
|
|
|
|
def test_authentication_and_specific_account_still_required(self):
|
|
url = '/api/v2/knowledge/search'
|
|
self.assertEqual(self.client.post(url, json={'query':'预约'}).status_code, 401)
|
|
self.assertEqual(self.client.post(url+'?account_id=-1', headers=self.headers, json={'query':'预约'}).status_code, 400)
|