Files
Athena-Deck/test_endpoint.py
T

243 lines
19 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,chat_upstream_error
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 UpstreamErrorTests(unittest.TestCase):
def test_context_limit_is_actionable_without_echoing_upstream_text(self):
raw=json.dumps({'error':{'type':'exceed_context_size_error','message':'private request text',
'n_prompt_tokens':13523,'n_ctx':8192}}).encode()
error=chat_upstream_error(400,raw)
self.assertEqual(error.code,'context_length_exceeded')
self.assertIn('13.523',str(error));self.assertIn('8.192',str(error))
self.assertNotIn('private request text',str(error))
def test_unknown_upstream_errors_stay_generic(self):
error=chat_upstream_error(400,b'{"error":{"message":"private request text"}}')
self.assertEqual(error.code,'upstream_error')
self.assertNotIn('private request text',str(error))
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 test_music_discovery_and_authenticated_wav_route(self):
self.rows.append(dict(id='music',name='Music',kind='music',runnable=True,blockers=[],parameters={},updated_at=1))
self.ep.music=Mock();self.ep.music.start.return_value=dict(id='a'*32);self.ep.music.status.return_value=dict(job=dict(id='a'*32,state='complete'));self.ep.music.audio.return_value=b'RIFFsynthetic-WAV'
self.enable('music')
self.assertEqual(self.request('/v1/audio/music/models')[1]['data'][0]['id'],'Music')
self.assertNotIn('Music',[x['id'] for x in self.request()[1]['data']])
conn=http.client.HTTPConnection('127.0.0.1',self.port,timeout=5)
conn.request('POST','/v1/audio/music',json.dumps(dict(model='Music',lyrics='Synthetic',style='pop')),{'Content-Type':'application/json','Authorization':'Bearer '+'A'*32})
response=conn.getresponse();self.assertEqual(response.status,200);self.assertEqual(response.getheader('Content-Type'),'audio/wav');self.assertEqual(response.read(),b'RIFFsynthetic-WAV');conn.close()
self.assertEqual(self.request('/v1/audio/music',dict(model='Music',lyrics='Synthetic',style='pop',unsupported=True))[0],400)
self.assertEqual(self.request('/v1/audio/music/models',token='bad')[0],401)
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_model_lists_separate_modalities(self):
self.rows[-1].update(runnable=True,blockers=[])
self.rows.append(dict(id='stt',name='stt',kind='stt',runnable=True,blockers=[],parameters={},updated_at=1))
for name in ('alpha','image','audio','stt'):self.enable(name)
for route,expected in [('/v1/models','alpha'),('/v1/images/models','athena-image'),('/v1/audio/speech/models','audio'),('/v1/audio/transcriptions/models','stt')]:
self.assertEqual([p['id'] for p in self.request(route)[1]['data']],[expected])
self.assertEqual(self.request(route,token='bad')[0],401)
self.rows[-1]['runnable']=False
self.assertEqual(self.request('/v1/audio/transcriptions/models')[1]['data'],[])
def test_reasoning_effort_forwarded_and_validated(self):
self.enable('alpha')
for effort in ('none','minimal','low','medium','high','xhigh','max','ultra',None):
status,_=self.request('/v1/chat/completions',dict(model='alpha',messages=[dict(role='user',content='synthetic')],reasoning_effort=effort))
self.assertEqual(status,200)
self.assertEqual(self.worker.requests[-1].get('reasoning_effort'),{'minimal':'low','max':'xhigh','ultra':'xhigh'}.get(effort,effort))
for effort in ([],True,3,'invalid'):
self.assertEqual(self.request('/v1/chat/completions',dict(model='alpha',messages=[dict(role='user',content='synthetic')],reasoning_effort=effort))[0],400)
def test_video_profiles_not_published(self):
self.rows.append(dict(id='video-old',name='old',kind='video',runnable=True,blockers=[],parameters={},updated_at=1))
self.assertFalse(any(p['kind']=='video' for p in self.ep.rows()))
with self.assertRaises(ValueError):self.ep.enable({'id':'video-old','enabled':True})
def test_video_mode_rejects_chat_without_starting_llama(self):
self.ep.video=Mock();self.ep.video.status.return_value={'service':{'ready':False}};self.ep.scheduler.gpu_mode='video'
code,body=self.request('/v1/chat/completions',dict(model='alpha',messages=[dict(role='user',content='synthetic')]))
self.assertEqual(code,503);self.assertEqual(body['error']['code'],'video_not_ready');self.assertFalse(self.worker.started)
self.ep.video=None
def test_video_mode_routes_original_api_even_on_models_path(self):
from unittest.mock import patch
self.ep.video=Mock();self.ep.video.status.return_value={'service':{'ready':True},'selected':'ltx'};self.ep.scheduler.gpu_mode='video'
def forward(handler,service):
self.assertEqual(service,'ltx');handler.send({'ltx_native':handler.path},207)
with patch('endpoint.relay',side_effect=forward):
code,body=self.request('/v1/models');self.assertEqual(code,207);self.assertEqual(body,{'ltx_native':'/v1/models'})
self.ep.scheduler.gpu_mode='llm';self.ep.video=None
def test_comfy_browser_session_only_in_video_and_same_origin(self):
class NativeVideo:
def status(self):return {'service':{'ready':True}}
def relay(self,handler):handler.send({'native':True})
self.ep.video=NativeVideo();self.ep.scheduler.gpu_mode='video';self.ep.video_browser_auth=lambda h:h.get('Cookie')=='deck_session=synthetic'
for origin,expected in [(f'http://127.0.0.1:{self.port}',200),('https://foreign.example',403)]:
conn=http.client.HTTPConnection('127.0.0.1',self.port)
conn.request('GET','/',headers={'Cookie':'deck_session=synthetic','Origin':origin})
response=conn.getresponse();self.assertEqual(response.status,expected);response.read();conn.close()
conn=http.client.HTTPConnection('127.0.0.1',self.port);conn.request('GET','/',headers={'Cookie':'deck_session=wrong'})
response=conn.getresponse();self.assertEqual(response.status,401);response.read();conn.close()
self.ep.scheduler.gpu_mode='llm';self.ep.video=None
def test_video_api_redirect_is_explicit_not_translation(self):
code,body=self.request('/v1/videos',{'prompt':'synthetic'})
self.assertEqual(code,410);self.assertEqual(body['error']['code'],'use_ltx_api')
def test_stt_endpoint_multipart_and_publication(self):
from test_stt import audio
self.rows.append(dict(id='stt',name='stt',kind='stt',runnable=True,blockers=[],parameters={}))
self.ep.stt=Mock();self.ep.stt.start.return_value={'id':'s'};self.ep.stt.status.return_value={'job':dict(id='s',state='complete',text='Testaufnahme')}
body=b'--x\r\nContent-Disposition: form-data; name="file"; filename="audio.wav"\r\n\r\n'+audio()+b'\r\n--x\r\nContent-Disposition: form-data; name="model"\r\n\r\nstt\r\n--x--\r\n'
def request():
c=http.client.HTTPConnection('127.0.0.1',self.port,timeout=5);c.request('POST','/v1/audio/transcriptions',body,{'Content-Type':'multipart/form-data; boundary=x','Authorization':'Bearer '+'A'*32});r=c.getresponse();result=r.status,json.loads(r.read());c.close();return result
self.assertEqual(request()[0],404)
self.enable('stt');self.assertEqual(request(),(200,{'text':'Testaufnahme'}))
self.assertEqual(self.ep.stt.start.call_args.args[0],'stt')
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'})
self.ep.enable(dict(id='image2',enabled=False));self.assertEqual(set(self.ep.config['enabled_profiles']),{'alpha','beta'})
def test_image_alias_follows_selection_and_profile_size(self):
self.rows[2]['parameters']['height']=512
self.rows.append(dict(self.rows[2],id='image2',name='image2'))
self.ep.images.start.return_value={'id':'job'};self.ep.images.status.return_value={'job':{'id':'job','state':'complete'}};self.ep.images.image.return_value=b'png'
self.assertEqual(self.request('/v1/images/generations',dict(model='athena-image',prompt='synthetic'))[0],404)
for name in ('image','image2'):
self.enable(name)
status,result=self.request('/v1/images/generations',dict(model='athena-image',prompt='synthetic',size='1536x1024'))
self.assertEqual(status,200);self.assertEqual(result['athena_deck']['size'],'512x512');self.assertEqual(self.ep.images.start.call_args.args[0],name)
self.assertEqual([m['id'] for m in self.request('/v1/images/models')[1]['data']],['athena-image'])
self.assertEqual(self.request('/v1/images/generations',dict(model='athena-image',prompt='synthetic',size=[]))[0],400)
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()