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}