1use std::collections::BTreeMap;
9
10use secrecy::SecretString;
11use serde::{Deserialize, Serialize};
12
13mod connection_lifecycle;
14mod outbound_batch;
15mod request_flow;
16mod request_tracker;
17mod server_events;
18mod sticky_replay;
19mod timers;
20
21use outbound_batch::{FlushMode, OutboundBatcher};
22use request_tracker::RequestTracker;
23use sticky_replay::StickyReplayState;
24use timers::RequestTimeoutId;
25
26use crate::{
27 shared::{
28 AvailableFeatures, DownloadStates, JsonPayload, RecordingState, RecordingStateUpdate,
29 StreamType, UserId, UserInfo,
30 },
31 signaling::{
32 AuthPayload, ClientBroadcastPayload, ClientEnvelope, ClientMessage, MAX_ENVELOPE_BATCH_LEN,
33 NegotiationUploadSlot, PeerSnapshot, RequestId, ServerEnvelope, StreamIntentPayload,
34 SubscribePayload, TrackBinding, WebSocketCloseCode, WelcomePayload, decode_envelope_batch,
35 },
36 wire::ServerMessage,
37};
38
39pub const RECOVERY_TIMER_ID: u32 = 1;
41const BATCH_FLUSH_TIMER_ID: u32 = 2;
42const INITIAL_RECOVERY_DELAY_MS: u32 = 1_000;
43const MAX_RECOVERY_DELAY_MS: u32 = 30_000;
44const BATCH_FLUSH_DELAY_MS: u32 = 100;
45const REQUEST_TIMEOUT_MS: u32 = 5_000;
46const MAX_OUTBOUND_BATCH_LEN: usize = 16;
47
48#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
52#[serde(tag = "kind", rename_all = "camelCase")]
53pub enum Command {
54 SendWebSocket {
56 frame: String,
57 },
58 ApplyNegotiation {
60 #[serde(rename = "requestId")]
61 request_id: RequestId,
62 #[serde(rename = "negotiationKind")]
63 kind: NegotiationKind,
64 sdp: String,
65 #[serde(rename = "uploadSlots")]
66 upload_slots: Vec<NegotiationUploadSlot>,
67 },
68 ClosePeerConnection,
69 CloseWebSocket {
70 code: u16,
71 },
72 EmitStateChange {
75 state: ConnectionState,
76 cause: Option<String>,
77 },
78 SetAvailableFeatures {
79 features: AvailableFeatures,
80 },
81 SetRecordingState {
82 state: RecordingState,
83 },
84 #[serde(rename = "emitUpdate")]
86 EmitEvent {
87 #[serde(
88 rename = "update",
89 serialize_with = "crate::host_bridge::serialize_protocol_event"
90 )]
91 event: ProtocolEvent,
92 },
93 BeginPendingRequest {
94 request: PendingRequest,
95 },
96 CompletePendingRequest {
98 #[serde(rename = "requestId")]
99 request_id: RequestId,
100 #[serde(rename = "timeoutTimerId")]
101 timeout_timer_id: u32,
102 ok: bool,
103 },
104 ScheduleTimer {
107 id: u32,
108 ms: u32,
109 },
110 CancelTimer {
111 id: u32,
112 },
113 Connect {
115 url: String,
116 },
117}
118
119pub(crate) type Commands = Vec<Command>;
120
121#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
122#[serde(rename_all = "snake_case")]
123pub enum ConnectionState {
124 Disconnected,
125 Connecting,
126 Authenticated,
127 Connected,
128 Recovering,
129 Closed,
130}
131
132#[derive(Debug, Clone, PartialEq, Eq)]
133pub enum ProtocolEvent {
134 PeerSnapshot {
135 peers: Vec<PeerSnapshot>,
136 },
137 TrackSnapshot {
138 bindings: Vec<TrackBinding>,
139 },
140 PeerInfo {
141 user_id: UserId,
142 info: UserInfo,
143 },
144 PeerLeft {
145 user_id: UserId,
146 },
147 Broadcast {
148 sender_id: UserId,
149 message: JsonPayload,
150 },
151 RecordingStateChanged {
152 state: RecordingStateUpdate,
153 },
154}
155
156#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
157#[serde(rename_all = "camelCase")]
158pub enum NegotiationKind {
159 Offer,
160 Renegotiate,
161}
162
163#[derive(Debug, Clone, Copy, PartialEq, Eq)]
164pub(crate) enum PendingRequestKind {
165 StartRecording,
166 StopRecording,
167}
168
169#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
170#[serde(rename_all = "camelCase")]
171pub struct PendingRequest {
172 pub request_id: RequestId,
173 pub timeout_timer_id: u32,
174 pub timeout_ms: u32,
175}
176
177#[derive(Debug, Clone, PartialEq, Eq)]
178struct ConnectContext {
179 url: String,
180 jwt: String,
181 room: Option<String>,
182}
183
184#[derive(Debug, Clone, PartialEq, Eq)]
185enum ProtocolPhase {
186 Disconnected,
187 Connecting,
188 Authenticated(Option<RequestId>),
189 Connected(Option<RequestId>),
190 Recovering,
191 Closed,
192}
193
194impl ProtocolPhase {
195 const fn connection_state(&self) -> ConnectionState {
196 match self {
197 Self::Disconnected => ConnectionState::Disconnected,
198 Self::Connecting => ConnectionState::Connecting,
199 Self::Authenticated(_) => ConnectionState::Authenticated,
200 Self::Connected(_) => ConnectionState::Connected,
201 Self::Recovering => ConnectionState::Recovering,
202 Self::Closed => ConnectionState::Closed,
203 }
204 }
205
206 const fn is_awaiting_welcome(&self) -> bool {
207 matches!(self, Self::Connecting | Self::Recovering)
208 }
209
210 const fn can_send_client_messages(&self) -> bool {
211 matches!(self, Self::Authenticated(_) | Self::Connected(_))
212 }
213
214 const fn can_enter_connected(&self) -> bool {
215 matches!(self, Self::Authenticated(None))
216 }
217
218 fn accept_negotiation(
219 &mut self,
220 request_id: &RequestId,
221 kind: NegotiationKind,
222 ) -> Result<(), NegotiationRejection> {
223 match (self, kind) {
224 (Self::Authenticated(pending), NegotiationKind::Offer)
225 | (Self::Connected(pending), NegotiationKind::Renegotiate) => {
226 if pending.is_some() {
227 return Err(NegotiationRejection::ProtocolError);
228 }
229 *pending = Some(request_id.clone());
230 Ok(())
231 }
232 (Self::Authenticated(_), NegotiationKind::Renegotiate)
233 | (Self::Connected(_), NegotiationKind::Offer) => {
234 Err(NegotiationRejection::ProtocolError)
235 }
236 (
237 Self::Disconnected | Self::Connecting | Self::Recovering | Self::Closed,
238 NegotiationKind::Offer | NegotiationKind::Renegotiate,
239 ) => Err(NegotiationRejection::Ignored),
240 }
241 }
242
243 fn resolve_negotiation(&mut self, request_id: &RequestId, kind: NegotiationKind) -> bool {
244 let pending = match (self, kind) {
245 (Self::Authenticated(pending), NegotiationKind::Offer)
246 | (Self::Connected(pending), NegotiationKind::Renegotiate) => pending,
247 (Self::Authenticated(_), NegotiationKind::Renegotiate)
248 | (Self::Connected(_), NegotiationKind::Offer)
249 | (
250 Self::Disconnected | Self::Connecting | Self::Recovering | Self::Closed,
251 NegotiationKind::Offer | NegotiationKind::Renegotiate,
252 ) => return false,
253 };
254 if pending.as_ref() != Some(request_id) {
255 return false;
256 }
257 *pending = None;
258 true
259 }
260}
261
262#[derive(Debug, Clone, Copy, PartialEq, Eq)]
263pub(super) enum NegotiationRejection {
264 Ignored,
265 ProtocolError,
266}
267
268#[derive(Debug, Clone, PartialEq, Eq)]
273pub struct ProtocolCore {
274 phase: ProtocolPhase,
276 track_users_by_mid: BTreeMap<String, UserId>,
281 sticky_replay: StickyReplayState,
289 connect_context: Option<ConnectContext>,
295 recovery_delay_ms: u32,
301 outbound_batch: OutboundBatcher,
308 request_tracker: RequestTracker,
314}
315
316impl Default for ProtocolCore {
317 fn default() -> Self {
318 Self::new()
319 }
320}
321
322impl ProtocolCore {
323 #[must_use]
329 pub fn new() -> Self {
330 Self {
331 phase: ProtocolPhase::Disconnected,
332 track_users_by_mid: BTreeMap::new(),
333 sticky_replay: StickyReplayState::new(),
334 connect_context: None,
335 recovery_delay_ms: INITIAL_RECOVERY_DELAY_MS,
336 outbound_batch: OutboundBatcher::new(),
337 request_tracker: RequestTracker::new(),
338 }
339 }
340
341 #[must_use]
342 pub const fn state(&self) -> ConnectionState {
343 self.phase.connection_state()
344 }
345
346 pub fn on_ws_open(&mut self) -> Vec<Command> {
351 if !self.phase.is_awaiting_welcome() {
352 return Vec::new();
353 }
354 let Some(connect_context) = self.connect_context.as_ref() else {
355 return Vec::new();
356 };
357 self.enqueue_client_message(
358 ClientMessage::Auth(AuthPayload {
359 jwt: SecretString::from(connect_context.jwt.clone()),
360 channel: connect_context.room.clone(),
361 }),
362 FlushMode::Immediate,
363 )
364 }
365
366 pub fn on_ws_message(&mut self, frame: &str) -> Vec<Command> {
372 let Ok(batch) = decode_envelope_batch(frame, MAX_ENVELOPE_BATCH_LEN) else {
373 return close_for_protocol_error();
374 };
375 let Ok(envelopes) = batch
376 .into_iter()
377 .map(ServerEnvelope::decode)
378 .collect::<Result<Vec<_>, _>>()
379 else {
380 return close_for_protocol_error();
381 };
382 let mut commands = Vec::new();
383 for envelope in envelopes {
384 match envelope {
385 ServerEnvelope::Message(message) => {
386 if self.phase.is_awaiting_welcome()
387 && !matches!(message, ServerMessage::Welcome(_))
388 {
389 return close_for_protocol_error();
390 }
391 commands.extend(server_events::handle_server_message(self, message));
392 }
393 ServerEnvelope::Request {
394 request_id,
395 request,
396 } => {
397 commands.extend(request_flow::handle_server_request(
398 self, request_id, request,
399 ));
400 }
401 ServerEnvelope::Response {
402 response_to,
403 response,
404 } => {
405 commands.extend(request_flow::handle_server_response(
406 self,
407 &response_to,
408 response,
409 ));
410 }
411 }
412 }
413 commands
414 }
415
416 fn accept_welcome(&mut self, payload: WelcomePayload) -> Commands {
417 if !self.phase.is_awaiting_welcome() {
418 return Vec::new();
419 }
420 let WelcomePayload {
421 features,
422 recording,
423 peers,
424 } = payload;
425 self.recovery_delay_ms = INITIAL_RECOVERY_DELAY_MS;
426 self.phase = ProtocolPhase::Authenticated(None);
427
428 let mut commands = vec![
429 Command::SetAvailableFeatures { features },
430 Command::SetRecordingState { state: recording },
431 Command::EmitStateChange {
432 state: self.phase.connection_state(),
433 cause: None,
434 },
435 ];
436 if !peers.is_empty() {
437 commands.push(Command::EmitEvent {
438 event: ProtocolEvent::PeerSnapshot { peers },
439 });
440 }
441 commands.extend(self.replay_session_state());
442 commands
443 }
444
445 pub fn on_transport_ready(&mut self) -> Vec<Command> {
451 if !self.phase.can_enter_connected() {
452 return Vec::new();
453 }
454 self.phase = ProtocolPhase::Connected(None);
455 let mut commands = vec![Command::EmitStateChange {
456 state: self.state(),
457 cause: None,
458 }];
459 commands.extend(self.replay_publication_state());
460 commands
461 }
462
463 pub fn publish(&mut self, stream_type: StreamType, active: bool) -> Vec<Command> {
468 self.sticky_replay.set_publish_active(stream_type, active);
469 if !matches!(&self.phase, ProtocolPhase::Connected(_)) {
470 return Vec::new();
471 }
472 let message = if active {
473 ClientMessage::Publish(StreamIntentPayload { stream_type })
474 } else {
475 ClientMessage::Unpublish(StreamIntentPayload { stream_type })
476 };
477 self.enqueue_client_message(message, FlushMode::Batched)
478 }
479
480 pub fn subscribe(&mut self, user_id: UserId, states: DownloadStates) -> Vec<Command> {
486 self.sticky_replay
487 .remember_subscription_states(&user_id, &states);
488 if !self.phase.can_send_client_messages() {
489 return Vec::new();
490 }
491 self.enqueue_client_message(
492 ClientMessage::Subscribe(SubscribePayload { user_id, states }),
493 FlushMode::Batched,
494 )
495 }
496
497 pub fn update_info(&mut self, info: UserInfo) -> Vec<Command> {
503 self.sticky_replay.remember_info(&info);
504 if !self.phase.can_send_client_messages() {
505 return Vec::new();
506 }
507 self.enqueue_client_message(ClientMessage::Info(info), FlushMode::Batched)
508 }
509
510 pub fn broadcast(&mut self, message: JsonPayload) -> Vec<Command> {
516 if !self.phase.can_send_client_messages() {
517 return Vec::new();
518 }
519 self.enqueue_client_message(
520 ClientMessage::Broadcast(ClientBroadcastPayload { message }),
521 FlushMode::Batched,
522 )
523 }
524
525 pub fn on_timer(&mut self, timer_id: u32) -> Vec<Command> {
531 if timer_id == RECOVERY_TIMER_ID {
532 return connection_lifecycle::handle_recovery_timer(self);
533 }
534 if timer_id == BATCH_FLUSH_TIMER_ID {
535 return self.outbound_batch.flush(false);
536 }
537 if let Some(commands) = RequestTimeoutId::try_from_raw(timer_id)
538 .and_then(|timeout_id| self.request_tracker.resolve_timeout(timeout_id))
539 {
540 return commands;
541 }
542 Vec::new()
543 }
544
545 fn enqueue_client_message(&mut self, message: ClientMessage, mode: FlushMode) -> Commands {
546 let Some(envelope) = ClientEnvelope::Message(message).into_envelope().ok() else {
547 return Vec::new();
548 };
549 self.outbound_batch.enqueue(envelope, mode)
550 }
551
552 fn clear_runtime_state(&mut self) {
553 self.track_users_by_mid.clear();
554 self.outbound_batch.clear();
555 self.request_tracker.clear();
556 }
557
558 fn teardown_runtime_state(&mut self) -> Commands {
564 let mut commands = self.outbound_batch.discard_pending();
565 commands.extend(self.request_tracker.fail_all());
566 if !self.track_users_by_mid.is_empty() {
567 self.track_users_by_mid.clear();
568 commands.push(Command::EmitEvent {
569 event: ProtocolEvent::TrackSnapshot {
570 bindings: Vec::new(),
571 },
572 });
573 }
574 commands
575 }
576
577 fn replay_session_state(&mut self) -> Commands {
579 if !self.phase.can_send_client_messages() {
580 return Vec::new();
581 }
582 let Some(replay_batch) = self.sticky_replay.replay_session_batch() else {
583 return Vec::new();
584 };
585 self.outbound_batch.extend(replay_batch);
586 self.outbound_batch.flush(true)
587 }
588
589 fn replay_publication_state(&mut self) -> Commands {
591 if !self.phase.can_send_client_messages() {
592 return Vec::new();
593 }
594 let replay_batch: Vec<_> = self
595 .sticky_replay
596 .active_publications()
597 .filter_map(|stream_type| {
598 ClientEnvelope::Message(ClientMessage::Publish(StreamIntentPayload { stream_type }))
599 .into_envelope()
600 .ok()
601 })
602 .collect();
603 if replay_batch.is_empty() {
604 return Vec::new();
605 }
606 self.outbound_batch.extend(replay_batch);
607 self.outbound_batch.flush(true)
608 }
609}
610
611fn empty_features() -> AvailableFeatures {
612 AvailableFeatures {
613 rtc: false,
614 transcription: false,
615 audio_recording: false,
616 video_recording: false,
617 }
618}
619
620fn close_for_protocol_error() -> Commands {
621 vec![Command::CloseWebSocket {
622 code: u16::from(WebSocketCloseCode::ProtocolError),
623 }]
624}
625
626fn next_recovery_delay(current_delay_ms: u32) -> u32 {
631 (current_delay_ms.saturating_mul(3) / 2).min(MAX_RECOVERY_DELAY_MS)
632}
633
634#[cfg(test)]
635#[path = "core/TESTS/mod.rs"]
636mod tests;