Add owned OpenAI endpoint and coordinated native model switching
This commit is contained in:
@@ -0,0 +1,119 @@
|
||||
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_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()
|
||||
Reference in New Issue
Block a user