Skip to content
File

Blob: archive/str0m-probe/sfu_probe.py

python255 lines
1"""Isolated str0m publisher / aiortc subscriber with synthetic audio and data."""
2 
3import asyncio
4from datetime import datetime, timedelta, timezone
5import importlib.util
6import json
7import os
8from pathlib import Path
9import socket
10 
11from aiortc import RTCConfiguration, RTCPeerConnection, RTCSessionDescription
12from aiortc.mediastreams import MediaStreamError
13from cryptography import x509
14from cryptography.hazmat.primitives import hashes, serialization
15from cryptography.hazmat.primitives.asymmetric import ec
16from cryptography.x509.oid import NameOID
17 
18ROOT = Path(__file__).resolve().parents[2]
19OUT = ROOT / "artifacts/archive/str0m-probe"
20OUT.mkdir(parents=True, exist_ok=True)
21os.umask(0o077)
22spec = importlib.util.spec_from_file_location("sfu_api", ROOT / "archive/sfu-bringup/probe-sfu.py")
23module = importlib.util.module_from_spec(spec)
24spec.loader.exec_module(module)
25 
26 
27async 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 
35def 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 
49async 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 
253if __name__ == "__main__":
254 raise SystemExit(asyncio.run(main()))