groq-chatbot / tests /test_context_manager.py
dromerosm's picture
Bound conversation context and retrieve document passages across all pages
3e56df9
Raw History Blame Contribute Delete
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()