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_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'),'max' if effort=='ultra' else 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.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_mode_active');self.assertFalse(self.worker.started) 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()