WIlfLin's picture
Add opt-in shared-material reranking with measured multi-question cache reuse
2183ffa verified
Raw History Blame Contribute Delete
5.51 kB
"""CPU contract tests for isolation, validation, cache evidence and lifecycle."""
import asyncio
import time
import unittest
import httpx
from pydantic import ValidationError
from shared_material import Backend, Material, MaterialStore, Question, QuestionsRequest, create_app
class Tokenizer:
def encode(self, text, **kwargs):
if text in ['no','yes','\n']:
return {'no':[11],'yes':[12],'\n':[13]}[text]
return [ord(c)+100 for c in text]
def apply_chat_template(self, messages, **kwargs):
text=''.join(f"<{m['role']}>{m['content']}</turn>" for m in messages)
if kwargs.get('add_generation_prompt'):
text+='<assistant>'
return self.encode(text)
class Contracts(unittest.TestCase):
def test_store_lru_expiry_and_delete(self):
store=MaterialStore(limit=2,ttl_s=100)
def m(key,age=0):return Material(key,(1,),1,0,time.monotonic()-age)
store.put(m('a'));store.put(m('b'));store.get('a');store.put(m('c'))
with self.assertRaises(KeyError):store.get('b')
store.put(m('old',101))
with self.assertRaises(KeyError):store.get('old')
store.delete('c')
with self.assertRaises(KeyError):store.get('c')
def test_request_bounds(self):
with self.assertRaises(ValidationError):Question(instructions='x',options=['yes',' '])
with self.assertRaises(ValidationError):QuestionsRequest(questions={})
q=Question(instructions='x',options=['a','b','c'])
with self.assertRaises(ValidationError):QuestionsRequest(questions={str(i):q for i in range(43)})
with self.assertRaises(ValidationError):QuestionsRequest(questions={'q':q},concurrency=5)
class APIContracts(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.calls=[]
async def transport(request):
import json
if request.url.path=='/health':return httpx.Response(200,json={})
body=json.loads(request.content);self.calls.append(body)
ids=body['prompt'];cached=len(ids)//800*800
return httpx.Response(200,json={'choices':[{'logprobs':{'top_logprobs':[{'token_id:11':-2.0,'token_id:12':-.1}]}}],
'usage':{'prompt_tokens':len(ids),'prompt_tokens_details':{'cached_tokens':cached}}})
self.engine=httpx.AsyncClient(transport=httpx.MockTransport(transport))
self.backend=Backend(Tokenizer(),client=self.engine)
self.api=httpx.AsyncClient(transport=httpx.ASGITransport(app=create_app(self.backend)),base_url='http://test')
async def asyncTearDown(self):
await self.api.aclose();await self.engine.aclose()
async def test_exact_prefix_and_separate_salts(self):
a=self.backend.material('blue parcel');b=self.backend.material('red parcel')
req=QuestionsRequest(questions={'x':Question(instructions='color?',options=['blue','red']),
'y':Question(instructions='other?',options=['yes','no'])})
compiled=self.backend.compile(a,req.questions)
self.assertEqual(len(a.prefix)%800,0)
self.assertTrue(all(tuple(ids[:len(a.prefix)])==a.prefix for _,_,ids in compiled))
self.assertNotEqual(a.id,b.id);self.assertNotEqual(a.prefix,b.prefix)
self.assertEqual(a.content_tokens+a.padding_tokens,len(a.prefix))
async def test_invalid_suffix_submits_nothing(self):
prepared=(await self.api.post('/v1/materials',json={'state':'blue'})).json()
before=len(self.calls)
response=await self.api.post('/v1/materials/'+prepared['material_id']+'/questions',json={
'questions':{'ok':{'instructions':'color','options':['blue','red']},
'too_long':{'instructions':'x'*5000,'options':['a','b']}}})
self.assertEqual(response.status_code,422);self.assertEqual(len(self.calls),before)
async def test_id_append_delete_and_ephemeral_cleanup(self):
p=(await self.api.post('/v1/materials',json={'state':'blue'})).json();key=p['material_id']
for prompt in ['color?','another question?']:
r=await self.api.post(f'/v1/materials/{key}/questions',json={'questions':{'q':{'instructions':prompt,'options':['blue','red']}}})
self.assertEqual(r.status_code,200);self.assertTrue(r.json()['metrics']['all_candidates_reused_entire_material'])
self.assertTrue(all(c['cache_salt']==key for c in self.calls))
self.assertEqual((await self.api.delete('/v1/materials/'+key)).status_code,200)
r=await self.api.post(f'/v1/materials/{key}/questions',json={'questions':{'q':{'instructions':'x','options':['a','b']}}})
self.assertEqual(r.status_code,404)
r=await self.api.post('/v1/rerank',json={'state':'blue','questions':{'q':{'instructions':'color','options':['blue','red']}}})
self.assertEqual(r.status_code,200);self.assertEqual(len(self.backend.store.items),0)
async def test_missing_cache_evidence_fails_closed(self):
async def bad(request):
return httpx.Response(200,json={'choices':[{'logprobs':{'top_logprobs':[{'token_id:11':-1,'token_id:12':-1}]}}],
'usage':{'prompt_tokens':1}})
async with httpx.AsyncClient(transport=httpx.MockTransport(bad)) as client:
backend=Backend(Tokenizer(),client=client)
with self.assertRaisesRegex(ValueError,'cached_tokens'):
await backend.complete([99],'salt')
if __name__=='__main__':unittest.main()