PyAgoraRTC 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,696 @@
1
+ """``AgoraSession``: one gateway WebSocket session, from join to close (``docs/architecture.md`` §2).
2
+
3
+ Composes the pure layers (``sdp``, ``session.messages``, ``session.recovery``) with a ``GatewayTransport``.
4
+ Every background task is owned (D13); every ending goes through ``_end`` and fires ``on_closed`` once (D14, D23).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import asyncio
10
+ from collections.abc import Mapping
11
+ from contextlib import suppress
12
+ from dataclasses import replace
13
+ import logging
14
+ import time
15
+ from typing import TYPE_CHECKING
16
+
17
+ from pyagorartc.ap.response import fingerprints_from_edge
18
+ from pyagorartc.const import (
19
+ DECLARED_SSRC_TIMEOUT_S,
20
+ EDGE_DOMAIN_SUFFIX,
21
+ GATEWAY_SEND_TIMEOUT_S,
22
+ KEEPALIVE_INTERVAL_S,
23
+ PING_INTERVAL_S,
24
+ )
25
+ from pyagorartc.exceptions import GatewayConnectError, JoinTimeoutError, SessionClosedError
26
+ from pyagorartc.models import CloseReason, IceCandidate, SessionOptions, fingerprint
27
+ from pyagorartc.sdp import answer_from_ortc, candidates_to_ortc, extract_inline_candidates, offer_to_ortc
28
+ from pyagorartc.session.messages import (
29
+ JOIN_ROLE,
30
+ FrameType,
31
+ build_join,
32
+ build_leave,
33
+ build_ping,
34
+ build_renew_token,
35
+ build_set_client_role,
36
+ build_subscribe,
37
+ build_unsubscribe,
38
+ describe_frame,
39
+ encode_frame,
40
+ is_quit,
41
+ new_process_id,
42
+ new_request_id,
43
+ parse_error,
44
+ parse_frame,
45
+ parse_join_result,
46
+ parse_notification,
47
+ parse_p2p_lost,
48
+ parse_p2p_ok,
49
+ parse_remote_stream,
50
+ parse_rtp_capability_change,
51
+ parse_user_event,
52
+ )
53
+ from pyagorartc.session.recovery import Keepalive, PeerRecovery, RenewDebounce
54
+ from pyagorartc.session.transport import WebsocketsTransport
55
+
56
+ if TYPE_CHECKING:
57
+ from collections.abc import Awaitable, Callable, Coroutine
58
+
59
+ from pyagorartc.ap.response import APResponse
60
+ from pyagorartc.models import ChannelCredentials, EdgeAddress, RemoteStream
61
+ from pyagorartc.session.messages import GatewayFrame, JoinResult, JsonObject
62
+ from pyagorartc.session.transport import GatewayConnection, GatewayTransport
63
+
64
+ _LOGGER = logging.getLogger(__name__)
65
+
66
+ # shipped (HA-Luba ``_send_set_client_role``; protocol.md §3.3)
67
+ _CLIENT_ROLE_LEVEL = 0
68
+
69
+ type Spawn = Callable[[Coroutine[object, object, None]], asyncio.Task[None]]
70
+
71
+
72
+ class AgoraSession:
73
+ """One offer, one join, one close (architecture §2); a new offer is a new session.
74
+
75
+ ``deadline`` is absolute seconds on ``clock``; ``sleep`` must measure the same clock (tests pass a
76
+ manual pair). ``spawn`` replaces the task factory for every background task (D13). Host callbacks
77
+ run in owned tasks; an exception they raise is logged, never propagated into the session.
78
+ """
79
+
80
+ def __init__( # noqa: PLR0913 - every host hook and seam is explicit and keyword-only
81
+ self,
82
+ creds: ChannelCredentials,
83
+ ap: APResponse,
84
+ *,
85
+ options: SessionOptions = SessionOptions(), # noqa: B008 - frozen dataclass, safe as a default
86
+ transport: GatewayTransport | None = None,
87
+ token_provider: Callable[[], Awaitable[str | None]] | None = None,
88
+ on_peer_left: Callable[[int], Awaitable[None]] | None = None,
89
+ on_closed: Callable[[CloseReason], Awaitable[None]] | None = None,
90
+ on_stream: Callable[[RemoteStream], Awaitable[None]] | None = None,
91
+ keepalive: Callable[[], Awaitable[bool]] | None = None,
92
+ keepalive_interval_s: float = KEEPALIVE_INTERVAL_S,
93
+ deadline: float | None = None,
94
+ clock: Callable[[], float] = time.monotonic,
95
+ wall_clock: Callable[[], float] = time.time,
96
+ sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
97
+ spawn: Spawn | None = None,
98
+ request_id_factory: Callable[[], str] = new_request_id,
99
+ ) -> None:
100
+ self._creds = creds
101
+ self._ap = ap
102
+ self._options = options
103
+ self._transport: GatewayTransport = transport if transport is not None else WebsocketsTransport()
104
+ self._token_provider = token_provider
105
+ self._on_peer_left = on_peer_left
106
+ self._on_closed = on_closed
107
+ self._on_stream = on_stream
108
+ self._keepalive = keepalive
109
+ self._keepalive_policy = Keepalive(interval_s=keepalive_interval_s, deadline=deadline)
110
+ self._clock = clock
111
+ self._wall_clock = wall_clock
112
+ self._sleep = sleep
113
+ self._spawn_factory = spawn
114
+ self._new_id = request_id_factory
115
+ self._recovery = PeerRecovery(clock=clock)
116
+ self._renew_debounce = RenewDebounce(window_s=options.renew_debounce_s, clock=clock)
117
+
118
+ self._candidates: list[IceCandidate] = []
119
+ self._tasks: set[asyncio.Task[None]] = set()
120
+ self._pending: dict[str, asyncio.Future[GatewayFrame]] = {}
121
+ self._join_waiter: asyncio.Future[GatewayFrame] | None = None
122
+ self._conn: GatewayConnection | None = None
123
+ self._edge: EdgeAddress | None = None
124
+ self._reader_stopped = False
125
+ self._join_started = False
126
+ self._joined = False
127
+ self._own_uid = creds.uid
128
+ self._token = creds.token
129
+ self._rtx = False
130
+ self._online: set[int] = set()
131
+ self._streams: dict[tuple[int, int], RemoteStream] = {}
132
+ self._subscriptions: dict[tuple[int, int], asyncio.Task[None]] = {}
133
+ self._stream_announced = asyncio.Event()
134
+ self._recovery_task: asyncio.Task[None] | None = None
135
+ self._unknown_types: set[str] = set()
136
+ self._close_reason: CloseReason | None = None
137
+ self._ender: asyncio.Task[object] | None = None
138
+ self._ended = asyncio.Event()
139
+
140
+ def __repr__(self) -> str:
141
+ return f"AgoraSession(channel={self._creds.channel_name!r}, uid={self._creds.uid}, state={self._state})"
142
+
143
+ @property
144
+ def _state(self) -> str:
145
+ if self._close_reason is not None:
146
+ return f"closed:{self._close_reason.value}"
147
+ if self._joined:
148
+ return "joined"
149
+ return "joining" if self._join_started else "new"
150
+
151
+ @property
152
+ def is_connected(self) -> bool:
153
+ """Whether the gateway socket is open and the session has not ended."""
154
+ return self._close_reason is None and self._conn is not None and self._conn.is_open
155
+
156
+ @property
157
+ def is_joined(self) -> bool:
158
+ """Whether the join succeeded and the session has not ended."""
159
+ return self._joined and self._close_reason is None
160
+
161
+ @property
162
+ def remote_users(self) -> frozenset[int]:
163
+ """Publishers in the channel (uids other than ours that pass ``target_uid``)."""
164
+ return frozenset(self._online)
165
+
166
+ @property
167
+ def remote_streams(self) -> tuple[RemoteStream, ...]:
168
+ """Video streams announced (in the join payload or since) that pass ``target_uid``, in arrival order."""
169
+ return tuple(self._streams.values())
170
+
171
+ @property
172
+ def close_reason(self) -> CloseReason | None:
173
+ """Why the session ended, or ``None`` while it runs."""
174
+ return self._close_reason
175
+
176
+ def add_ice_candidate(self, candidate: IceCandidate | str) -> None:
177
+ """Queue a viewer candidate for the join ORTC; once ``join`` has started it is ignored (D11, D23, Q4)."""
178
+ if self._join_started:
179
+ _LOGGER.debug("Ignoring an ICE candidate added after join: the gateway has no trickle message (Q4)")
180
+ return
181
+ if isinstance(candidate, str):
182
+ candidate = IceCandidate(candidate)
183
+ self._candidates.append(replace(candidate, candidate=candidate.candidate.removeprefix("a=")))
184
+
185
+ async def join(self, offer_sdp: str, session_id: str) -> str:
186
+ """Join the channel and return the answer SDP for ``offer_sdp``.
187
+
188
+ Raises:
189
+ SessionClosedError: ``join`` was already called, or the session ended while joining.
190
+ GatewayConnectError: The AP named no gateway edge or none accepted a connection, the socket took no
191
+ join frame within ``GATEWAY_SEND_TIMEOUT_S``, or it closed before the result.
192
+ JoinRejectedError: The gateway refused the join.
193
+ JoinTimeoutError: No join result within ``options.join_timeout_s``.
194
+ SdpError: The offer or the gateway ORTC could not be translated.
195
+
196
+ On any failure the socket is closed, owned tasks are cancelled and ``on_closed(JOIN_FAILED)`` fires
197
+ before the exception propagates (D9, D23).
198
+
199
+ """
200
+ if self._join_started or self._close_reason is not None:
201
+ raise SessionClosedError("an AgoraSession joins once; create a new session for a new offer")
202
+ self._join_started = True
203
+ try:
204
+ return await self._join(offer_sdp, session_id)
205
+ except BaseException:
206
+ await self._end(CloseReason.JOIN_FAILED)
207
+ raise
208
+
209
+ async def renew_token(self, token: str | None = None) -> None:
210
+ """Send ``renew_token``: ``token``, else ``token_provider()``, else the last token sent (D8).
211
+
212
+ Raises:
213
+ SessionClosedError: The session is not joined.
214
+ GatewayConnectError: The socket is gone.
215
+
216
+ """
217
+ if not self.is_joined:
218
+ raise SessionClosedError("renew_token needs a joined session")
219
+ if token is None and self._token_provider is not None:
220
+ token = await self._token_provider()
221
+ if token and token != self._token:
222
+ self._token = token
223
+ _LOGGER.info("Channel token rotated for %s (fingerprint %s)", self._creds.channel_name, fingerprint(token))
224
+ await self._send(build_renew_token(self._token, request_id=self._new_id()))
225
+
226
+ async def close(self) -> None:
227
+ """End the session: cancel and await every owned task, ``leave`` when joined (bounded), close the socket.
228
+
229
+ Idempotent and safe after a failed join or from inside ``on_closed`` (it never awaits its own task).
230
+ """
231
+ if self._close_reason is None:
232
+ await self._end(CloseReason.CLOSED_BY_HOST)
233
+ return
234
+ current = asyncio.current_task()
235
+ if current is self._ender:
236
+ return
237
+ await self._ended.wait()
238
+ if others := [task for task in self._tasks if task is not current]:
239
+ await asyncio.gather(*others, return_exceptions=True)
240
+
241
+ async def _join(self, offer_sdp: str, session_id: str) -> str:
242
+ ortc = offer_to_ortc(offer_sdp, dtls_role=self._options.ortc_dtls_role)
243
+ ortc["iceParameters"]["candidates"] = candidates_to_ortc(
244
+ _dedupe([*extract_inline_candidates(offer_sdp), *self._candidates])
245
+ )
246
+ self._edge, conn = await self._connect()
247
+ if self._close_reason is not None:
248
+ await conn.close()
249
+ self._ensure_running()
250
+ self._conn = conn
251
+ self._spawn(self._message_loop(conn))
252
+
253
+ request_id = self._new_id()
254
+ waiter: asyncio.Future[GatewayFrame] = asyncio.get_running_loop().create_future()
255
+ self._join_waiter = waiter
256
+ self._pending[request_id] = waiter
257
+ timer = self._spawn(self._expire_join(waiter))
258
+ try:
259
+ await self._send_bounded(
260
+ build_join(
261
+ self._creds,
262
+ ortc,
263
+ self._ap.to_ap_response(),
264
+ options=self._options,
265
+ session_id=session_id,
266
+ process_id=new_process_id(),
267
+ client_ts_ms=self._wall_ms(),
268
+ request_id=request_id,
269
+ )
270
+ )
271
+ frame = await waiter
272
+ finally:
273
+ self._pending.pop(request_id, None)
274
+ self._join_waiter = None
275
+ timer.cancel()
276
+ join = parse_join_result(frame)
277
+ if self._reader_stopped:
278
+ # The result arrived, but the message loop stopped before this task resumed to act on it.
279
+ raise GatewayConnectError("gateway socket closed before the join completed")
280
+ _LOGGER.debug("Joined channel %s as uid %s (cid %s)", self._creds.channel_name, join.uid, join.cid)
281
+ await self._on_joined(join)
282
+
283
+ remote_video = None
284
+ if self._options.declare_remote_video_ssrc:
285
+ remote_video = await self._wait_for_remote_video()
286
+ return answer_from_ortc(
287
+ self._with_edge_fingerprint(join.ortc), offer_sdp, options=self._options, remote_video=remote_video
288
+ )
289
+
290
+ async def _on_joined(self, join: JoinResult) -> None:
291
+ self._joined = True
292
+ self._own_uid = join.uid if join.uid is not None else self._creds.uid
293
+ self._rtx = join.offers_rtx
294
+ self._recovery.reset()
295
+ # Presence seen before the result named our uid.
296
+ self._online.discard(self._own_uid)
297
+ for key in [k for k in self._streams if k[0] == self._own_uid]:
298
+ del self._streams[key]
299
+ if self._options.send_set_client_role:
300
+ # D6: off by default; Mammotion mowers leave the channel when they see it.
301
+ await self._send(
302
+ build_set_client_role(
303
+ JOIN_ROLE, _CLIENT_ROLE_LEVEL, client_ts_ms=self._wall_ms(), request_id=self._new_id()
304
+ )
305
+ )
306
+ self._ensure_running()
307
+ for stream in join.existing_streams:
308
+ if self._wanted(stream.uid) and stream.uid != self._own_uid:
309
+ self._online.add(stream.uid)
310
+ self._announce(stream)
311
+ for stream in list(self._streams.values()):
312
+ self._maybe_subscribe(stream)
313
+ self._spawn(self._ping_loop())
314
+ if self._keepalive is not None or self._keepalive_policy.deadline is not None:
315
+ self._spawn(self._keepalive_loop())
316
+
317
+ def _with_edge_fingerprint(self, ortc: Mapping[str, object]) -> Mapping[str, object]:
318
+ """D26: when the gateway ORTC carries no fingerprint, the AP's for the connected edge (else any gateway edge)."""
319
+ dtls = ortc.get("dtlsParameters", {})
320
+ if not isinstance(dtls, Mapping) or _has_fingerprint(dtls):
321
+ return ortc
322
+ for edge in [*([self._edge] if self._edge is not None else []), *self._ap.get_gateway_addresses()]:
323
+ if fingerprints := fingerprints_from_edge(edge):
324
+ _LOGGER.debug("Gateway ORTC has no DTLS fingerprint; using the AP's for edge %s", edge.ip)
325
+ return {**ortc, "dtlsParameters": {**dtls, "fingerprints": fingerprints}}
326
+ return ortc
327
+
328
+ async def _connect(self) -> tuple[EdgeAddress, GatewayConnection]:
329
+ if not (edges := self._ap.get_gateway_addresses()):
330
+ raise GatewayConnectError("the access point returned no gateway edge")
331
+ for edge in edges:
332
+ url = _edge_url(edge)
333
+ try:
334
+ return edge, await self._transport.connect(
335
+ url, timeout_s=self._options.connect_timeout_s, verify_ssl=self._options.verify_ssl
336
+ )
337
+ except GatewayConnectError as exc:
338
+ _LOGGER.debug("Gateway edge %s refused the connection: %s", url, exc)
339
+ raise GatewayConnectError(f"no gateway edge accepted a connection ({len(edges)} tried)")
340
+
341
+ async def _expire_join(self, waiter: asyncio.Future[GatewayFrame]) -> None:
342
+ await self._sleep(self._options.join_timeout_s)
343
+ if not waiter.done():
344
+ waiter.set_exception(JoinTimeoutError(f"no join result within {self._options.join_timeout_s}s"))
345
+
346
+ async def _wait_for_remote_video(self) -> RemoteStream | None:
347
+ if not self._streams:
348
+ timer = self._spawn(self._expire_stream_wait())
349
+ try:
350
+ await self._stream_announced.wait()
351
+ finally:
352
+ timer.cancel()
353
+ self._ensure_running()
354
+ if not self._streams:
355
+ _LOGGER.debug(
356
+ "No video stream announced within %ss; answering without a remote SSRC", DECLARED_SSRC_TIMEOUT_S
357
+ )
358
+ return None
359
+ return next(iter(self._streams.values()))
360
+
361
+ async def _expire_stream_wait(self) -> None:
362
+ await self._sleep(DECLARED_SSRC_TIMEOUT_S)
363
+ self._stream_announced.set()
364
+
365
+ async def _message_loop(self, conn: GatewayConnection) -> None:
366
+ """Read until the socket closes; however the loop stops (a handler raising included), a joined session ends."""
367
+ try:
368
+ while self._close_reason is None:
369
+ try:
370
+ text = await conn.recv()
371
+ except GatewayConnectError:
372
+ return
373
+ if (frame := parse_frame(text)) is None:
374
+ _LOGGER.debug("Ignoring a gateway frame that is not a JSON object")
375
+ continue
376
+ await self._dispatch(frame)
377
+ finally:
378
+ self._reader_stopped = True
379
+ if (waiter := self._join_waiter) is not None and not waiter.done():
380
+ waiter.set_exception(GatewayConnectError("gateway socket closed before the join result"))
381
+ if self._joined and self._close_reason is None:
382
+ await self._end(CloseReason.SOCKET_CLOSED)
383
+
384
+ async def _dispatch(self, frame: GatewayFrame) -> None:
385
+ if frame.id is not None:
386
+ if (future := self._pending.pop(frame.id, None)) is not None and not future.done():
387
+ future.set_result(frame)
388
+ else:
389
+ _LOGGER.debug("Response matches no pending request (%s): %s", frame.result, describe_frame(frame))
390
+ return
391
+ if (handler := _HANDLERS.get(frame.type or "")) is None:
392
+ if (name := frame.type or "<none>") not in self._unknown_types:
393
+ self._unknown_types.add(name)
394
+ _LOGGER.debug("Ignoring unhandled gateway frame: %s", describe_frame(frame))
395
+ return
396
+ await handler(self, frame)
397
+
398
+ async def _on_add_video_stream(self, frame: GatewayFrame) -> None:
399
+ if (stream := parse_remote_stream(frame.message)) is None:
400
+ _LOGGER.debug("Ignoring on_add_video_stream without an int uid and ssrcId")
401
+ return
402
+ if stream.uid == self._own_uid or not self._wanted(stream.uid):
403
+ _LOGGER.debug(
404
+ "Ignoring a stream from uid %s (ours %s, target %s)",
405
+ stream.uid,
406
+ self._own_uid,
407
+ self._options.target_uid,
408
+ )
409
+ return
410
+ self._announce(stream)
411
+ self._maybe_subscribe(stream)
412
+
413
+ async def _on_user_online(self, frame: GatewayFrame) -> None:
414
+ if (event := parse_user_event(frame.message)) is None:
415
+ _LOGGER.debug("Ignoring on_user_online without a uid")
416
+ return
417
+ if event.uid == self._own_uid or not self._wanted(event.uid):
418
+ return
419
+ self._online.add(event.uid)
420
+ for stream in [s for s in self._streams.values() if s.uid == event.uid]:
421
+ self._maybe_subscribe(stream)
422
+
423
+ async def _on_user_offline(self, frame: GatewayFrame) -> None:
424
+ if (event := parse_user_event(frame.message)) is None:
425
+ _LOGGER.debug("Ignoring on_user_offline without a uid")
426
+ return
427
+ uid = event.uid
428
+ if uid == self._own_uid or not self._wanted(uid):
429
+ return
430
+ _LOGGER.debug("Peer %s left the channel (reason %s)", uid, event.reason)
431
+ self._online.discard(uid)
432
+ for key in [k for k in self._streams if k[0] == uid]:
433
+ del self._streams[key]
434
+ if subscriptions := [self._subscriptions.pop(k) for k in list(self._subscriptions) if k[0] == uid]:
435
+ for task in subscriptions:
436
+ task.cancel()
437
+ with suppress(GatewayConnectError):
438
+ await self._send(build_unsubscribe(uid, request_id=self._new_id()))
439
+ if self._on_peer_left is None or not self.is_joined:
440
+ return
441
+ if self._recovery_task is not None:
442
+ self._recovery_task.cancel()
443
+ self._recovery_task = None
444
+ if (delay := self._recovery.peer_left(uid)) is not None:
445
+ self._recovery_task = self._spawn(self._recover_peer(uid, delay))
446
+
447
+ async def _recover_peer(self, uid: int, delay: float) -> None:
448
+ await self._sleep(delay)
449
+ # Past the debounce a later departure must not cancel the host's on_peer_left mid-call.
450
+ if self._recovery_task is asyncio.current_task():
451
+ self._recovery_task = None
452
+ if not self.is_joined or self._on_peer_left is None:
453
+ return
454
+ if self._recovery.should_recover(uid, peer_present=uid in self._online):
455
+ await self._run_callback("on_peer_left", self._on_peer_left(uid))
456
+
457
+ async def _on_token_will_expire(self, _frame: GatewayFrame) -> None:
458
+ if self.is_joined and self._renew_debounce.should_send():
459
+ self._spawn(self._renew_on_expiry())
460
+
461
+ async def _renew_on_expiry(self) -> None:
462
+ try:
463
+ await self.renew_token()
464
+ except (GatewayConnectError, SessionClosedError):
465
+ self._renew_debounce.clear()
466
+ _LOGGER.debug("renew_token could not be sent; the next will_expire retries")
467
+ except Exception:
468
+ self._renew_debounce.clear()
469
+ raise
470
+
471
+ async def _on_token_did_expire(self, _frame: GatewayFrame) -> None:
472
+ _LOGGER.warning("Gateway reports the channel token expired for %s", self._creds.channel_name)
473
+ self._renew_debounce.clear()
474
+
475
+ async def _on_notification(self, frame: GatewayFrame) -> None:
476
+ notification = parse_notification(frame.message)
477
+ if not is_quit(notification):
478
+ _LOGGER.debug("Gateway notification %s (code %s)", notification.action, notification.code)
479
+ return
480
+ _LOGGER.warning(
481
+ "Gateway quit the session on %s (code %s, %s)",
482
+ self._creds.channel_name,
483
+ notification.code,
484
+ notification.detail,
485
+ )
486
+ if self._joined:
487
+ await self._end(CloseReason.GATEWAY_QUIT)
488
+
489
+ async def _on_p2p_lost(self, frame: GatewayFrame) -> None:
490
+ lost = parse_p2p_lost(frame)
491
+ if not self._options.end_on_p2p_lost:
492
+ # D22: Mammotion ignored it on purpose; Q13 asks what it means for a subscriber.
493
+ _LOGGER.debug("Gateway reported p2p_lost (code %s, %s); ignored", lost.code, lost.error)
494
+ return
495
+ _LOGGER.warning("Gateway reported p2p_lost (code %s, %s); ending the session", lost.code, lost.error)
496
+ if self._joined:
497
+ await self._end(CloseReason.P2P_LOST)
498
+
499
+ async def _on_p2p_ok(self, frame: GatewayFrame) -> None:
500
+ ok = parse_p2p_ok(frame.message)
501
+ if ok.uid is not None and ok.uid != self._own_uid:
502
+ _LOGGER.debug("p2p_ok names uid %s, not ours (%s)", ok.uid, self._own_uid)
503
+
504
+ async def _on_rtp_capability_change(self, frame: GatewayFrame) -> None:
505
+ caps = parse_rtp_capability_change(frame.message)
506
+ _LOGGER.debug("Gateway RTP capabilities changed: video codecs %s", caps.video_codecs)
507
+
508
+ async def _on_error(self, frame: GatewayFrame) -> None:
509
+ error = parse_error(frame)
510
+ _LOGGER.warning("Gateway error event (code %s): %s", error.code, error.message)
511
+
512
+ def _announce(self, stream: RemoteStream) -> None:
513
+ self._streams.setdefault((stream.uid, stream.ssrc), stream)
514
+ self._stream_announced.set()
515
+
516
+ def _maybe_subscribe(self, stream: RemoteStream) -> None:
517
+ # Subscribe once both the stream and its publisher are known; they arrive in either order (protocol.md §3.2).
518
+ key = (stream.uid, stream.ssrc)
519
+ if not self.is_joined or key in self._subscriptions or stream.uid not in self._online:
520
+ return
521
+ self._subscriptions[key] = self._spawn(self._subscribe(stream))
522
+
523
+ async def _subscribe(self, stream: RemoteStream) -> None:
524
+ key = (stream.uid, stream.ssrc)
525
+ attempts = self._options.subscribe_retry_attempts
526
+ for attempt in range(attempts + 1):
527
+ if attempt > 0 and self._subscriptions.get(key) is not asyncio.current_task():
528
+ return
529
+ request_id = self._new_id()
530
+ ack: asyncio.Future[GatewayFrame] | None = None
531
+ if attempt < attempts:
532
+ ack = asyncio.get_running_loop().create_future()
533
+ self._pending[request_id] = ack
534
+ try:
535
+ await self._send(
536
+ build_subscribe(stream, codec=self._options.client_codec, rtx=self._rtx, request_id=request_id)
537
+ )
538
+ except GatewayConnectError:
539
+ self._pending.pop(request_id, None)
540
+ _LOGGER.debug(
541
+ "Subscribe to uid %s could not be sent; the message loop reports the closed socket", stream.uid
542
+ )
543
+ return
544
+ if attempt == 0 and self._on_stream is not None:
545
+ await self._run_callback("on_stream", self._on_stream(stream))
546
+ if ack is None:
547
+ return
548
+ await self._sleep(self._options.subscribe_retry_delay_s)
549
+ self._pending.pop(request_id, None)
550
+ if ack.done() and not ack.cancelled() and ack.result().ok:
551
+ return
552
+ ack.cancel()
553
+ _LOGGER.debug("Subscribe to uid %s not acknowledged; retry %s/%s", stream.uid, attempt + 1, attempts)
554
+
555
+ async def _ping_loop(self) -> None:
556
+ while True:
557
+ await self._sleep(PING_INTERVAL_S)
558
+ try:
559
+ await self._send(build_ping(self._new_id()))
560
+ except GatewayConnectError:
561
+ _LOGGER.debug("Ping could not be sent; the message loop reports the closed socket")
562
+ return
563
+
564
+ async def _keepalive_loop(self) -> None:
565
+ policy = self._keepalive_policy
566
+ keepalive = self._keepalive
567
+ while self._close_reason is None:
568
+ if policy.deadline_reached(self._clock()):
569
+ _LOGGER.info("Session deadline reached for %s; ending the stream", self._creds.channel_name)
570
+ await self._end(CloseReason.DEADLINE)
571
+ return
572
+ if keepalive is not None and await self._run_callback("keepalive", keepalive()) is False:
573
+ _LOGGER.debug("Keep-alive callback returned False; keep-alive stopped")
574
+ keepalive = None
575
+ now = self._clock()
576
+ if keepalive is not None:
577
+ await self._sleep(policy.delay(now))
578
+ elif policy.deadline is not None:
579
+ await self._sleep(max(0.0, policy.deadline - now))
580
+ else:
581
+ return
582
+
583
+ async def _end(self, reason: CloseReason) -> None:
584
+ if self._close_reason is not None:
585
+ return
586
+ self._close_reason = reason
587
+ current = asyncio.current_task()
588
+ self._ender = current
589
+ _LOGGER.debug("Ending session on %s: %s", self._creds.channel_name, reason.value)
590
+ if (waiter := self._join_waiter) is not None and not waiter.done():
591
+ waiter.set_exception(SessionClosedError(f"session ended while joining ({reason.value})"))
592
+ self._stream_announced.set()
593
+ if others := [task for task in self._tasks if task is not current]:
594
+ for task in others:
595
+ task.cancel()
596
+ await asyncio.gather(*others, return_exceptions=True)
597
+ for future in self._pending.values():
598
+ future.cancel()
599
+ self._pending.clear()
600
+ if (conn := self._conn) is not None:
601
+ if reason is CloseReason.CLOSED_BY_HOST and self._joined and conn.is_open:
602
+ with suppress(GatewayConnectError):
603
+ await self._send_bounded(build_leave(self._new_id()))
604
+ with suppress(GatewayConnectError):
605
+ await conn.close()
606
+ self._online.clear()
607
+ self._streams.clear()
608
+ self._subscriptions.clear()
609
+ try:
610
+ if self._on_closed is not None:
611
+ await self._run_callback("on_closed", self._on_closed(reason))
612
+ finally:
613
+ self._ended.set()
614
+
615
+ async def _send(self, frame: JsonObject) -> None:
616
+ if (conn := self._conn) is None:
617
+ raise SessionClosedError("the session has no gateway socket")
618
+ await conn.send(encode_frame(frame))
619
+ _LOGGER.debug("Sent %s", describe_frame(frame))
620
+
621
+ async def _send_bounded(self, frame: JsonObject) -> None:
622
+ """``_send`` within ``GATEWAY_SEND_TIMEOUT_S``; a stalled socket raises ``GatewayConnectError``."""
623
+ try:
624
+ async with asyncio.timeout(GATEWAY_SEND_TIMEOUT_S):
625
+ await self._send(frame)
626
+ except TimeoutError as exc:
627
+ raise GatewayConnectError(f"gateway socket accepted no frame within {GATEWAY_SEND_TIMEOUT_S}s") from exc
628
+
629
+ def _spawn(self, coro: Coroutine[object, object, None]) -> asyncio.Task[None]:
630
+ if self._close_reason is not None:
631
+ coro.close()
632
+ raise SessionClosedError(f"session ended ({self._close_reason.value}); no new task is started")
633
+ task = self._spawn_factory(coro) if self._spawn_factory else asyncio.get_running_loop().create_task(coro)
634
+ self._tasks.add(task)
635
+ task.add_done_callback(self._task_done)
636
+ return task
637
+
638
+ def _task_done(self, task: asyncio.Task[None]) -> None:
639
+ self._tasks.discard(task)
640
+ if not task.cancelled() and (exc := task.exception()) is not None:
641
+ _LOGGER.warning("Session task failed on %s", self._creds.channel_name, exc_info=exc)
642
+
643
+ async def _run_callback[T](self, name: str, awaitable: Awaitable[T]) -> T | None:
644
+ try:
645
+ return await awaitable
646
+ except Exception: # noqa: BLE001 - a host callback must not break the session (D14)
647
+ _LOGGER.warning("Host callback %s raised; ignored", name, exc_info=True)
648
+ return None
649
+
650
+ def _ensure_running(self) -> None:
651
+ if self._close_reason is not None:
652
+ raise SessionClosedError(f"session ended while joining ({self._close_reason.value})")
653
+
654
+ def _wanted(self, uid: int) -> bool:
655
+ return self._options.target_uid is None or uid == self._options.target_uid
656
+
657
+ def _wall_ms(self) -> int:
658
+ return int(self._wall_clock() * 1000)
659
+
660
+
661
+ def _edge_url(edge: EdgeAddress) -> str:
662
+ return f"wss://{edge.ip.replace('.', '-')}{EDGE_DOMAIN_SUFFIX}:{edge.port}"
663
+
664
+
665
+ def _has_fingerprint(dtls: Mapping[str, object]) -> bool:
666
+ fingerprints = dtls.get("fingerprints")
667
+ return any(
668
+ isinstance(entry, Mapping) and entry.get("fingerprint")
669
+ for entry in (fingerprints if isinstance(fingerprints, list) else [])
670
+ )
671
+
672
+
673
+ def _dedupe(candidates: list[IceCandidate]) -> list[IceCandidate]:
674
+ seen: set[str] = set()
675
+ unique = []
676
+ for candidate in candidates:
677
+ if candidate.candidate not in seen:
678
+ seen.add(candidate.candidate)
679
+ unique.append(candidate)
680
+ return unique
681
+
682
+
683
+ type _Handler = Callable[[AgoraSession, GatewayFrame], Awaitable[None]]
684
+
685
+ _HANDLERS: Mapping[str, _Handler] = {
686
+ FrameType.ON_ADD_VIDEO_STREAM: AgoraSession._on_add_video_stream, # noqa: SLF001
687
+ FrameType.ON_USER_ONLINE: AgoraSession._on_user_online, # noqa: SLF001
688
+ FrameType.ON_USER_OFFLINE: AgoraSession._on_user_offline, # noqa: SLF001
689
+ FrameType.ON_TOKEN_PRIVILEGE_WILL_EXPIRE: AgoraSession._on_token_will_expire, # noqa: SLF001
690
+ FrameType.ON_TOKEN_PRIVILEGE_DID_EXPIRE: AgoraSession._on_token_did_expire, # noqa: SLF001
691
+ FrameType.ON_NOTIFICATION: AgoraSession._on_notification, # noqa: SLF001
692
+ FrameType.ON_P2P_LOST: AgoraSession._on_p2p_lost, # noqa: SLF001
693
+ FrameType.ON_P2P_OK: AgoraSession._on_p2p_ok, # noqa: SLF001
694
+ FrameType.ON_RTP_CAPABILITY_CHANGE: AgoraSession._on_rtp_capability_change, # noqa: SLF001
695
+ FrameType.ERROR: AgoraSession._on_error, # noqa: SLF001
696
+ }