posttrain-arena / relay_agent.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
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)