Spaces:
Running
Running
Download tests/test_context_manager.py from dromerosm/groq-chatbot: direct link, hf CLI and curl.
- Browser
- Download file 6.54 kB
-
https://huggingface.co/spaces/dromerosm/groq-chatbot/resolve/main/tests/test_context_manager.py
- Command line
-
hf download hf://spaces/dromerosm/groq-chatbot/tests/test_context_manager.py
-
curl -L -o test_context_manager.py https://huggingface.co/spaces/dromerosm/groq-chatbot/resolve/main/tests/test_context_manager.py
6.54 kB
| import copy | |
| import io | |
| import unittest | |
| from pypdf import PdfWriter | |
| from pypdf.generic import DictionaryObject, NameObject, DecodedStreamObject | |
| from context_manager import (index_document, load_document, prepare_request, prompt_tokens, | |
| rank_excerpts, estimate_tokens, MAX_DOCUMENT_BYTES) | |
| def synthetic_pdf(): | |
| writer = PdfWriter() | |
| for number in range(1, 13): | |
| page = writer.add_blank_page(width=600, height=800) | |
| font = DictionaryObject({NameObject('/Type'): NameObject('/Font'), | |
| NameObject('/Subtype'): NameObject('/Type1'), | |
| NameObject('/BaseFont'): NameObject('/Helvetica')}) | |
| page[NameObject('/Resources')] = DictionaryObject({NameObject('/Font'): DictionaryObject({NameObject('/F1'): font})}) | |
| text = ('Routine warehouse inspection and inventory notes. ' * 70 if number < 12 | |
| else 'Project Aurora release code is ZX731. Release date is 17 November 2027.') | |
| stream = DecodedStreamObject() | |
| stream.set_data(f'BT /F1 10 Tf 30 700 Td ({text}) Tj ET'.encode()) | |
| page[NameObject('/Contents')] = writer._add_object(stream) | |
| result = io.BytesIO() | |
| writer.write(result) | |
| return result.getvalue() | |
| class ContextTests(unittest.TestCase): | |
| def setUp(self): | |
| self.model = {'context_window': 4096, 'max_completion_tokens': 2048, | |
| 'rate_limits': {'tpm': 6000}} | |
| def test_real_pdf_last_page_is_retrieved_instead_of_initial_prefix(self): | |
| document = load_document(synthetic_pdf(), 'synthetic.pdf') | |
| self.assertEqual(len(document['pages']), 12) | |
| self.assertGreater(estimate_tokens(' '.join(document['pages'][:11])), 1200) | |
| ranked = rank_excerpts(document, 'What is Project Aurora release code?') | |
| self.assertEqual(ranked[0]['page'], 12) | |
| request = prepare_request([], 'What is Project Aurora release code?', self.model, 1024, document) | |
| self.assertIn('ZX731', request['messages'][0]['content']) | |
| self.assertEqual(request['sources'][0]['location'], 'page 12') | |
| self.assertFalse(request['no_document_match']) | |
| self.assertLessEqual(request['prompt_tokens'] + request['max_tokens'] + request['margin'], 4096) | |
| def test_long_history_is_trimmed_as_complete_turns_without_mutation(self): | |
| history = [] | |
| for i in range(100): | |
| history += [{'role': 'user', 'content': f'Question {i}: ' + 'example ' * 80}, | |
| {'role': 'assistant', 'content': f'Answer {i}: ' + 'response ' * 80}] | |
| original = copy.deepcopy(history) | |
| request = prepare_request(history, 'What did we discuss most recently?', self.model, 1024) | |
| self.assertGreater(request['omitted_messages'], 0) | |
| self.assertEqual(request['omitted_messages'] % 2, 0) | |
| self.assertEqual(request['messages'][0]['role'], 'user') | |
| self.assertIn('Answer 99', request['messages'][-2]['content']) | |
| self.assertEqual(history, original) | |
| self.assertLessEqual(prompt_tokens(request['messages']) + request['max_tokens'] + request['margin'], 4096) | |
| def test_tpm_and_output_and_context_limits_are_independent(self): | |
| for context, tpm, output_cap, requested in [(8192, 2000, 4000, 3000), (1024, 8000, 100, 500), (4096, None, 2048, 9000)]: | |
| model = {'context_window': context, 'max_completion_tokens': output_cap, 'rate_limits': {'tpm': tpm}} | |
| request = prepare_request([], 'Hello', model, requested) | |
| self.assertLessEqual(request['max_tokens'], output_cap) | |
| self.assertLessEqual(request['max_tokens'], requested) | |
| self.assertLessEqual(request['prompt_tokens'] + request['max_tokens'] + request['margin'], min(context, tpm or context)) | |
| def test_oversized_question_is_rejected_without_truncation(self): | |
| with self.assertRaisesRegex(ValueError, 'too large'): | |
| prepare_request([], 'question ' * 10000, self.model, 512) | |
| def test_image_allowance_is_counted_and_payload_preserved(self): | |
| content = [{'type': 'text', 'text': 'Describe this image'}, | |
| {'type': 'image_url', 'image_url': {'url': 'data:image/png;base64,abc'}}] | |
| model = {'context_window': 16000, 'max_completion_tokens': 2000} | |
| request = prepare_request([], content, model, 512) | |
| self.assertGreater(request['prompt_tokens'], 4096) | |
| self.assertEqual(request['messages'][-1]['content'], content) | |
| with self.assertRaises(ValueError): | |
| prepare_request([], content, self.model, 512) | |
| def test_no_match_is_explicit_and_does_not_inject_irrelevant_prefix(self): | |
| doc = index_document(['Warehouse inventory records.'], 'inventory.txt') | |
| request = prepare_request([], 'Meteorite tungsten isotope?', self.model, 512, doc) | |
| self.assertTrue(request['no_document_match']) | |
| self.assertEqual(request['sources'], []) | |
| def test_txt_and_markdown_keep_end_and_use_section_labels(self): | |
| for name in ['notes.txt', 'notes.md']: | |
| data = ('Routine inventory. ' * 10000 + '\nAurora release code ZX731').encode() | |
| doc = load_document(data, name) | |
| self.assertIn('ZX731', doc['pages'][0]) | |
| request = prepare_request([], 'Aurora release code?', self.model, 512, doc) | |
| self.assertTrue(any('ZX731' in s['text'] for s in request['sources'])) | |
| self.assertTrue(request['sources'][0]['location'].startswith('section')) | |
| def test_followup_can_retrieve_subject_from_previous_question(self): | |
| doc = index_document(['Aurora release code ZX731 and delivery on Tuesday.'], 'notes.txt') | |
| request = prepare_request([{'role': 'user', 'content': 'Tell me about Aurora release.'}], | |
| 'And when?', self.model, 512, doc) | |
| self.assertFalse(request['no_document_match']) | |
| def test_blank_and_oversized_documents_are_rejected(self): | |
| with self.assertRaisesRegex(ValueError, 'No readable text'): | |
| load_document(b'', 'empty.txt') | |
| with self.assertRaisesRegex(ValueError, '10 MB'): | |
| load_document(b'x' * (MAX_DOCUMENT_BYTES + 1), 'large.txt') | |
| def test_missing_metadata_uses_bounded_fallback(self): | |
| request = prepare_request([], 'Hello', {}, 100000) | |
| self.assertLessEqual(request['max_tokens'], 4000) | |
| self.assertLessEqual(request['prompt_tokens'] + request['max_tokens'] + request['margin'], 8192) | |
| if __name__ == '__main__': | |
| unittest.main() | |