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.py from benchflow/posttrain-arena: direct link, hf CLI and curl.
- Browser
- Download file 7.08 kB
-
https://huggingface.co/spaces/benchflow/posttrain-arena/resolve/main/relay.py
- Command line
-
hf download hf://spaces/benchflow/posttrain-arena/relay.py
-
curl -L -o relay.py https://huggingface.co/spaces/benchflow/posttrain-arena/resolve/main/relay.py
7.08 kB
| """Rollout ingress through the Space's own port: a job connects *outbound* over WebSocket, and | |
| OpenCode inside the sandboxes calls the Space over HTTPS. No inbound tunnel into the job. | |
| Job side (relay_agent.py) -> wss://<space>/relay/connect {run_id, relay_key}, authenticated with a | |
| BenchFlow editor HF token. Sandbox side -> https://<space>/relay/<run_id>/v1/... with the per-run | |
| relay key as bearer. Frames are JSON with base64 bodies; responses stream chunk by chunk. | |
| """ | |
| import asyncio, base64, hashlib, hmac, json, re, secrets, time | |
| from fastapi import APIRouter, HTTPException, Request, WebSocket, WebSocketDisconnect | |
| from fastapi.concurrency import run_in_threadpool | |
| from fastapi.responses import StreamingResponse | |
| import auth | |
| router = APIRouter() | |
| RUN_ID = re.compile(r'^[a-z0-9][a-z0-9-]{6,63}$') | |
| MAX_BODY = 16 * 1024 * 1024 | |
| HEAD_TIMEOUT = 900 | |
| CHUNK_TIMEOUT = 900 | |
| FORWARD_HEADERS = ('content-type', 'accept', 'authorization', 'user-agent', 'x-request-id') | |
| RETURN_HEADERS = ('content-type', 'x-request-id') | |
| class Relay: | |
| def __init__(self, run_id, key_hash, websocket): | |
| self.run_id, self.key_hash, self.websocket = run_id, key_hash, websocket | |
| self.pending = {} | |
| self.send_lock = asyncio.Lock() | |
| self.connected_at = time.time() | |
| self.requests = 0 | |
| self.reconnects = 0 | |
| RELAYS: dict[str, Relay] = {} | |
| def key_hash(value): return hashlib.sha256(value.encode()).hexdigest() | |
| def editor_token(token): | |
| user = auth.identity(token) | |
| if not any(o.get('name') == 'benchflow' and o.get('roleInOrg') in ('admin', 'write') for o in user.get('orgs', [])): | |
| raise HTTPException(403, 'Only BenchFlow editor jobs may register a relay.') | |
| return user['name'] | |
| async def connect(websocket: WebSocket): | |
| await websocket.accept() | |
| try: | |
| hello = await asyncio.wait_for(websocket.receive_json(), timeout=30) | |
| run_id, relay_key, token = str(hello.get('run_id', '')), str(hello.get('relay_key', '')), str(hello.get('token', '')) | |
| if not RUN_ID.fullmatch(run_id) or len(relay_key) < 32 or not token: | |
| await websocket.send_json({'ok': False, 'error': 'run_id, relay_key and token are required'}); await websocket.close(code=4400); return | |
| try: owner = await run_in_threadpool(editor_token, token) | |
| except HTTPException as error: | |
| await websocket.send_json({'ok': False, 'error': error.detail}); await websocket.close(code=4401); return | |
| except (WebSocketDisconnect, asyncio.TimeoutError, ValueError): | |
| return | |
| previous = RELAYS.get(run_id) | |
| relay = Relay(run_id, key_hash(relay_key), websocket) | |
| if previous is not None: | |
| # A reconnect keeps the run's in-flight requests: the agent finishes them and answers over the new socket. | |
| relay.pending = previous.pending | |
| relay.requests = previous.requests | |
| relay.reconnects = previous.reconnects + 1 | |
| previous.pending = {} | |
| try: await previous.websocket.close(code=4409) | |
| except Exception: pass | |
| RELAYS[run_id] = relay | |
| await websocket.send_json({'ok': True, 'run_id': run_id, 'owner': owner}) | |
| try: | |
| while True: | |
| frame = await websocket.receive_json() | |
| queue = relay.pending.get(frame.get('id')) | |
| if queue is not None: queue.put_nowait(frame) | |
| except (WebSocketDisconnect, RuntimeError): | |
| pass | |
| finally: | |
| if RELAYS.get(run_id) is relay: | |
| # Keep pending requests alive for a grace period so a reconnecting agent can still answer them. | |
| asyncio.get_running_loop().call_later(RECONNECT_GRACE, lambda: asyncio.ensure_future(expire(run_id, relay))) | |
| RECONNECT_GRACE = 120 | |
| async def expire(run_id, relay): | |
| if RELAYS.get(run_id) is relay: | |
| del RELAYS[run_id] | |
| await fail_all(relay, 'relay disconnected') | |
| async def fail_all(relay, reason): | |
| for queue in list(relay.pending.values()): | |
| queue.put_nowait({'type': 'error', 'error': reason}) | |
| def authorized(relay, request): | |
| header = request.headers.get('authorization', '') | |
| if not header.startswith('Bearer '): raise HTTPException(401, 'Use the run relay key as a bearer token.') | |
| if not hmac.compare_digest(key_hash(header[7:]), relay.key_hash): raise HTTPException(403, 'Relay key does not match this run.') | |
| async def forward(run_id: str, path: str, request: Request): | |
| relay = RELAYS.get(run_id) | |
| if relay is None: raise HTTPException(503, 'The run is not connected to the relay yet. Retry shortly.') | |
| authorized(relay, request) | |
| body = await request.body() | |
| if len(body) > MAX_BODY: raise HTTPException(413, 'Request body exceeds the relay limit.') | |
| request_id = secrets.token_hex(12) | |
| queue: asyncio.Queue = asyncio.Queue() | |
| relay.pending[request_id] = queue | |
| relay.requests += 1 | |
| target = '/' + path + ('?' + request.url.query if request.url.query else '') | |
| frame = {'id': request_id, 'method': request.method, 'path': target, | |
| 'headers': {k: v for k, v in request.headers.items() if k.lower() in FORWARD_HEADERS}, | |
| 'body': base64.b64encode(body).decode()} | |
| try: | |
| async with relay.send_lock: | |
| await relay.websocket.send_json(frame) | |
| except (RuntimeError, WebSocketDisconnect, OSError): | |
| relay.pending.pop(request_id, None) | |
| raise HTTPException(503, 'The run is reconnecting to the relay. Retry shortly.') from None | |
| try: | |
| head = await asyncio.wait_for(queue.get(), timeout=HEAD_TIMEOUT) | |
| except asyncio.TimeoutError: | |
| relay.pending.pop(request_id, None) | |
| raise HTTPException(504, 'The run did not answer through the relay in time.') from None | |
| if head.get('type') != 'head': | |
| relay.pending.pop(request_id, None) | |
| raise HTTPException(502, 'Relay error: ' + str(head.get('error', 'no response head'))) | |
| status = int(head.get('status', 502)) | |
| headers = {k: v for k, v in (head.get('headers') or {}).items() if k.lower() in RETURN_HEADERS} | |
| async def stream(): | |
| try: | |
| while True: | |
| frame = await asyncio.wait_for(queue.get(), timeout=CHUNK_TIMEOUT) | |
| kind = frame.get('type') | |
| if kind == 'chunk': yield base64.b64decode(frame.get('data', '')) | |
| elif kind == 'end': return | |
| else: return | |
| except asyncio.TimeoutError: | |
| return | |
| finally: | |
| relay.pending.pop(request_id, None) | |
| return StreamingResponse(stream(), status_code=status, headers=headers, media_type=headers.get('content-type')) | |
| def status(): | |
| return {'connected': [{'run_id': r.run_id, 'since': r.connected_at, 'requests': r.requests, 'in_flight': len(r.pending), 'reconnects': r.reconnects, 'open': r.websocket.client_state.name == 'CONNECTED'} for r in RELAYS.values()], | |
| 'note': 'Jobs connect outbound to /relay/connect; sandboxes call /relay/<run_id>/v1/... with the per-run key.'} | |