File
Blob: archive/str0m-probe/sfu_probe.py
| 1 | """Isolated str0m publisher / aiortc subscriber with synthetic audio and data.""" |
| 2 | |
| 3 | import asyncio |
| 4 | from datetime import datetime, timedelta, timezone |
| 5 | import importlib.util |
| 6 | import json |
| 7 | import os |
| 8 | from pathlib import Path |
| 9 | import socket |
| 10 | |
| 11 | from aiortc import RTCConfiguration, RTCPeerConnection, RTCSessionDescription |
| 12 | from aiortc.mediastreams import MediaStreamError |
| 13 | from cryptography import x509 |
| 14 | from cryptography.hazmat.primitives import hashes, serialization |
| 15 | from cryptography.hazmat.primitives.asymmetric import ec |
| 16 | from cryptography.x509.oid import NameOID |
| 17 | |
| 18 | ROOT = Path(__file__).resolve().parents[2] |
| 19 | OUT = ROOT / "artifacts/archive/str0m-probe" |
| 20 | OUT.mkdir(parents=True, exist_ok=True) |
| 21 | os.umask(0o077) |
| 22 | spec = importlib.util.spec_from_file_location("sfu_api", ROOT / "archive/sfu-bringup/probe-sfu.py") |
| 23 | module = importlib.util.module_from_spec(spec) |
| 24 | spec.loader.exec_module(module) |
| 25 | |
| 26 | |
| 27 | async def until(predicate, timeout=20): |
| 28 | deadline = asyncio.get_running_loop().time() + timeout |
| 29 | while not predicate(): |
| 30 | if asyncio.get_running_loop().time() >= deadline: |
| 31 | raise TimeoutError("bounded condition") |
| 32 | await asyncio.sleep(0.02) |
| 33 | |
| 34 | |
| 35 | def certificate(): |
| 36 | key = ec.generate_private_key(ec.SECP256R1()) |
| 37 | name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "str0m feasibility probe")]) |
| 38 | now = datetime.now(timezone.utc) |
| 39 | cert = (x509.CertificateBuilder().subject_name(name).issuer_name(name) |
| 40 | .public_key(key.public_key()).serial_number(x509.random_serial_number()) |
| 41 | .not_valid_before(now - timedelta(minutes=1)).not_valid_after(now + timedelta(days=1)) |
| 42 | .sign(key, hashes.SHA256())) |
| 43 | cert_path, key_path = OUT / "probe-cert.der", OUT / "probe-key.der" |
| 44 | cert_path.write_bytes(cert.public_bytes(serialization.Encoding.DER)) |
| 45 | key_path.write_bytes(key.private_bytes(serialization.Encoding.DER, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())) |
| 46 | return cert_path, key_path |
| 47 | |
| 48 | |
| 49 | async def main(): |
| 50 | api = module.SFU() |
| 51 | subscriber = RTCPeerConnection(RTCConfiguration(iceServers=[])) |
| 52 | allocated = [] |
| 53 | tracks = [] |
| 54 | processes = [] |
| 55 | consumers = [] |
| 56 | peer_events = [] |
| 57 | telemetry, spectrum, replies = [], [], [] |
| 58 | audio = {"frames": 0, "samples": 0, "rates": set(), "channels": set()} |
| 59 | result = {"board_access": False, "cleanup": []} |
| 60 | stage = "initialize" |
| 61 | |
| 62 | def journal(): |
| 63 | (OUT / "sfu-private.json").write_text(json.dumps({"channels": allocated, "tracks": tracks}, indent=2) + "\n") |
| 64 | |
| 65 | async def send(value): |
| 66 | process.stdin.write((json.dumps(value) + "\n").encode()) |
| 67 | await process.stdin.drain() |
| 68 | |
| 69 | async def event(wanted, timeout=20): |
| 70 | deadline = asyncio.get_running_loop().time() + timeout |
| 71 | while True: |
| 72 | line = await asyncio.wait_for(process.stdout.readline(), max(0.1, deadline - asyncio.get_running_loop().time())) |
| 73 | if not line: |
| 74 | raise RuntimeError("publisher exited") |
| 75 | value = json.loads(line) |
| 76 | if value.get("event") == "failed": |
| 77 | raise RuntimeError("publisher failed") |
| 78 | if value.get("event") != "offer": |
| 79 | peer_events.append(value) |
| 80 | if value.get("event") == wanted: |
| 81 | return value |
| 82 | |
| 83 | async def create_channels(sid, configs): |
| 84 | response = await api.call(f"/sessions/{sid}/datachannels/new", {"dataChannels": configs}) |
| 85 | for channel in response.get("dataChannels", []): |
| 86 | if isinstance(channel.get("id"), int): |
| 87 | allocated.append({"sessionId": sid, "id": channel["id"]}) |
| 88 | journal() |
| 89 | if len(response.get("dataChannels", [])) != len(configs) or any(x.get("errorCode") for x in response["dataChannels"]): |
| 90 | raise RuntimeError("channel registration failed") |
| 91 | return {item["dataChannelName"]: item["id"] for item in response["dataChannels"]} |
| 92 | |
| 93 | async def consume(track): |
| 94 | try: |
| 95 | while True: |
| 96 | frame = await track.recv() |
| 97 | audio["frames"] += 1 |
| 98 | audio["samples"] += frame.samples |
| 99 | audio["rates"].add(frame.sample_rate) |
| 100 | audio["channels"].add(len(frame.layout.channels)) |
| 101 | except (MediaStreamError, asyncio.CancelledError): |
| 102 | pass |
| 103 | |
| 104 | @subscriber.on("track") |
| 105 | def on_track(track): |
| 106 | if track.kind == "audio": |
| 107 | consumers.append(asyncio.create_task(consume(track))) |
| 108 | |
| 109 | try: |
| 110 | cert, key = certificate() |
| 111 | probe = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) |
| 112 | probe.connect(("1.1.1.1", 80)) |
| 113 | local_ip = probe.getsockname()[0] |
| 114 | probe.close() |
| 115 | stage = "start publisher" |
| 116 | process = await asyncio.create_subprocess_exec( |
| 117 | str(Path(__file__).parent / "target/release/str0m-probe"), local_ip, str(cert), str(key), |
| 118 | stdin=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.DEVNULL) |
| 119 | processes.append(process) |
| 120 | offer = await event("offer") |
| 121 | result["offer_has_stereo_opus"] = "opus/48000/2" in offer["sdp"].lower() |
| 122 | source = (await api.call("/sessions/new"))["sessionId"] |
| 123 | stage = "register publisher audio" |
| 124 | published = await api.call(f"/sessions/{source}/tracks/new", { |
| 125 | "sessionDescription": {"type": "offer", "sdp": offer["sdp"]}, |
| 126 | "tracks": [{"location": "local", "mid": offer["mid"], "trackName": "str0m-probe-audio"}], |
| 127 | }) |
| 128 | tracks.extend({"sessionId": source, "mid": item["mid"]} for item in published.get("tracks", []) if item.get("mid")) |
| 129 | result["published_audio"] = [{k: item[k] for k in ("mid", "trackName", "errorCode") if k in item} for item in published.get("tracks", [])] |
| 130 | journal() |
| 131 | if any(item.get("errorCode") for item in published.get("tracks", [])): |
| 132 | raise RuntimeError("publisher audio rejected") |
| 133 | await send({"command": "answer", "sdp": published["sessionDescription"]["sdp"]}) |
| 134 | stage = "publisher connection" |
| 135 | await event("connected") |
| 136 | stage = "publisher channels" |
| 137 | local = await create_channels(source, [ |
| 138 | {"location": "local", "dataChannelName": "padding-a"}, |
| 139 | {"location": "local", "dataChannelName": "robot", "ordered": True}, |
| 140 | {"location": "local", "dataChannelName": "padding-b"}, |
| 141 | {"location": "local", "dataChannelName": "spectrum", "ordered": False, "maxRetransmits": 0}, |
| 142 | ]) |
| 143 | result["publisher_ids"] = {name: local[name] for name in ("robot", "spectrum")} |
| 144 | await send({"command": "channels", "robot": local["robot"], "spectrum": local["spectrum"]}) |
| 145 | |
| 146 | stage = "subscriber connection" |
| 147 | sink, transport = await api.establish() |
| 148 | await subscriber.setRemoteDescription(RTCSessionDescription(**transport["sessionDescription"])) |
| 149 | await subscriber.setLocalDescription(await subscriber.createAnswer()) |
| 150 | await api.call(f"/sessions/{sink}/renegotiate", {"sessionDescription": {"type": "answer", "sdp": subscriber.localDescription.sdp}}, method="PUT") |
| 151 | await until(lambda: subscriber.connectionState == "connected") |
| 152 | remote = await create_channels(sink, [ |
| 153 | {"location": "remote", "sessionId": source, "dataChannelName": "robot", "ordered": True, "waitForAck": True, "canReply": True}, |
| 154 | {"location": "remote", "sessionId": source, "dataChannelName": "spectrum", "ordered": False, "maxRetransmits": 0, "waitForAck": True}, |
| 155 | ]) |
| 156 | result["subscriber_ids"] = remote |
| 157 | robot = subscriber.createDataChannel("robot", negotiated=True, id=remote["robot"], ordered=True) |
| 158 | band = subscriber.createDataChannel("spectrum", negotiated=True, id=remote["spectrum"], ordered=False, maxRetransmits=0) |
| 159 | |
| 160 | @robot.on("message") |
| 161 | def robot_message(message): |
| 162 | item = json.loads(message) |
| 163 | (telemetry if item.get("kind") == "telemetry" else replies).append(item) |
| 164 | |
| 165 | @band.on("message") |
| 166 | def spectrum_message(message): |
| 167 | if isinstance(message, bytes) and len(message) == 4: |
| 168 | spectrum.append(int.from_bytes(message, "little")) |
| 169 | |
| 170 | await until(lambda: robot.readyState == "open" and band.readyState == "open") |
| 171 | robot.send("ready") |
| 172 | band.send("ready") |
| 173 | stage = "audio warmup" |
| 174 | await send({"command": "warmup", "frames": 25}) |
| 175 | result["warmup"] = (await event("warmup_done"))["stats"] |
| 176 | stage = "subscriber audio" |
| 177 | pulled = await api.call(f"/sessions/{sink}/tracks/new", {"tracks": [{"location": "remote", "sessionId": source, "trackName": "str0m-probe-audio"}]}) |
| 178 | tracks.extend({"sessionId": sink, "mid": item["mid"]} for item in pulled.get("tracks", []) if item.get("mid")) |
| 179 | result["pulled_audio"] = [{k: item[k] for k in ("mid", "trackName", "errorCode") if k in item} for item in pulled.get("tracks", [])] |
| 180 | journal() |
| 181 | if any(item.get("errorCode") for item in pulled.get("tracks", [])): |
| 182 | raise RuntimeError("subscriber audio rejected") |
| 183 | if pulled.get("sessionDescription"): |
| 184 | await subscriber.setRemoteDescription(RTCSessionDescription(**pulled["sessionDescription"])) |
| 185 | await subscriber.setLocalDescription(await subscriber.createAnswer()) |
| 186 | await api.call(f"/sessions/{sink}/renegotiate", {"sessionDescription": {"type": "answer", "sdp": subscriber.localDescription.sdp}}, method="PUT") |
| 187 | else: |
| 188 | raise RuntimeError("subscriber audio offer missing") |
| 189 | stage = "publisher channel readiness" |
| 190 | while not {"robot", "spectrum"}.issubset({e.get("label") for e in peer_events if e.get("event") == "channel_open"}): |
| 191 | await event("channel_open") |
| 192 | await asyncio.sleep(0.3) |
| 193 | stage = "synthetic workload" |
| 194 | frames = 300 |
| 195 | await send({"command": "run", "frames": frames, "drop_every": 5}) |
| 196 | for sequence in range(5): |
| 197 | robot.send(json.dumps({"kind": "command", "sequence": sequence})) |
| 198 | await asyncio.sleep(0.1) |
| 199 | done = await event("run_done") |
| 200 | await until(lambda: len(telemetry) == frames // 5 and len(replies) == 5, timeout=15) |
| 201 | await send({"command": "stats"}) |
| 202 | stats = (await event("stats"))["stats"] |
| 203 | result.update( |
| 204 | sender_stats=stats, |
| 205 | audio={k: sorted(v) if isinstance(v, set) else v for k, v in audio.items()}, |
| 206 | reliable_received=len(telemetry), |
| 207 | reliable_in_order=[x["sequence"] for x in telemetry] == list(range(frames // 5)), |
| 208 | unreliable_received=len(spectrum), |
| 209 | unreliable_unique=len(set(spectrum)), |
| 210 | command_round_trips=len(replies), |
| 211 | peer_events=peer_events, |
| 212 | ) |
| 213 | result["passed"] = ( |
| 214 | result["offer_has_stereo_opus"] and result["reliable_in_order"] and len(replies) == 5 |
| 215 | and 0 < len(spectrum) < frames // 2 and len(spectrum) == len(set(spectrum)) |
| 216 | and audio["frames"] > 200 and audio["rates"] == {48000} and audio["channels"] == {2} |
| 217 | and stats["deliberately_dropped"] > 0 |
| 218 | ) |
| 219 | except Exception as error: |
| 220 | result["passed"] = False |
| 221 | result["error"] = {"stage": stage, "type": type(error).__name__} |
| 222 | finally: |
| 223 | for item in reversed(allocated): |
| 224 | try: |
| 225 | await api.call(f"/sessions/{item['sessionId']}/datachannels/close", {"dataChannels": [{"id": item["id"]}]}, method="PUT") |
| 226 | result["cleanup"].append({"kind": "channel", "closed": True}) |
| 227 | except Exception: |
| 228 | result["cleanup"].append({"kind": "channel", "closed": False}) |
| 229 | for item in reversed(tracks): |
| 230 | try: |
| 231 | await api.call(f"/sessions/{item['sessionId']}/tracks/close", {"tracks": [{"mid": item["mid"]}], "force": True}, method="PUT") |
| 232 | result["cleanup"].append({"kind": "track", "closed": True}) |
| 233 | except Exception: |
| 234 | result["cleanup"].append({"kind": "track", "closed": False}) |
| 235 | for process in processes: |
| 236 | if process.returncode is None: |
| 237 | try: |
| 238 | await send({"command": "stop"}) |
| 239 | await asyncio.wait_for(process.wait(), timeout=5) |
| 240 | except Exception: |
| 241 | process.kill() |
| 242 | await process.wait() |
| 243 | await subscriber.close() |
| 244 | for consumer in consumers: |
| 245 | consumer.cancel() |
| 246 | await asyncio.gather(*consumers, return_exceptions=True) |
| 247 | result["cleanup_complete"] = all(x["closed"] for x in result["cleanup"]) |
| 248 | (OUT / "sfu-result.json").write_text(json.dumps(result, indent=2) + "\n") |
| 249 | print(json.dumps({k: v for k, v in result.items() if k != "peer_events"}, indent=2)) |
| 250 | return 0 if result.get("passed") and result["cleanup_complete"] else 1 |
| 251 | |
| 252 | |
| 253 | if __name__ == "__main__": |
| 254 | raise SystemExit(asyncio.run(main())) |