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