Files
Athena-Deck/test_endpoint.py
T

139 lines
9.9 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_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()