File size: 3,419 Bytes
66a180a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
#!/usr/bin/env python3
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[code] = {"sender": ws|None, "receiver": ws|None, "event": asyncio.Event()}
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