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()