Skip to main content

vox_rtc_server/
gateway.rs

1use crate::{
2    ControlledSession, VoxRtcControlSession, VoxRtcServerClient, VoxRtcServerClientOptions,
3    WireEvent,
4};
5use axum::{
6    Router,
7    extract::{
8        State, WebSocketUpgrade,
9        ws::{CloseFrame, Message, WebSocket},
10    },
11    http::{HeaderMap, StatusCode},
12    response::{IntoResponse, Response},
13    routing::get,
14};
15use futures_util::{SinkExt, StreamExt};
16use serde_json::{Map, Value, json};
17use std::{
18    future::Future,
19    pin::Pin,
20    sync::{
21        Arc,
22        atomic::{AtomicBool, AtomicUsize, Ordering},
23    },
24};
25use tokio::sync::{Notify, mpsc, watch};
26
27const DEFAULT_PATH: &str = "/api/vox/rtc";
28const MAX_SAFE_INTEGER: u64 = 9_007_199_254_740_991;
29
30pub type GatewayHookFuture = Pin<Box<dyn Future<Output = std::result::Result<(), String>> + Send>>;
31pub type SessionCreatedHook = Arc<dyn Fn(GatewaySessionContext) -> GatewayHookFuture + Send + Sync>;
32pub type SessionClosedHook = Arc<dyn Fn(GatewayClosedContext) -> GatewayHookFuture + Send + Sync>;
33pub type GatewayErrorHook = Arc<dyn Fn(String) + Send + Sync>;
34
35#[derive(Clone)]
36pub struct GatewaySessionContext {
37    pub headers: HeaderMap,
38    pub session: VoxRtcControlSession,
39}
40
41#[derive(Clone)]
42pub struct GatewayClosedContext {
43    pub headers: HeaderMap,
44    pub session: VoxRtcControlSession,
45    pub reason: String,
46}
47
48pub struct GatewayOptions {
49    pub vox_http_base: String,
50    pub api_key: Option<String>,
51    pub path: String,
52    pub on_session_created: Option<SessionCreatedHook>,
53    pub on_session_closed: Option<SessionClosedHook>,
54    pub on_error: Option<GatewayErrorHook>,
55}
56
57impl GatewayOptions {
58    pub fn new(vox_http_base: impl Into<String>) -> Self {
59        Self {
60            vox_http_base: vox_http_base.into(),
61            api_key: None,
62            path: DEFAULT_PATH.to_owned(),
63            on_session_created: None,
64            on_session_closed: None,
65            on_error: None,
66        }
67    }
68}
69
70#[derive(Clone)]
71pub struct VoxRtcGateway {
72    state: Arc<GatewayState>,
73}
74
75struct GatewayState {
76    client: VoxRtcServerClient,
77    options: GatewayOptions,
78    shutting_down: AtomicBool,
79    shutdown: watch::Sender<Option<String>>,
80    active: AtomicUsize,
81    drained: Notify,
82}
83
84struct ActiveGuard(Arc<GatewayState>);
85
86impl Drop for ActiveGuard {
87    fn drop(&mut self) {
88        if self.0.active.fetch_sub(1, Ordering::AcqRel) == 1 {
89            self.0.drained.notify_waiters();
90        }
91    }
92}
93
94#[derive(Debug)]
95struct ClientMessage {
96    id: String,
97    kind: String,
98    data: Map<String, Value>,
99}
100
101#[derive(Debug)]
102struct CommandError {
103    message: String,
104    typed: bool,
105}
106
107impl CommandError {
108    fn plain(message: impl Into<String>) -> Self {
109        Self {
110            message: message.into(),
111            typed: false,
112        }
113    }
114
115    fn invalid(message: impl Into<String>) -> Self {
116        Self {
117            message: message.into(),
118            typed: true,
119        }
120    }
121}
122
123impl VoxRtcGateway {
124    pub fn new(options: GatewayOptions) -> crate::Result<Self> {
125        let mut client_options = VoxRtcServerClientOptions::new(&options.vox_http_base);
126        client_options.api_key = options.api_key.clone();
127        let client = VoxRtcServerClient::with_options(client_options)?;
128        Ok(Self::with_client(options, client))
129    }
130
131    pub fn with_client(options: GatewayOptions, client: VoxRtcServerClient) -> Self {
132        let (shutdown, _) = watch::channel(None);
133        Self {
134            state: Arc::new(GatewayState {
135                client,
136                options,
137                shutting_down: AtomicBool::new(false),
138                shutdown,
139                active: AtomicUsize::new(0),
140                drained: Notify::new(),
141            }),
142        }
143    }
144
145    pub fn router(&self) -> Router {
146        let path = normalize_path(&self.state.options.path);
147        Router::new()
148            .route(&path, get(upgrade))
149            .with_state(self.state.clone())
150    }
151
152    pub async fn close(&self, reason: impl Into<String>) {
153        if !self.state.shutting_down.swap(true, Ordering::AcqRel) {
154            let _ = self.state.shutdown.send(Some(reason.into()));
155        }
156        loop {
157            let notified = self.state.drained.notified();
158            tokio::pin!(notified);
159            notified.as_mut().enable();
160            if self.state.active.load(Ordering::Acquire) == 0 {
161                break;
162            }
163            notified.await;
164        }
165        self.state.client.disconnect().await;
166    }
167}
168
169async fn upgrade(
170    State(state): State<Arc<GatewayState>>,
171    headers: HeaderMap,
172    socket: WebSocketUpgrade,
173) -> Response {
174    if state.shutting_down.load(Ordering::Acquire) {
175        return (
176            StatusCode::SERVICE_UNAVAILABLE,
177            "RTC gateway is shutting down",
178        )
179            .into_response();
180    }
181    socket
182        .on_upgrade(move |socket| serve(socket, headers, state))
183        .into_response()
184}
185
186async fn serve(socket: WebSocket, headers: HeaderMap, state: Arc<GatewayState>) {
187    state.active.fetch_add(1, Ordering::AcqRel);
188    let _active = ActiveGuard(state.clone());
189    let controlled = match state.client.create_controlled_session().await {
190        Ok(value) => value,
191        Err(error) => {
192            report(&state, error.to_string());
193            return;
194        }
195    };
196    let context = GatewaySessionContext {
197        headers: headers.clone(),
198        session: controlled.session.clone(),
199    };
200    let (event_tx, mut event_rx) = mpsc::unbounded_channel();
201    let _listener = controlled.session.on_event(move |event| {
202        let _ = event_tx.send(event);
203    });
204    let mut shutdown = state.shutdown.subscribe();
205    let (mut sender, mut receiver) = socket.split();
206    let shutdown_reason = { shutdown.borrow().clone() };
207    if let Some(reason) = shutdown_reason {
208        cleanup(&state, &headers, &controlled.session, false, &reason).await;
209        let _ = sender
210            .send(Message::Close(Some(CloseFrame {
211                code: 1001,
212                reason: reason.into(),
213            })))
214            .await;
215        return;
216    }
217    if let Some(hook) = &state.options.on_session_created
218        && let Err(error) = hook(context.clone()).await
219    {
220        report(&state, error.clone());
221        let gateway_error = CommandError::plain(error);
222        let _ = send_error(&mut sender, None, &gateway_error).await;
223        cleanup(
224            &state,
225            &headers,
226            &controlled.session,
227            false,
228            "session_created_hook_failed",
229        )
230        .await;
231        let _ = sender
232            .send(Message::Close(Some(CloseFrame {
233                code: 1011,
234                reason: "RTC gateway setup failed".into(),
235            })))
236            .await;
237        return;
238    }
239
240    if send_json(
241        &mut sender,
242        None,
243        "gateway.ready",
244        bootstrap_json(&controlled),
245    )
246    .await
247    .is_err()
248    {
249        cleanup(
250            &state,
251            &headers,
252            &controlled.session,
253            false,
254            "gateway_ready_failed",
255        )
256        .await;
257        return;
258    }
259    let mut pending_offer_id: Option<String> = None;
260    let mut generated_offer = false;
261    let mut rtc_close_requested = false;
262    let mut reason = "browser_disconnected".to_owned();
263
264    loop {
265        tokio::select! {
266            changed = shutdown.changed() => {
267                if changed.is_ok()
268                    && let Some(value) = shutdown.borrow().clone()
269                {
270                    reason = value;
271                    break;
272                }
273            }
274            message = receiver.next() => {
275                match message {
276                    Some(Ok(Message::Text(text))) => {
277                        match parse_message(text.as_str()) {
278                            Ok(message) => {
279                                if let Err(error) = handle_message(
280                                    &controlled.session,
281                                    &message,
282                                    &mut pending_offer_id,
283                                    &mut generated_offer,
284                                    &mut rtc_close_requested,
285                                ).await {
286                                    if pending_offer_id.as_deref() == Some(message.id.as_str()) {
287                                        pending_offer_id = None;
288                                    }
289                                    if send_error(&mut sender, Some(&message.id), &error).await.is_err() {
290                                        break;
291                                    }
292                                }
293                            }
294                            Err(error) => {
295                                if send_error(&mut sender, None, &error).await.is_err() {
296                                    break;
297                                }
298                            }
299                        }
300                    }
301                    Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
302                    Some(Ok(_)) => {
303                        let error = CommandError::plain("RTC gateway accepts text JSON only");
304                        if send_error(&mut sender, None, &error).await.is_err() {
305                            break;
306                        }
307                    }
308                }
309            }
310            event = event_rx.recv() => {
311                let Some(event) = event else { break; };
312                let correlates = event.r#type == "rtc.answer" || event.r#type == "rtc.signaling_error";
313                let request_id = if correlates { pending_offer_id.take() } else { None };
314                let closed_reason = if event.r#type == "rtc.session.closed" {
315                    event_reason(&event)
316                } else {
317                    None
318                };
319                if send_json(&mut sender, request_id.as_deref(), &event.r#type, Value::Object(event.data)).await.is_err() {
320                    break;
321                }
322                if event.r#type == "rtc.session.closed" {
323                    rtc_close_requested = true;
324                    reason = closed_reason.unwrap_or_else(|| "session_closed".to_owned());
325                    let _ = sender.send(Message::Close(Some(CloseFrame { code: 1000, reason: reason.clone().into() }))).await;
326                    break;
327                }
328            }
329        }
330    }
331    cleanup(
332        &state,
333        &headers,
334        &controlled.session,
335        rtc_close_requested,
336        &reason,
337    )
338    .await;
339}
340
341async fn handle_message(
342    session: &VoxRtcControlSession,
343    message: &ClientMessage,
344    pending_offer_id: &mut Option<String>,
345    generated_offer: &mut bool,
346    rtc_close_requested: &mut bool,
347) -> std::result::Result<(), CommandError> {
348    match message.kind.as_str() {
349        "rtc.offer" => {
350            if pending_offer_id.is_some() {
351                return Err(CommandError::plain("An RTC offer is already pending"));
352            }
353            let generation = generation(&message.data, "rtc.offer", false)?;
354            let offer = offer(message.data.get("offer"))?;
355            *pending_offer_id = Some(message.id.clone());
356            session
357                .send_offer(
358                    offer,
359                    message.data.get("restart") == Some(&Value::Bool(true)),
360                    generation,
361                )
362                .await
363                .map_err(|error| CommandError::plain(error.to_string()))?;
364            *generated_offer = generation.is_some();
365            Ok(())
366        }
367        "rtc.ice_candidate" => {
368            let generation = generation(&message.data, "rtc.ice_candidate", *generated_offer)?;
369            let candidate = candidate(message.data.get("candidate"))?;
370            session
371                .send_ice_candidate(candidate, generation)
372                .await
373                .map_err(|error| CommandError::plain(error.to_string()))
374        }
375        "rtc.close" => {
376            *rtc_close_requested = true;
377            let reason = message
378                .data
379                .get("reason")
380                .and_then(Value::as_str)
381                .unwrap_or("client_closed");
382            session
383                .close_rtc(reason)
384                .await
385                .map_err(|error| CommandError::plain(error.to_string()))
386        }
387        value => Err(CommandError::plain(format!(
388            "Unsupported RTC gateway message type: {value}"
389        ))),
390    }
391}
392
393async fn cleanup(
394    state: &Arc<GatewayState>,
395    headers: &HeaderMap,
396    session: &VoxRtcControlSession,
397    rtc_close_requested: bool,
398    reason: &str,
399) {
400    if !rtc_close_requested && let Err(error) = session.close_rtc(reason).await {
401        report(state, error.to_string());
402    }
403    if let Err(error) = session.close().await {
404        report(state, error.to_string());
405    }
406    if let Some(hook) = &state.options.on_session_closed
407        && let Err(error) = hook(GatewayClosedContext {
408            headers: headers.clone(),
409            session: session.clone(),
410            reason: reason.to_owned(),
411        })
412        .await
413    {
414        report(state, error);
415    }
416}
417
418fn parse_message(text: &str) -> std::result::Result<ClientMessage, CommandError> {
419    let value: Value = serde_json::from_str(text)
420        .map_err(|_| CommandError::plain("RTC gateway message must be valid JSON"))?;
421    let object = value
422        .as_object()
423        .ok_or_else(|| CommandError::plain("RTC gateway message must be an object"))?;
424    let id = object
425        .get("id")
426        .and_then(Value::as_str)
427        .map(str::trim)
428        .filter(|value| !value.is_empty());
429    let kind = object
430        .get("type")
431        .and_then(Value::as_str)
432        .map(str::trim)
433        .filter(|value| !value.is_empty());
434    let data = object.get("data").and_then(Value::as_object);
435    match (id, kind, data) {
436        (Some(id), Some(kind), Some(data)) => Ok(ClientMessage {
437            id: id.to_owned(),
438            kind: kind.to_owned(),
439            data: data.clone(),
440        }),
441        _ => Err(CommandError::plain(
442            "RTC gateway message requires id, type, and object data",
443        )),
444    }
445}
446
447fn generation(
448    data: &Map<String, Value>,
449    command: &str,
450    required: bool,
451) -> std::result::Result<Option<u64>, CommandError> {
452    let Some(value) = data.get("generation") else {
453        if required {
454            return Err(CommandError::invalid(format!(
455                "{command} requires generation for a generated RTC negotiation"
456            )));
457        }
458        return Ok(None);
459    };
460    let value = value
461        .as_u64()
462        .filter(|value| *value > 0 && *value <= MAX_SAFE_INTEGER)
463        .ok_or_else(|| {
464            CommandError::invalid(format!(
465                "{command} generation must be a positive safe integer"
466            ))
467        })?;
468    Ok(Some(value))
469}
470
471fn offer(value: Option<&Value>) -> std::result::Result<Value, CommandError> {
472    let object = value
473        .and_then(Value::as_object)
474        .ok_or_else(|| CommandError::plain("rtc.offer requires a non-empty SDP offer"))?;
475    if object.get("type").and_then(Value::as_str) != Some("offer")
476        || object
477            .get("sdp")
478            .and_then(Value::as_str)
479            .map(str::trim)
480            .unwrap_or_default()
481            .is_empty()
482    {
483        return Err(CommandError::plain(
484            "rtc.offer requires a non-empty SDP offer",
485        ));
486    }
487    Ok(json!({ "type": "offer", "sdp": object["sdp"] }))
488}
489
490fn candidate(value: Option<&Value>) -> std::result::Result<Option<Value>, CommandError> {
491    let Some(value) = value else {
492        return Ok(None);
493    };
494    if value.is_null() {
495        return Ok(None);
496    }
497    let object = value.as_object().ok_or_else(|| {
498        CommandError::plain("rtc.ice_candidate requires a candidate object or null")
499    })?;
500    let text = object
501        .get("candidate")
502        .and_then(Value::as_str)
503        .ok_or_else(|| {
504            CommandError::plain("rtc.ice_candidate requires a candidate object or null")
505        })?;
506    Ok(Some(json!({
507        "candidate": text,
508        "sdpMid": object.get("sdpMid").cloned().unwrap_or(Value::Null),
509        "sdpMLineIndex": object.get("sdpMLineIndex").cloned().unwrap_or(Value::Null),
510        "usernameFragment": object.get("usernameFragment").cloned().unwrap_or(Value::Null),
511    })))
512}
513
514fn bootstrap_json(controlled: &ControlledSession) -> Value {
515    json!({
516        "session": {
517            "sessionId": controlled.bootstrap.session_id,
518            "expiresAt": controlled.bootstrap.expires_at,
519            "attachTtlSeconds": controlled.bootstrap.attach_ttl_seconds,
520            "iceServers": controlled.bootstrap.ice_servers,
521        }
522    })
523}
524
525fn event_reason(event: &WireEvent) -> Option<String> {
526    event
527        .data
528        .get("reason")
529        .and_then(Value::as_str)
530        .map(ToOwned::to_owned)
531}
532
533async fn send_error<S>(
534    sender: &mut S,
535    id: Option<&str>,
536    error: &CommandError,
537) -> std::result::Result<(), ()>
538where
539    S: futures_util::Sink<Message> + Unpin,
540{
541    let mut data = json!({ "message": error.message });
542    if error.typed {
543        data["code"] = Value::String("command_invalid".to_owned());
544    }
545    send_json(sender, id, "gateway.error", data).await
546}
547
548async fn send_json<S>(
549    sender: &mut S,
550    id: Option<&str>,
551    kind: &str,
552    data: Value,
553) -> std::result::Result<(), ()>
554where
555    S: futures_util::Sink<Message> + Unpin,
556{
557    let mut payload = json!({ "type": kind, "data": data });
558    if let Some(id) = id {
559        payload["id"] = Value::String(id.to_owned());
560    }
561    sender
562        .send(Message::Text(payload.to_string().into()))
563        .await
564        .map_err(|_| ())
565}
566
567fn normalize_path(path: &str) -> String {
568    let trimmed = path.trim().trim_matches('/');
569    if trimmed.is_empty() {
570        "/".to_owned()
571    } else {
572        format!("/{trimmed}")
573    }
574}
575
576fn report(state: &GatewayState, error: String) {
577    if let Some(hook) = &state.options.on_error {
578        hook(error);
579    }
580}
581
582#[cfg(test)]
583mod tests {
584    use super::*;
585
586    #[test]
587    fn parses_generated_offer_and_candidate_generations_verbatim() {
588        let offer = parse_message(
589            r#"{"id":"offer-1","type":"rtc.offer","data":{"offer":{"type":"offer","sdp":"sdp"},"generation":1}}"#,
590        )
591        .expect("valid offer");
592        assert_eq!(
593            generation(&offer.data, "rtc.offer", false).unwrap(),
594            Some(1)
595        );
596
597        let candidate_message = parse_message(
598            r#"{"id":"candidate-1","type":"rtc.ice_candidate","data":{"candidate":null,"generation":2}}"#,
599        )
600        .expect("valid candidate");
601        assert_eq!(
602            generation(&candidate_message.data, "rtc.ice_candidate", true).unwrap(),
603            Some(2)
604        );
605        assert_eq!(
606            candidate(candidate_message.data.get("candidate")).unwrap(),
607            None
608        );
609    }
610
611    #[test]
612    fn generated_negotiation_requires_candidate_generation() {
613        let data = Map::from_iter([("candidate".to_owned(), Value::Null)]);
614        let error = generation(&data, "rtc.ice_candidate", true).unwrap_err();
615        assert!(error.typed);
616        assert_eq!(
617            error.message,
618            "rtc.ice_candidate requires generation for a generated RTC negotiation"
619        );
620    }
621
622    #[test]
623    fn rejects_malformed_generation_but_preserves_legacy_mode() {
624        let legacy = Map::new();
625        assert_eq!(generation(&legacy, "rtc.offer", false).unwrap(), None);
626
627        for value in [
628            json!(0),
629            json!(1.5),
630            json!("2"),
631            json!(9007199254740992_u64),
632        ] {
633            let data = Map::from_iter([("generation".to_owned(), value)]);
634            let error = generation(&data, "rtc.offer", false).unwrap_err();
635            assert!(error.typed);
636            assert!(error.message.contains("positive safe integer"));
637        }
638    }
639
640    #[test]
641    fn validates_offer_and_candidate_shape() {
642        assert_eq!(
643            offer(Some(&json!({"type": "offer", "sdp": "abc"}))).unwrap(),
644            json!({"type": "offer", "sdp": "abc"})
645        );
646        assert!(offer(Some(&json!({"type": "answer", "sdp": "abc"}))).is_err());
647        assert!(candidate(Some(&json!({"candidate": 3}))).is_err());
648        assert_eq!(
649            candidate(Some(&json!({
650                "candidate": "candidate:one",
651                "sdpMid": "audio",
652                "sdpMLineIndex": 0
653            })))
654            .unwrap()
655            .unwrap()["sdpMLineIndex"],
656            json!(0)
657        );
658    }
659
660    #[tokio::test]
661    async fn shutdown_waiter_cannot_miss_the_final_active_session() {
662        let options = GatewayOptions::new("http://vox.test");
663        let client = VoxRtcServerClient::new("http://vox.test").unwrap();
664        let gateway = VoxRtcGateway::with_client(options, client);
665        gateway.state.active.store(1, Ordering::Release);
666        let guard = ActiveGuard(gateway.state.clone());
667        let closing = tokio::spawn({
668            let gateway = gateway.clone();
669            async move { gateway.close("test_shutdown").await }
670        });
671
672        tokio::task::yield_now().await;
673        drop(guard);
674        tokio::time::timeout(std::time::Duration::from_secs(1), closing)
675            .await
676            .expect("gateway close completed")
677            .expect("close task succeeded");
678    }
679}