148 lines
11 KiB
Python
148 lines
11 KiB
Python
import contextlib
|
|
import hashlib
|
|
import http.client
|
|
import json
|
|
import socket
|
|
import tempfile
|
|
import threading
|
|
import time
|
|
import unittest
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from http.server import BaseHTTPRequestHandler,ThreadingHTTPServer
|
|
from unittest.mock import Mock
|
|
from endpoint import Endpoint
|
|
from inference import Scheduler,InferenceError
|
|
|
|
class FakeWorker:
|
|
def __init__(self):self.name=None;self.started=[];self.stopped=[];self.requests=[];self.block=threading.Event();self.block.set()
|
|
def stop(self):
|
|
if self.name:self.stopped.append(self.name)
|
|
self.name=None
|
|
def ensure(self,p):self.name=p['name'];self.started.append(self.name)
|
|
def status(self):return dict(state='ready' if self.name else 'stopped',profile_name=self.name,profile_id=self.name,port=0,error=None,gpus=[])
|
|
def connect(self):return http.client.HTTPConnection('127.0.0.1',self.http.server_port,timeout=3),'internal-test-only'
|
|
|
|
class Upstream(BaseHTTPRequestHandler):
|
|
def log_message(self,*args):pass
|
|
def do_POST(self):
|
|
w=self.server.worker;data=json.loads(self.rfile.read(int(self.headers['Content-Length'])));w.requests.append(data)
|
|
if data.get('stream'):
|
|
self.send_response(200);self.send_header('Content-Type','text/event-stream');self.end_headers();self.wfile.write(b'data: {"choices":[]}\n\n');self.wfile.flush();w.block.wait(3);self.wfile.write(b'data: [DONE]\n\n')
|
|
else:
|
|
body=json.dumps(dict(id='test',model=data['model'],choices=[{'message':{'role':'assistant','content':'synthetic'}}])).encode();self.send_response(200);self.send_header('Content-Length',str(len(body)));self.end_headers();self.wfile.write(body)
|
|
|
|
class EndpointTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.tmp=tempfile.TemporaryDirectory();self.worker=FakeWorker();self.worker.http=ThreadingHTTPServer(('127.0.0.1',0),Upstream);self.worker.http.worker=self.worker;threading.Thread(target=self.worker.http.serve_forever,daemon=True).start()
|
|
self.rows=[dict(id=n,name=n,kind='chat',revision=1,updated_at=1,runnable=True,blockers=[],parameters=dict(slots=2,temperature=.2,top_p=.8,top_k=20)) for n in ('alpha','beta')]
|
|
self.rows.append(dict(id='image',name='image',kind='image',revision=1,updated_at=1,runnable=True,blockers=[],parameters=dict(width=512,height=512)))
|
|
self.rows.append(dict(id='audio',name='audio',kind='audio',revision=1,updated_at=1,runnable=False,blockers=['No worker'],parameters={}))
|
|
self.profiles=SimpleNamespace(status=lambda:dict(profiles=self.rows));self.record={'api_token_hash':hashlib.sha256(b'A'*32).hexdigest()}
|
|
self.images=Mock();self.images.status.return_value=dict(job=None)
|
|
self.ep=Endpoint(self.tmp.name,self.profiles,self.worker,Scheduler(self.worker),self.images,SimpleNamespace(read=lambda:self.record),1)
|
|
with socket.socket() as s:s.bind(('127.0.0.1',0));self.port=s.getsockname()[1]
|
|
self.ep.configure({'port':self.port});self.ep.start()
|
|
def tearDown(self):
|
|
self.worker.block.set();self.ep.close();self.worker.http.shutdown();self.worker.http.server_close();self.tmp.cleanup()
|
|
def request(self,path='/v1/models',data=None,token='A'*32):
|
|
conn=http.client.HTTPConnection('127.0.0.1',self.port,timeout=5);headers={'Content-Type':'application/json','Authorization':'Bearer '+token}
|
|
conn.request('POST' if data is not None else 'GET',path,body=json.dumps(data) if data is not None else None,headers=headers);r=conn.getresponse();body=r.read();conn.close();return r.status,(body if data and data.get('stream') else json.loads(body))
|
|
def enable(self,name):self.ep.enable(dict(id=name,enabled=True))
|
|
def test_explicit_publication_auth_rotation_and_missing_audio(self):
|
|
self.assertEqual(self.request()[1]['data'],[]);self.assertEqual(self.request(token='bad')[0],401)
|
|
self.enable('alpha');self.assertEqual([p['id'] for p in self.request()[1]['data']],['alpha']);self.assertFalse(self.worker.started)
|
|
with self.assertRaises(ValueError):self.enable('audio')
|
|
self.assertEqual(self.request('/v1/audio/speech',{})[0],501)
|
|
self.record['api_token_hash']=hashlib.sha256(b'B'*32).hexdigest();self.assertEqual(self.request()[0],401);self.assertEqual(self.request(token='B'*32)[0],200)
|
|
def test_tts_endpoint_returns_wav_and_requires_publication(self):
|
|
self.rows[-1].update(runnable=True,blockers=[])
|
|
self.ep.tts=Mock();self.ep.tts.start.return_value={'id':'a'*32};self.ep.tts.status.return_value={'job':{'id':'a'*32,'state':'complete'}};self.ep.tts.audio.return_value=b'RIFF-test-WAVE'
|
|
self.assertEqual(self.request('/v1/audio/speech',{'model':'audio','input':'hello'})[0],404)
|
|
self.enable('audio')
|
|
self.assertEqual(self.request('/v1/audio/speech',{'model':'audio','input':'hello','response_format':'mp3'})[0],400)
|
|
c=http.client.HTTPConnection('127.0.0.1',self.port,timeout=5);c.request('POST','/v1/audio/speech',json.dumps({'model':'audio','input':'hello','voice':'Ryan','response_format':'wav'}),{'Content-Type':'application/json','Authorization':'Bearer '+'A'*32});r=c.getresponse()
|
|
self.assertEqual(r.status,200);self.assertEqual(r.getheader('Content-Type'),'audio/wav');self.assertEqual(r.read(),b'RIFF-test-WAVE');c.close()
|
|
self.assertEqual(self.ep.tts.start.call_args.kwargs['speaker'],'Ryan')
|
|
def test_single_image_selection_preserves_multiple_llms(self):
|
|
self.rows.append(dict(self.rows[2],id='image2',name='image2'))
|
|
self.enable('alpha');self.enable('beta');self.enable('image');self.enable('image2')
|
|
self.assertEqual(set(self.ep.config['enabled_profiles']),{'alpha','beta','image2'})
|
|
self.assertEqual({r['id'] for r in self.request()[1]['data']},{'alpha','beta','image2'})
|
|
self.ep.enable(dict(id='image2',enabled=False));self.assertEqual(set(self.ep.config['enabled_profiles']),{'alpha','beta'})
|
|
def test_invalid_image_selection_keeps_previous(self):
|
|
self.rows.append(dict(self.rows[2],id='badimage',name='badimage',runnable=False,blockers=['missing']))
|
|
self.enable('image')
|
|
with self.assertRaises(ValueError):self.enable('badimage')
|
|
self.assertEqual(self.ep.config['enabled_profiles'],['image'])
|
|
def test_vision_needs_projector_and_rejects_remote_urls(self):
|
|
self.enable('alpha')
|
|
data=dict(model='alpha',messages=[{'role':'user','content':[{'type':'image_url','image_url':{'url':'https://example.com/a.png'}}]}])
|
|
self.assertEqual(self.request('/v1/chat/completions',data)[0],400)
|
|
self.rows[0]['parameters']['vision_projector']='synthetic'
|
|
self.assertEqual(self.request('/v1/chat/completions',data)[0],400)
|
|
data['messages'][0]['content'][0]['image_url']['url']='data:image/png;base64,aGVsbG8='
|
|
self.assertEqual(self.request('/v1/chat/completions',data)[0],200)
|
|
def test_profile_switch_and_sampling_defaults(self):
|
|
for name in ('alpha','beta'):self.enable(name)
|
|
for name in ('alpha','beta'):
|
|
status,data=self.request('/v1/chat/completions',dict(model=name,messages=[dict(role='user',content='test')]))
|
|
self.assertEqual(status,200);self.assertEqual(data['model'],name)
|
|
self.assertEqual(self.worker.stopped,['alpha']);self.assertEqual(self.worker.requests[-1]['temperature'],.2)
|
|
self.assertEqual(self.ep.status()['counts']['llm']['enabled'],2)
|
|
def test_unknown_or_unsupported_request_does_not_load(self):
|
|
self.enable('alpha')
|
|
for req in [dict(model='unknown',messages=[{}]),dict(model='alpha',messages=[dict(content=[dict(type='image_url')])]),dict(model='alpha',messages=[{}],cache_file='/tmp/foo')]:
|
|
self.assertGreaterEqual(self.request('/v1/chat/completions',req)[0],400)
|
|
self.assertFalse(self.worker.started)
|
|
def test_stream_lease_blocks_switch_until_stream_finishes(self):
|
|
self.enable('alpha');self.enable('beta');self.worker.block.clear()
|
|
conn=http.client.HTTPConnection('127.0.0.1',self.port,timeout=5)
|
|
conn.request('POST','/v1/chat/completions',json.dumps(dict(model='alpha',messages=[{}],stream=True)),headers={'Content-Type':'application/json','Authorization':'Bearer '+'A'*32});r=conn.getresponse();self.assertEqual(r.status,200)
|
|
result=[];thread=threading.Thread(target=lambda:result.append(self.request('/v1/chat/completions',dict(model='beta',messages=[{}]))));thread.start()
|
|
deadline=time.monotonic()+2
|
|
while not self.ep.scheduler.status()['waiting_requests'] and time.monotonic()<deadline:time.sleep(.01)
|
|
self.assertEqual(self.worker.name,'alpha');self.assertEqual(self.worker.stopped,[])
|
|
self.worker.block.set();self.assertIn(b'[DONE]',r.read());conn.close();thread.join(3);self.assertEqual(result[0][0],200);self.assertEqual(self.worker.name,'beta')
|
|
def test_port_conflict_and_graceful_stop(self):
|
|
with self.assertRaises(ValueError):self.ep.configure(dict(port=self.port+1))
|
|
self.ep.stop();deadline=time.monotonic()+3
|
|
while self.ep.state!='stopped' and time.monotonic()<deadline:time.sleep(.02)
|
|
self.assertEqual(self.ep.state,'stopped')
|
|
with socket.socket() as sock:
|
|
sock.bind(('127.0.0.1',self.port))
|
|
with self.assertRaises(ValueError):self.ep.start()
|
|
self.ep.start();self.assertTrue(self.ep.status()['reachable'])
|
|
def test_image_unloads_chat_and_returns_openai_shape(self):
|
|
self.enable('alpha');self.enable('image');self.request('/v1/chat/completions',dict(model='alpha',messages=[{}]))
|
|
self.images.start.return_value={'id':'new'};self.images.status.return_value={'job':{'id':'new','state':'complete'}};self.images.image.return_value=b'synthetic-png'
|
|
status,data=self.request('/v1/images/generations',dict(model='image',prompt='synthetic',response_format='b64_json'))
|
|
self.assertEqual(status,200);self.assertIn('b64_json',data['data'][0]);self.assertIsNone(self.worker.name);self.images.start.assert_called_once_with('image','synthetic',reserved=True)
|
|
def test_disabled_and_edited_queued_profiles_are_rejected(self):
|
|
self.enable('alpha');self.enable('beta');self.worker.block.clear()
|
|
first=threading.Thread(target=lambda:self.request('/v1/chat/completions',dict(model='alpha',messages=[{}],stream=True)));first.start()
|
|
deadline=time.monotonic()+2
|
|
while not self.worker.requests and time.monotonic()<deadline:time.sleep(.01)
|
|
result=[];second=threading.Thread(target=lambda:result.append(self.request('/v1/chat/completions',dict(model='beta',messages=[{}]))));second.start()
|
|
deadline=time.monotonic()+2
|
|
while not self.ep.scheduler.status()['waiting_requests'] and time.monotonic()<deadline:time.sleep(.01)
|
|
self.ep.enable(dict(id='beta',enabled=False));self.worker.block.set();first.join(4);second.join(4);self.assertEqual(result[0][0],503);self.assertNotIn('beta',self.worker.started)
|
|
|
|
class SchedulerTests(unittest.TestCase):
|
|
def test_parallel_slots_and_fifo_switch(self):
|
|
w=FakeWorker();s=Scheduler(w)
|
|
with s.lease('a',2):
|
|
with s.lease('a',2):self.assertEqual(s.status()['active_requests'],2)
|
|
with self.assertRaises(InferenceError):
|
|
with s.lease('b',timeout=0):pass
|
|
with s.lease('b'):self.assertEqual(s.status()['active_requests'],1)
|
|
self.assertEqual(s.status()['active_requests'],0)
|
|
def test_failed_prepare_releases_queue_and_can_retry(self):
|
|
s=Scheduler(FakeWorker())
|
|
with self.assertRaises(ValueError):
|
|
with s.lease('bad',prepare=lambda:(_ for _ in ()).throw(ValueError('bad'))):pass
|
|
with s.lease('good'):self.assertFalse(s.status()['switching'])
|
|
self.assertEqual(s.status()['active_requests'],0)
|
|
|
|
if __name__=='__main__':unittest.main()
|