Spaces:
Running
Running
Xiangyi Li
Relay: reconnects keep in-flight requests (pending transfers to the new socket; agent outbox survives sessions); status reports reconnects and socket state
5124490 Download relay_agent.py from benchflow/posttrain-arena: direct link, hf CLI and curl.
- Browser
- Download file 4.12 kB
-
https://huggingface.co/spaces/benchflow/posttrain-arena/resolve/main/relay_agent.py
- Command line
-
hf download hf://spaces/benchflow/posttrain-arena/relay_agent.py
-
curl -L -o relay_agent.py https://huggingface.co/spaces/benchflow/posttrain-arena/resolve/main/relay_agent.py
4.12 kB
| """Job-side relay agent: one outbound WebSocket to the Space, requests proxied to the loopback model bridge. | |
| Usage (inside the HF Job): | |
| RELAY_WS_URL=wss://<space>/relay/connect RELAY_RUN_ID=<run> RELAY_KEY=<key> HF_TOKEN=<editor token> \ | |
| python relay_agent.py [--upstream http://127.0.0.1:8001] | |
| Prints RELAY_CONNECTED once registered; reconnects forever on failure. | |
| """ | |
| import argparse, asyncio, base64, json, os, sys, time | |
| import httpx | |
| import websockets | |
| CHUNK = 64 * 1024 | |
| # Responses go through one outbound queue that survives reconnects: a request that was in flight when the | |
| # socket dropped is still answered (the Space keeps its pending entry for a grace period). | |
| OUTBOX: asyncio.Queue = asyncio.Queue() | |
| TASKS: set = set() | |
| async def emit(frame): | |
| await OUTBOX.put(json.dumps(frame)) | |
| async def handle(frame, upstream): | |
| request_id = frame.get('id') | |
| try: | |
| body = base64.b64decode(frame.get('body') or '') | |
| headers = {k: v for k, v in (frame.get('headers') or {}).items() if k.lower() != 'host'} | |
| async with httpx.AsyncClient(timeout=httpx.Timeout(900.0, connect=30.0)) as client: | |
| async with client.stream(frame.get('method', 'GET'), upstream + frame.get('path', '/'), headers=headers, content=body) as response: | |
| await emit({'id': request_id, 'type': 'head', 'status': response.status_code, | |
| 'headers': {k: v for k, v in response.headers.items() if k.lower() in ('content-type', 'x-request-id')}}) | |
| async for chunk in response.aiter_bytes(CHUNK): | |
| await emit({'id': request_id, 'type': 'chunk', 'data': base64.b64encode(chunk).decode()}) | |
| await emit({'id': request_id, 'type': 'end'}) | |
| except Exception as error: | |
| await emit({'id': request_id, 'type': 'error', 'error': str(error)[:300]}) | |
| async def sender(ws): | |
| while True: | |
| message = await OUTBOX.get() | |
| try: | |
| await ws.send(message) | |
| except Exception: | |
| # Socket gone: put the message back for the next session and stop this sender. | |
| OUTBOX.put_nowait(message) | |
| raise | |
| async def session(url, run_id, key, token, upstream): | |
| headers = {'Authorization': 'Bearer ' + token} | |
| try: | |
| connect = websockets.connect(url, additional_headers=headers, max_size=64 * 1024 * 1024, ping_interval=20, ping_timeout=60) | |
| except TypeError: | |
| connect = websockets.connect(url, extra_headers=headers, max_size=64 * 1024 * 1024, ping_interval=20, ping_timeout=60) | |
| async with connect as ws: | |
| await ws.send(json.dumps({'run_id': run_id, 'relay_key': key, 'token': token})) | |
| ack = json.loads(await asyncio.wait_for(ws.recv(), timeout=60)) | |
| if not ack.get('ok'): | |
| raise RuntimeError('relay refused: ' + str(ack.get('error'))) | |
| print('RELAY_CONNECTED run=' + run_id + ' in_flight=' + str(len(TASKS)), flush=True) | |
| send_task = asyncio.create_task(sender(ws)) | |
| try: | |
| async for message in ws: | |
| try: frame = json.loads(message) | |
| except ValueError: continue | |
| task = asyncio.create_task(handle(frame, upstream)) | |
| TASKS.add(task); task.add_done_callback(TASKS.discard) | |
| finally: | |
| send_task.cancel() | |
| try: await send_task | |
| except BaseException: pass | |
| async def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--upstream', default='http://127.0.0.1:8001') | |
| args = parser.parse_args() | |
| url, run_id, key, token = os.environ['RELAY_WS_URL'], os.environ['RELAY_RUN_ID'], os.environ['RELAY_KEY'], os.environ['HF_TOKEN'] | |
| delay = 2 | |
| while True: | |
| started = time.time() | |
| try: | |
| await session(url, run_id, key, token, args.upstream) | |
| except Exception as error: | |
| print('relay disconnected: ' + str(error)[:200], flush=True) | |
| delay = 2 if time.time() - started > 60 else min(delay * 2, 30) | |
| await asyncio.sleep(delay) | |
| if __name__ == '__main__': | |
| try: asyncio.run(main()) | |
| except KeyboardInterrupt: sys.exit(0) | |