posttrain-arena / relay.py
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
Raw History Blame Contribute Delete
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']
@router.websocket('/relay/connect')
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.')
@router.api_route('/relay/{run_id}/{path:path}', methods=['GET', 'POST'])
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'))
@router.get('/api/relay/status')
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.'}