| |
| import asyncio |
| import json |
| import re |
|
|
| from fastapi import FastAPI, WebSocket, WebSocketDisconnect |
| from fastapi.responses import PlainTextResponse |
|
|
| app = FastAPI() |
| CODE_RE = re.compile(r"^\d{9}$") |
|
|
| |
| PENDING = {} |
| PENDING_LOCK = asyncio.Lock() |
|
|
|
|
| @app.get("/") |
| async def root(): |
| return PlainTextResponse("ok") |
|
|
|
|
| async def register(code, role, ws): |
| async with PENDING_LOCK: |
| entry = PENDING.get(code) |
| if entry is None: |
| entry = {"sender": None, "receiver": None, "event": asyncio.Event()} |
| PENDING[code] = entry |
| if entry.get(role) is not None: |
| return None, "role already connected" |
| entry[role] = ws |
| if entry.get("sender") is not None and entry.get("receiver") is not None: |
| entry["event"].set() |
| return entry, None |
|
|
|
|
| async def unregister(code, role, ws): |
| async with PENDING_LOCK: |
| entry = PENDING.get(code) |
| if entry is None: |
| return |
| if entry.get(role) is ws: |
| entry[role] = None |
| if entry.get("sender") is None and entry.get("receiver") is None: |
| PENDING.pop(code, None) |
|
|
|
|
| async def forward(src, dst): |
| try: |
| while True: |
| msg = await src.receive() |
| if msg.get("type") == "websocket.disconnect": |
| break |
| if msg.get("bytes") is not None: |
| await dst.send_bytes(msg["bytes"]) |
| elif msg.get("text") is not None: |
| await dst.send_text(msg["text"]) |
| except WebSocketDisconnect: |
| pass |
| except Exception: |
| pass |
|
|
|
|
| @app.websocket("/ws") |
| async def ws_relay(ws: WebSocket): |
| await ws.accept() |
| code = None |
| role = None |
| entry = None |
| try: |
| raw = await ws.receive_text() |
| try: |
| payload = json.loads(raw) |
| except json.JSONDecodeError: |
| await ws.send_text(json.dumps({"error": "invalid json"})) |
| return |
| role = payload.get("role") |
| code = payload.get("code") |
| if role not in ("sender", "receiver"): |
| await ws.send_text(json.dumps({"error": "invalid role"})) |
| return |
| if not isinstance(code, str) or not CODE_RE.match(code): |
| await ws.send_text(json.dumps({"error": "invalid code"})) |
| return |
| entry, err = await register(code, role, ws) |
| if err: |
| await ws.send_text(json.dumps({"error": err})) |
| return |
| await ws.send_text(json.dumps({"status": "waiting"})) |
| await entry["event"].wait() |
| other = entry["receiver"] if role == "sender" else entry["sender"] |
| if other is None: |
| await ws.send_text(json.dumps({"error": "peer missing"})) |
| return |
| await ws.send_text(json.dumps({"status": "paired"})) |
| await other.send_text(json.dumps({"status": "paired"})) |
| task_a = asyncio.create_task(forward(ws, other)) |
| task_b = asyncio.create_task(forward(other, ws)) |
| done, pending = await asyncio.wait( |
| {task_a, task_b}, return_when=asyncio.FIRST_COMPLETED |
| ) |
| for task in pending: |
| task.cancel() |
| finally: |
| if code and role: |
| await unregister(code, role, ws) |
| try: |
| await ws.close() |
| except Exception: |
| pass |
|
|