File size: 3,383 Bytes
a0a9254 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 | """Persistent local inference API, without a game UI, trainer or external model calls."""
import argparse,json,time
from http.server import HTTPServer,BaseHTTPRequestHandler
from pathlib import Path
from urllib.parse import urlsplit
from predictor import DecisionPredictor,unique_object,reject_nonfinite
from schema import validate_request
ROOT=Path(__file__).resolve().parents[1]
def make_handler(engine):
class Handler(BaseHTTPRequestHandler):
def reply(self,status,data):
raw=json.dumps(data,ensure_ascii=False,allow_nan=False).encode()
self.send_response(status);self.send_header('Content-Type','application/json; charset=utf-8')
self.send_header('Content-Length',str(len(raw)));self.send_header('Cache-Control','no-store')
self.end_headers();self.wfile.write(raw)
def do_GET(self):
if urlsplit(self.path).path in ['/', '/api/health']:
self.reply(200,{'service':'NanoJev-Web','version':'1.0.0-web','ready':True,'model':'browser-head-v5','device':str(engine.device),'precision':engine.precision,'provider_calls':0})
else:self.reply(404,{'error':'Not found'})
def do_POST(self):
if urlsplit(self.path).path!='/api/evaluate':return self.reply(404,{'error':'Not found'})
try:
n=int(self.headers.get('Content-Length',0))
if not 0<n<=2_000_000:raise ValueError('Expected 1–2000000 request bytes')
origin=self.headers.get('Origin')
if origin and urlsplit(origin).netloc!=self.headers.get('Host'):
return self.reply(403,{'error':'Cross-origin requests are disabled'})
body=json.loads(self.rfile.read(n),object_pairs_hook=unique_object,parse_constant=reject_nonfinite)
validate_request(body);start=time.perf_counter();result=engine.predict(body,batch_questions=1)
result['execution']['server_evaluation_seconds']=time.perf_counter()-start
self.reply(200,result)
except (ValueError,KeyError,TypeError) as e:self.reply(400,{'error':str(e)})
except Exception:
import traceback;traceback.print_exc()
self.reply(500,{'error':'Local inference failed; no external fallback was used'})
def log_message(self,fmt,*args):print(fmt%args,flush=True)
return Handler
def main():
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--model',type=Path,default=ROOT/'model');p.add_argument('--port',type=int,default=8774)
p.add_argument('--device',choices=['mps','cpu'],default='mps');p.add_argument('--max-length',type=int,default=768)
a=p.parse_args()
if not 1024<=a.port<=65535:p.error('Use an unprivileged port from 1024 to 65535')
# Bind first: an occupied port fails before loading another copy of the model.
server=HTTPServer(('127.0.0.1',a.port),BaseHTTPRequestHandler)
try:
engine=DecisionPredictor(a.model,device_name=a.device,precision='fp32',max_length=a.max_length)
server.RequestHandlerClass=make_handler(engine)
print(json.dumps({'service':'NanoJev-Web','ready':True,'url':f'http://127.0.0.1:{a.port}','device':a.device}),flush=True)
server.serve_forever()
except KeyboardInterrupt:pass
finally:server.server_close()
if __name__=='__main__':main()
|