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