File size: 18,520 Bytes
cc3f990
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9d0e73
cc3f990
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9d0e73
cc3f990
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
"""Loopback-only browser playground for a fixed catalog of trusted checkpoints."""
import argparse
from concurrent.futures import ThreadPoolExecutor
import gc
import hashlib
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import json
from pathlib import Path
import re
import signal
import subprocess
import sys
import threading
import time
from urllib.parse import urlsplit

ROOT = Path(__file__).resolve().parent
MAX_BODY_BYTES = 128 * 1024
ASSETS = {'/': ('playground.html', 'text/html; charset=utf-8'),
          '/playground.html': ('playground.html', 'text/html; charset=utf-8'),
          '/playground.css': ('playground.css', 'text/css; charset=utf-8'),
          '/playground.js': ('playground.js', 'text/javascript; charset=utf-8')}
CSP = ("default-src 'none'; script-src 'self'; style-src 'self'; connect-src 'self'; "
       "img-src 'self' data:; base-uri 'none'; frame-ancestors 'none'; form-action 'none'")


class RequestError(Exception):
    def __init__(self, status, code, message):
        super().__init__(message)
        self.status, self.code, self.message = status, code, message


def checksum(path):
    value = hashlib.sha256()
    with Path(path).open('rb') as handle:
        for chunk in iter(lambda: handle.read(1024 * 1024), b''):
            value.update(chunk)
    return value.hexdigest()


def load_catalog(path, max_tokens=1024):
    if not isinstance(max_tokens, int) or isinstance(max_tokens, bool) or not 1 <= max_tokens <= 32768:
        raise ValueError('Invalid inference token limit')
    data = json.loads(Path(path).read_text())
    models = data.get('models')
    if not isinstance(models, list) or not models:
        raise ValueError('Model catalog must contain a nonempty models list')
    catalog = {}
    for model in models:
        if not isinstance(model, dict) or not re.fullmatch(r'[A-Za-z0-9][A-Za-z0-9_.-]{0,63}', model.get('id', '')):
            raise ValueError('Invalid catalog model ID')
        if model['id'] in catalog:
            raise ValueError('Duplicate catalog model ID')
        if any(not isinstance(model.get(key), str) or not model[key].strip() for key in ('label', 'description', 'checkpoint')):
            raise ValueError('Catalog needs a label, description and checkpoint')
        if type(model.get('calibrated')) is not bool or type(model.get('checkpoint_step')) is not int or model['checkpoint_step'] < 0:
            raise ValueError('Catalog needs explicit calibration and checkpoint-step metadata')
        checkpoint = Path(model['checkpoint'])
        if not checkpoint.is_absolute() or checkpoint.is_symlink() or not checkpoint.is_file():
            raise ValueError('Catalog checkpoint must be an absolute regular file')
        actual_hash = checksum(checkpoint)
        if model.get('sha256', actual_hash) != actual_hash:
            raise ValueError('Catalog checkpoint checksum mismatch')
        catalog[model['id']] = {key: model[key] for key in ('id', 'label', 'description', 'calibrated', 'checkpoint_step')}
        catalog[model['id']].update(checkpoint=str(checkpoint.resolve()), sha256=actual_hash, max_tokens=max_tokens)
    default = data.get('default_model', next(iter(catalog)))
    if default not in catalog:
        raise ValueError('Default model is absent from catalog')
    return catalog, default


def local_backend(checkpoint, device, max_tokens):
    from jev_harness import LocalBackend
    return LocalBackend(checkpoint, device=device, max_tokens=max_tokens)


def memory_preflight(device):
    Path('/proc/self/oom_score_adj').write_text('0')
    subprocess.run(['free', '-b'], check=True, capture_output=True, text=True, timeout=10)
    available = int(next(line.split()[1] for line in Path('/proc/meminfo').read_text().splitlines()
                         if line.startswith('MemAvailable:'))) * 1024
    if device == 'cuda':
        # Inspect all GPU allocations; coexistence is allowed when RAM is sufficient.
        subprocess.run(['nvidia-smi', '--query-compute-apps=pid,process_name,used_memory', '--format=csv'],
                       check=True, capture_output=True, text=True, timeout=10)
    if available < 24 * 2**30:
        raise RequestError(503, 'low_memory', 'At least 24 GiB of available system memory is required to load a model.')


def release_memory(device):
    gc.collect()
    torch = sys.modules.get('torch')
    if device == 'cuda' and torch is not None and torch.cuda.is_initialized():
        torch.cuda.empty_cache()


def release_exception_frames(error):
    # Future/HTTP error objects must not retain failed scorer frames or tensors.
    error.__traceback__ = None
    error.__context__ = None
    error.__cause__ = None


class ModelService:
    def __init__(self, catalog, default_model, device='cuda', backend_factory=local_backend,
                 preflight=memory_preflight, cleanup=release_memory):
        self.catalog, self.default_model, self.device = catalog, default_model, device
        self.backend_factory, self.preflight, self.cleanup = backend_factory, preflight, cleanup
        self.lock, self.state_lock = threading.Lock(), threading.Lock()
        self.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='playground-model')
        self.backend = None
        self.state = {'status': 'idle', 'loaded_model_id': None}

    def status(self):
        with self.state_lock:
            return dict(self.state)

    def set_state(self, status, **values):
        with self.state_lock:
            self.state.update(status=status, **values)

    def models(self):
        keys = ('id', 'label', 'description', 'calibrated', 'max_tokens', 'checkpoint_step')
        return {'models': [{key: model[key] for key in keys} for model in self.catalog.values()],
                'default_model': self.default_model}

    def unload(self):
        self.backend = None
        self.set_state('loading', loaded_model_id=None)
        self.cleanup(self.device)

    def validate(self, request):
        if not isinstance(request, dict) or set(request) != {'model', 'text', 'question', 'options'}:
            raise RequestError(422, 'invalid_request', 'Provide model, text, question and options.')
        if not isinstance(request['model'], str) or request['model'] not in self.catalog:
            raise RequestError(422, 'unknown_model', 'Choose a model from the configured catalog.')
        for key in ('text', 'question'):
            if not isinstance(request[key], str) or not request[key].strip():
                raise RequestError(422, 'invalid_request', 'Context and question must be nonempty text.')
        options = request['options']
        if (not isinstance(options, list) or not 2 <= len(options) <= 255 or
                any(not isinstance(option, str) or not option.strip() for option in options) or
                len(set(options)) != len(options)):
            raise RequestError(422, 'invalid_options', 'Provide 2–255 distinct, nonempty answer options.')

    def score(self, request):
        self.validate(request)
        if not self.lock.acquire(blocking=False):
            raise RequestError(503, 'busy', 'A model is loading or scoring. Try again when it finishes.')
        try:
            pending = self.executor.submit(self._score, request)
        except BaseException:
            self.lock.release()
            raise
        return pending.result()

    def close(self):
        self.executor.submit(self.unload).result()
        self.executor.shutdown(wait=True)
        self.set_state('idle')

    def _score(self, request):
        started = time.perf_counter()
        model = self.catalog[request['model']]
        try:
            if self.status()['loaded_model_id'] != model['id']:
                self.unload()
                try:
                    self.preflight(self.device)
                    if checksum(model['checkpoint']) != model['sha256']:
                        raise ValueError('Configured artifact changed')
                    self.backend = self.backend_factory(model['checkpoint'], device=self.device, max_tokens=model['max_tokens'])
                    if self.backend.calibrated is not model['calibrated'] or self.backend.max_tokens != model['max_tokens']:
                        raise ValueError('Loaded artifact does not match catalog metadata')
                except Exception as error:
                    if not isinstance(error, RequestError):
                        print(json.dumps({'event': 'model_load_failed', 'model_id': model['id'],
                                          'error_type': type(error).__name__}), flush=True)
                    # Drop construction tracebacks before releasing partial allocations.
                    release_exception_frames(error)
                    self.unload()
                    self.set_state('error')
                    if isinstance(error, RequestError):
                        raise error from None
                    raise RequestError(503, 'model_unavailable', 'The configured model could not be loaded.') from None
                self.set_state('ready', loaded_model_id=model['id'])
            payload = {'model': 'opensysone', 'state': request['text'], 'questions': {'answer': {
                'type': 'choice', 'instructions': request['question'],
                'criteria': {option: None for option in request['options']}}}}
            self.set_state('scoring')
            try:
                result = self.backend(payload)
                from jev_harness import validate_response
                validate_response(payload, result)
                probabilities = result['answers']['answer']['probabilities']
                tokens = result['usage']['input_tokens']
                if type(tokens) is not int or tokens < 0:
                    raise ValueError('Invalid backend token count')
            except ValueError as error:
                self.set_state('error')
                match = re.fullmatch(r'Input has (\d+) tokens; limit is (\d+); no truncation', str(error))
                release_exception_frames(error)
                if match:
                    raise RequestError(422, 'token_limit', f'Input has {match[1]} tokens; the limit is {match[2]}. Shorten the text, question or options; nothing was truncated.') from None
                print(json.dumps({'event': 'model_score_failed', 'model_id': model['id'],
                                  'error_type': type(error).__name__}), flush=True)
                raise RequestError(500, 'scoring_failed', 'The model could not produce a valid result.') from None
            except Exception as error:
                self.set_state('error')
                release_exception_frames(error)
                print(json.dumps({'event': 'model_score_failed', 'model_id': model['id'],
                                  'error_type': type(error).__name__}), flush=True)
                raise RequestError(500, 'scoring_failed', 'The model could not score this request.') from None
            self.set_state('ready')
            return {'model': {key: model[key] for key in ('id', 'label', 'calibrated', 'checkpoint_step')},
                    'probabilities': [{'option': option, 'probability': probabilities[option]} for option in request['options']],
                    'elapsed_seconds': time.perf_counter() - started, 'input_tokens': tokens}
        finally:
            self.lock.release()


def loopback_authority(value):
    try:
        parsed = urlsplit('//' + value)
        if (parsed.hostname not in ('127.0.0.1', 'localhost', '::1') or parsed.username is not None or
                parsed.password is not None or parsed.path or parsed.query or parsed.fragment):
            return None
        return parsed.hostname, parsed.port or 80
    except (ValueError, TypeError):
        return None


def create_server(service, port=7466, web_root=ROOT / 'web'):
    web_root = Path(web_root)
    class Handler(BaseHTTPRequestHandler):
        def setup(self):
            super().setup()
            self.connection.settimeout(15)

        def send_data(self, status, content, content_type):
            self.send_response(status)
            for key, value in (('Content-Type', content_type), ('Content-Length', str(len(content))),
                               ('Cache-Control', 'no-store'), ('X-Content-Type-Options', 'nosniff'),
                               ('Content-Security-Policy', CSP), ('Referrer-Policy', 'no-referrer')):
                self.send_header(key, value)
            self.end_headers()
            try:
                self.wfile.write(content)
            except (BrokenPipeError, ConnectionResetError):
                pass

        def json(self, status, value):
            self.send_data(status, json.dumps(value, allow_nan=False).encode(), 'application/json; charset=utf-8')

        def error(self, error):
            self.json(error.status, {'error': error.message, 'code': error.code})

        def send_error(self, code, message=None, explain=None):
            self.json(code, {'error': 'HTTP request rejected.', 'code': 'http_error'})

        def check_origin(self, post=False):
            hosts = self.headers.get_all('Host', [])
            authority = loopback_authority(hosts[0]) if len(hosts) == 1 else None
            if authority is None:
                raise RequestError(403, 'forbidden_origin', 'Use a loopback address or local SSH tunnel.')
            if post:
                origins = self.headers.get_all('Origin', [])
                if len(origins) > 1:
                    raise RequestError(403, 'forbidden_origin', 'The request must come from this playground.')
                if origins:
                    try:
                        parsed = urlsplit(origins[0])
                    except ValueError:
                        raise RequestError(403, 'forbidden_origin', 'The request must come from this playground.') from None
                    if (parsed.scheme not in ('http', 'https') or parsed.path or parsed.query or parsed.fragment or
                            loopback_authority(parsed.netloc) != authority):
                        raise RequestError(403, 'forbidden_origin', 'The request must come from this playground.')
                if self.headers.get('Sec-Fetch-Site') in ('cross-site', 'same-site'):
                    raise RequestError(403, 'forbidden_origin', 'The request must come from this playground.')

        def do_GET(self):
            try:
                self.check_origin()
                path = urlsplit(self.path).path
                if path == '/api/models':
                    return self.json(200, service.models())
                if path == '/api/status':
                    return self.json(200, service.status())
                if path not in ASSETS:
                    raise RequestError(404, 'not_found', 'Not found.')
                filename, content_type = ASSETS[path]
                try:
                    content = (web_root / filename).read_bytes()
                except OSError:
                    raise RequestError(404, 'not_found', 'Playground asset is unavailable.') from None
                self.send_data(200, content, content_type)
            except RequestError as error:
                self.error(error)
            except Exception:
                self.error(RequestError(500, 'internal_error', 'The request could not be completed.'))

        def do_POST(self):
            try:
                self.check_origin(post=True)
                if urlsplit(self.path).path != '/api/score':
                    raise RequestError(404, 'not_found', 'Not found.')
                if self.headers.get_content_type() != 'application/json':
                    raise RequestError(415, 'content_type', 'Use application/json.')
                lengths = self.headers.get_all('Content-Length', [])
                if self.headers.get('Transfer-Encoding') or len(lengths) != 1 or not lengths[0].isdigit():
                    raise RequestError(400, 'invalid_length', 'A valid Content-Length is required.')
                length = int(lengths[0])
                if not 0 < length <= MAX_BODY_BYTES:
                    raise RequestError(413, 'body_too_large', 'Request body must be 1 byte to 128 KiB.')
                raw = self.rfile.read(length)
                if len(raw) != length:
                    raise RequestError(400, 'invalid_body', 'Request body was incomplete.')
                try:
                    request = json.loads(raw)
                except (ValueError, UnicodeError):
                    raise RequestError(400, 'invalid_json', 'Request body must be valid JSON.') from None
                self.json(200, service.score(request))
            except RequestError as error:
                self.error(error)
            except (TimeoutError, OSError):
                self.error(RequestError(408, 'request_timeout', 'The request could not be read in time.'))
            except Exception:
                self.error(RequestError(500, 'internal_error', 'The request could not be completed.'))

        def log_message(self, format, *args):
            # Do not persist URLs, bodies, prompt text, or request headers.
            pass

    server = ThreadingHTTPServer(('127.0.0.1', port), Handler)
    server.daemon_threads = True
    return server


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--models', required=True, help='Trusted server-side JSON model catalog')
    parser.add_argument('--port', type=int, default=7466)
    parser.add_argument('--device', choices=('cuda', 'cpu'), default='cuda')
    parser.add_argument('--max-tokens', type=int, default=1024)
    args = parser.parse_args()
    if not 1 <= args.port <= 65535:
        parser.error('Port must be between 1 and 65535')
    Path('/proc/self/oom_score_adj').write_text('0')
    catalog, default = load_catalog(args.models, args.max_tokens)
    service = ModelService(catalog, default, args.device)
    server = create_server(service, args.port)
    signal.signal(signal.SIGTERM, lambda signum, frame: threading.Thread(target=server.shutdown, daemon=True).start())
    print(json.dumps({'event': 'playground_listening', 'host': '127.0.0.1', 'port': args.port,
                      'model_ids': list(catalog)}), flush=True)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        pass
    finally:
        server.server_close()
        service.close()


if __name__ == '__main__':
    main()