1use std::{
9 collections::BTreeMap,
10 future::pending,
11 io,
12 sync::{
13 Arc, Mutex,
14 atomic::{AtomicUsize, Ordering},
15 },
16};
17
18use tokio::sync::{
19 mpsc::{self, error::TrySendError},
20 watch,
21};
22
23use super::UserCloseReason;
24use crate::engine::{
25 JsonPayload, RecordingStateUpdate, UserId, UserInfo, metrics::RuntimeMetrics,
26 source_model::UserStreamId, sync::lock_unpoisoned,
27};
28
29pub const MAX_BROADCAST_PAYLOAD_BYTES: usize = 16 * 1024;
30
31const ROOM_EVENT_QUEUE_BYTES: usize = 1024;
32const BROADCAST_QUEUE_OVERHEAD_BYTES: usize = 256;
33const TRACK_PROJECTION_QUEUE_OVERHEAD_BYTES: usize = 256;
34
35#[derive(Debug, Clone, PartialEq, Eq)]
36pub struct BroadcastPayload {
37 message: Arc<JsonPayload>,
38 byte_len: usize,
39}
40
41impl BroadcastPayload {
42 pub fn try_new(message: JsonPayload) -> Result<Self, BroadcastPayloadError> {
51 let byte_len = serialized_json_len(&message)?;
52 if byte_len > MAX_BROADCAST_PAYLOAD_BYTES {
53 return Err(BroadcastPayloadError::TooLarge {
54 actual: byte_len,
55 limit: MAX_BROADCAST_PAYLOAD_BYTES,
56 });
57 }
58 Ok(Self {
59 message: Arc::new(message),
60 byte_len,
61 })
62 }
63
64 #[must_use]
65 pub const fn byte_len(&self) -> usize {
66 self.byte_len
67 }
68
69 #[must_use]
70 pub fn to_json(&self) -> JsonPayload {
71 self.message.as_ref().clone()
72 }
73}
74
75#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76pub enum BroadcastPayloadError {
77 TooLarge { actual: usize, limit: usize },
78 JsonSerialization,
79}
80
81#[derive(Debug, Clone, PartialEq, Eq)]
82pub enum RoomEventMessage {
83 Broadcast {
84 sender_id: UserId,
85 message: BroadcastPayload,
86 },
87 UserJoined {
88 user_id: UserId,
89 info: UserInfo,
90 },
91 UserDeparted {
92 user_id: UserId,
93 },
94 UserInfoChanged(BTreeMap<UserId, UserInfo>),
95 RecordingStateChanged(RecordingStateUpdate),
96}
97
98impl RoomEventMessage {
99 #[must_use]
100 pub(super) fn queued_bytes(&self) -> usize {
101 match self {
102 Self::Broadcast { message, .. } => message
103 .byte_len()
104 .saturating_add(BROADCAST_QUEUE_OVERHEAD_BYTES),
105 Self::UserInfoChanged(snapshot) => {
106 ROOM_EVENT_QUEUE_BYTES.saturating_mul(snapshot.len())
107 }
108 Self::UserJoined { .. }
109 | Self::UserDeparted { .. }
110 | Self::RecordingStateChanged(_) => ROOM_EVENT_QUEUE_BYTES,
111 }
112 }
113}
114
115#[derive(Debug, Default)]
116struct JsonByteCounter {
117 len: usize,
118}
119
120impl io::Write for JsonByteCounter {
121 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
122 self.len = self.len.saturating_add(buf.len());
123 Ok(buf.len())
124 }
125
126 fn flush(&mut self) -> io::Result<()> {
127 Ok(())
128 }
129}
130
131fn serialized_json_len(value: &JsonPayload) -> Result<usize, BroadcastPayloadError> {
132 let mut counter = JsonByteCounter::default();
133 serde_json::to_writer(&mut counter, value)
134 .map_err(|_error| BroadcastPayloadError::JsonSerialization)?;
135 Ok(counter.len)
136}
137
138pub const DEFAULT_USER_OUTBOUND_QUEUE_CAPACITY: usize = 128;
139pub const DEFAULT_USER_OUTBOUND_QUEUE_BYTE_CAPACITY: usize =
140 DEFAULT_USER_OUTBOUND_QUEUE_CAPACITY * MAX_BROADCAST_PAYLOAD_BYTES;
141
142pub(super) type OutboundSender = UserOutboundSender;
143
144#[derive(Debug, Clone, PartialEq, Eq)]
145pub struct RemoteTrackProjection {
146 pub consumer_mid: String,
147 pub user_id: UserId,
148 pub stream_id: UserStreamId,
149 pub producer_active: bool,
150}
151
152#[derive(Debug, Clone, PartialEq, Eq)]
153pub struct RemoteTrackSnapshot {
154 pub tracks: Vec<RemoteTrackProjection>,
155 pub requires_negotiation: bool,
156}
157
158impl RemoteTrackSnapshot {
159 fn queued_bytes(&self) -> usize {
160 self.tracks
161 .iter()
162 .fold(ROOM_EVENT_QUEUE_BYTES, |bytes, track| {
163 let user_id_bytes = match &track.user_id {
164 UserId::Integer(_) => 0,
165 UserId::String(value) => value.len(),
166 };
167 bytes
168 .saturating_add(TRACK_PROJECTION_QUEUE_OVERHEAD_BYTES)
169 .saturating_add(track.consumer_mid.len())
170 .saturating_add(user_id_bytes)
171 .saturating_add(track.stream_id.as_str().len())
172 })
173 }
174}
175
176#[derive(Debug, Clone)]
177pub(in crate::engine::room) struct VersionedRemoteTrackSnapshot {
178 pub(in crate::engine::room) snapshot: RemoteTrackSnapshot,
179 pub(in crate::engine::room) revision: u64,
180}
181
182#[derive(Debug, Clone)]
184pub enum UserOutbound {
185 Message(RoomEventMessage),
186 RemoteTracks(RemoteTrackSnapshot),
187 Close(UserCloseReason),
188}
189
190impl UserOutbound {
191 #[must_use]
192 pub(super) fn queued_bytes(&self) -> usize {
193 match self {
194 Self::Message(message) => message.queued_bytes(),
195 Self::RemoteTracks(snapshot) => snapshot.queued_bytes(),
196 Self::Close(_) => ROOM_EVENT_QUEUE_BYTES,
197 }
198 }
199}
200
201#[derive(Debug, Clone, Copy, PartialEq, Eq)]
203pub struct UserOutboundOverflow {
204 kind: UserOutboundOverflowKind,
205 message_capacity: usize,
206 byte_capacity: usize,
207 queued_bytes: usize,
208 message_bytes: usize,
209}
210
211#[derive(Debug, Clone, Copy, PartialEq, Eq)]
213pub enum UserOutboundOverflowKind {
214 MessageCount,
215 QueuedBytes,
216}
217
218impl UserOutboundOverflow {
219 const fn new(
220 kind: UserOutboundOverflowKind,
221 message_capacity: usize,
222 byte_capacity: usize,
223 queued_bytes: usize,
224 message_bytes: usize,
225 ) -> Self {
226 Self {
227 kind,
228 message_capacity,
229 byte_capacity,
230 queued_bytes,
231 message_bytes,
232 }
233 }
234
235 #[must_use]
236 pub const fn capacity(self) -> usize {
237 self.message_capacity
238 }
239
240 #[must_use]
241 pub const fn kind(self) -> UserOutboundOverflowKind {
242 self.kind
243 }
244
245 #[must_use]
246 pub const fn byte_capacity(self) -> usize {
247 self.byte_capacity
248 }
249
250 #[must_use]
251 pub const fn queued_bytes(self) -> usize {
252 self.queued_bytes
253 }
254
255 #[must_use]
256 pub const fn message_bytes(self) -> usize {
257 self.message_bytes
258 }
259}
260
261#[derive(Debug, Clone, Copy, PartialEq, Eq)]
263pub enum UserOutboundSendError {
264 Full(UserOutboundOverflow),
265 Closed,
266}
267
268#[derive(Debug)]
270pub enum UserOutboundEvent {
271 Message(UserOutbound),
272 Overflow(UserOutboundOverflow),
273 Closed,
274}
275
276#[derive(Debug, Clone, Copy, PartialEq, Eq)]
278pub struct UserOutboundQueueLimits {
279 message_capacity: usize,
280 byte_capacity: usize,
281}
282
283impl UserOutboundQueueLimits {
284 #[must_use]
285 pub fn new(message_capacity: usize, byte_capacity: usize) -> Self {
286 Self {
287 message_capacity: message_capacity.max(1),
288 byte_capacity: byte_capacity.max(1),
289 }
290 }
291
292 #[must_use]
293 pub const fn message_capacity(self) -> usize {
294 self.message_capacity
295 }
296
297 #[must_use]
298 pub const fn byte_capacity(self) -> usize {
299 self.byte_capacity
300 }
301}
302
303impl Default for UserOutboundQueueLimits {
304 fn default() -> Self {
305 Self::new(
306 DEFAULT_USER_OUTBOUND_QUEUE_CAPACITY,
307 DEFAULT_USER_OUTBOUND_QUEUE_BYTE_CAPACITY,
308 )
309 }
310}
311
312#[derive(Debug)]
313struct QueuedUserOutbound {
314 outbound: UserOutbound,
315 bytes: usize,
316}
317
318#[derive(Debug, Clone)]
319pub struct UserOutboundSender {
320 messages: mpsc::Sender<QueuedUserOutbound>,
321 overflow: watch::Sender<Option<UserOutboundOverflow>>,
322 metrics: Arc<RuntimeMetrics>,
323 limits: UserOutboundQueueLimits,
324 queued_bytes: Arc<AtomicUsize>,
325 latest_track_snapshot: Arc<Mutex<Option<VersionedRemoteTrackSnapshot>>>,
326}
327
328#[derive(Debug)]
329pub struct UserOutboundReceiver {
330 messages: mpsc::Receiver<QueuedUserOutbound>,
331 overflow: watch::Receiver<Option<UserOutboundOverflow>>,
332 metrics: Arc<RuntimeMetrics>,
333 queued_bytes: Arc<AtomicUsize>,
334}
335
336impl UserOutboundSender {
337 #[must_use]
338 pub fn channel(capacity: usize, metrics: Arc<RuntimeMetrics>) -> (Self, UserOutboundReceiver) {
339 Self::channel_with_limits(
340 UserOutboundQueueLimits::new(capacity, DEFAULT_USER_OUTBOUND_QUEUE_BYTE_CAPACITY),
341 metrics,
342 )
343 }
344
345 #[must_use]
346 pub fn channel_with_limits(
347 limits: UserOutboundQueueLimits,
348 metrics: Arc<RuntimeMetrics>,
349 ) -> (Self, UserOutboundReceiver) {
350 let (messages_tx, messages_rx) = mpsc::channel(limits.message_capacity());
351 let (overflow_tx, overflow_rx) = watch::channel(None);
352 let queued_bytes = Arc::new(AtomicUsize::new(0));
353 (
354 Self {
355 messages: messages_tx,
356 overflow: overflow_tx,
357 metrics: Arc::clone(&metrics),
358 limits,
359 queued_bytes: Arc::clone(&queued_bytes),
360 latest_track_snapshot: Arc::new(Mutex::new(None)),
361 },
362 UserOutboundReceiver {
363 messages: messages_rx,
364 overflow: overflow_rx,
365 metrics,
366 queued_bytes,
367 },
368 )
369 }
370
371 pub fn send(&self, outbound: UserOutbound) -> Result<(), UserOutboundSendError> {
382 self.enqueue(outbound)
383 }
384
385 pub(in crate::engine::room) fn send_remote_tracks(
388 &self,
389 snapshot: VersionedRemoteTrackSnapshot,
390 ) -> Result<(), UserOutboundSendError> {
391 let revision = snapshot.revision;
392 {
393 let mut latest = lock_unpoisoned(&self.latest_track_snapshot);
394 if self.messages.is_closed() {
395 return Err(UserOutboundSendError::Closed);
396 }
397 if latest
398 .as_ref()
399 .is_none_or(|current| revision > current.revision)
400 {
401 self.enqueue(UserOutbound::RemoteTracks(snapshot.snapshot.clone()))?;
402 *latest = Some(snapshot);
403 return Ok(());
404 }
405 if let Some(current) = latest.as_mut()
406 && revision < current.revision
407 && snapshot.snapshot.requires_negotiation
408 {
409 current.snapshot.requires_negotiation = true;
410 self.enqueue(UserOutbound::RemoteTracks(current.snapshot.clone()))?;
411 }
412 }
413 Ok(())
414 }
415
416 fn enqueue(&self, outbound: UserOutbound) -> Result<(), UserOutboundSendError> {
417 let bytes = outbound.queued_bytes();
418 self.reserve_bytes(bytes)?;
419 match self
420 .messages
421 .try_send(QueuedUserOutbound { outbound, bytes })
422 {
423 Ok(()) => {
424 self.metrics.add_ws_outbound_queued_messages(1);
425 Ok(())
426 }
427 Err(TrySendError::Full(_outbound)) => {
428 self.release_bytes(bytes);
429 let overflow = self.mark_overflow(
430 UserOutboundOverflowKind::MessageCount,
431 self.queued_bytes.load(Ordering::Acquire),
432 bytes,
433 );
434 Err(UserOutboundSendError::Full(overflow))
435 }
436 Err(TrySendError::Closed(_outbound)) => {
437 self.release_bytes(bytes);
438 Err(UserOutboundSendError::Closed)
439 }
440 }
441 }
442
443 fn reserve_bytes(&self, bytes: usize) -> Result<(), UserOutboundSendError> {
444 let byte_capacity = self.limits.byte_capacity();
445 let mut queued = self.queued_bytes.load(Ordering::Acquire);
446 loop {
447 let Some(next) = queued
448 .checked_add(bytes)
449 .filter(|next| *next <= byte_capacity)
450 else {
451 let overflow =
452 self.mark_overflow(UserOutboundOverflowKind::QueuedBytes, queued, bytes);
453 return Err(UserOutboundSendError::Full(overflow));
454 };
455 match self.queued_bytes.compare_exchange_weak(
456 queued,
457 next,
458 Ordering::AcqRel,
459 Ordering::Acquire,
460 ) {
461 Ok(_previous) => return Ok(()),
462 Err(current) => queued = current,
463 }
464 }
465 }
466
467 fn release_bytes(&self, bytes: usize) {
468 self.queued_bytes.fetch_sub(bytes, Ordering::AcqRel);
469 }
470
471 fn mark_overflow(
472 &self,
473 kind: UserOutboundOverflowKind,
474 queued_bytes: usize,
475 message_bytes: usize,
476 ) -> UserOutboundOverflow {
477 let overflow = UserOutboundOverflow::new(
478 kind,
479 self.limits.message_capacity(),
480 self.limits.byte_capacity(),
481 queued_bytes,
482 message_bytes,
483 );
484 self.metrics.record_ws_outbound_queue_overflow();
485 let _ = self.overflow.send(Some(overflow));
486 overflow
487 }
488}
489
490impl UserOutboundReceiver {
491 #[must_use]
492 pub fn has_overflowed(&self) -> bool {
493 self.overflow.borrow().is_some()
494 }
495
496 pub async fn recv(&mut self) -> Option<UserOutbound> {
500 self.messages
501 .recv()
502 .await
503 .map(|message| self.record_received(message))
504 }
505
506 pub fn try_recv(&mut self) -> Result<UserOutbound, mpsc::error::TryRecvError> {
515 self.messages
516 .try_recv()
517 .map(|message| self.record_received(message))
518 }
519
520 pub async fn recv_event(&mut self) -> UserOutboundEvent {
522 if let Some(overflow) = *self.overflow.borrow_and_update() {
523 return UserOutboundEvent::Overflow(overflow);
524 }
525 tokio::select! {
526 biased;
527 overflow = wait_for_overflow(&mut self.overflow) => {
528 UserOutboundEvent::Overflow(overflow)
529 }
530 message = self.messages.recv() => {
531 message.map_or(UserOutboundEvent::Closed, |message| {
532 UserOutboundEvent::Message(self.record_received(message))
533 })
534 }
535 }
536 }
537
538 fn record_received(&self, message: QueuedUserOutbound) -> UserOutbound {
539 self.queued_bytes.fetch_sub(message.bytes, Ordering::AcqRel);
540 self.metrics.add_ws_outbound_queued_messages(-1);
541 message.outbound
542 }
543}
544
545impl Drop for UserOutboundReceiver {
546 fn drop(&mut self) {
547 let mut pending = 0_i64;
548 let mut bytes = 0_usize;
549 while let Ok(message) = self.messages.try_recv() {
550 pending = pending.saturating_add(1);
551 bytes = bytes.saturating_add(message.bytes);
552 }
553 if pending > 0 {
554 self.metrics.add_ws_outbound_queued_messages(-pending);
555 }
556 if bytes > 0 {
557 self.queued_bytes.fetch_sub(bytes, Ordering::AcqRel);
558 }
559 }
560}
561
562async fn wait_for_overflow(
563 overflow: &mut watch::Receiver<Option<UserOutboundOverflow>>,
564) -> UserOutboundOverflow {
565 loop {
566 if let Some(overflow) = *overflow.borrow_and_update() {
567 return overflow;
568 }
569 if overflow.changed().await.is_err() {
570 pending::<()>().await;
573 }
574 }
575}
576
577#[derive(Debug, Clone)]
578pub(super) struct MessageFanout {
579 recipients: Vec<OutboundSender>,
580 message: RoomEventMessage,
581}
582
583impl MessageFanout {
584 pub(super) fn emit(self) {
585 for recipient in self.recipients {
586 let _ = recipient.send(UserOutbound::Message(self.message.clone()));
587 }
588 }
589}
590
591pub(super) fn fanout_all(
592 recipients: impl IntoIterator<Item = OutboundSender>,
593 message: RoomEventMessage,
594) -> MessageFanout {
595 MessageFanout {
596 recipients: recipients.into_iter().collect(),
597 message,
598 }
599}