active_call/handler/
peer.rs1use crate::app::AppState;
14use crate::call::active_call::CallParams;
15use crate::config::Config;
16use axum::extract::ws::{Message as AxumMessage, WebSocket};
17use futures::{SinkExt, StreamExt};
18use std::time::Duration;
19use tokio::net::TcpStream;
20use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
21use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
22use tracing::{debug, info, warn};
23
24type PeerStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
25
26pub async fn try_forward(
32 app_state: &AppState,
33 session_id: &str,
34 params: &CallParams,
35) -> Option<PeerStream> {
36 if params.forward.is_some() {
37 debug!(
38 session_id,
39 "forward is set; refusing to hop to another peer"
40 );
41 return None;
42 }
43 if app_state.config.peers.is_empty() {
44 return None;
45 }
46
47 let query = params.to_forward_query();
48 for peer in &app_state.config.peers {
49 let Some(base) = Config::peer_ws_endpoint(peer) else {
50 warn!(peer, "skipping invalid peer address");
51 continue;
52 };
53 let url = format!("{}/call?{}", base.trim_end_matches('/'), query);
54 debug!(session_id, %url, "attempting peer forward");
55 match tokio::time::timeout(
56 Duration::from_secs(3),
57 tokio_tungstenite::connect_async(&url),
58 )
59 .await
60 {
61 Ok(Ok((ws, _resp))) => {
62 info!(session_id, %url, "peer accepted forwarded websocket");
63 return Some(ws);
64 }
65 Ok(Err(e)) => {
66 warn!(session_id, %url, "peer forward connection failed: {}", e);
67 }
68 Err(_) => {
69 warn!(session_id, %url, "peer forward timed out");
70 }
71 }
72 }
73 None
74}
75
76pub async fn tunnel(client: WebSocket, peer: PeerStream) {
79 let (mut client_sink, mut client_stream) = client.split();
80 let (mut peer_sink, mut peer_stream) = peer.split();
81
82 let reason = loop {
83 tokio::select! {
84 msg = client_stream.next() => {
85 match msg {
86 Some(Ok(m)) => {
87 if let Some(t) = axum_to_tungstenite(m)
88 && let Err(e) = peer_sink.send(t).await
89 {
90 break format!("forward to peer failed: {}", e);
91 }
92 }
93 Some(Err(e)) => {
94 break format!("client websocket error: {}", e);
95 }
96 None => {
97 break "client websocket closed".to_string();
98 }
99 }
100 }
101 msg = peer_stream.next() => {
102 match msg {
103 Some(Ok(m)) => {
104 if let Some(a) = tungstenite_to_axum(m)
105 && let Err(e) = client_sink.send(a).await
106 {
107 break format!("forward to client failed: {}", e);
108 }
109 }
110 Some(Err(e)) => {
111 break format!("peer websocket error: {}", e);
112 }
113 None => {
114 break "peer websocket closed".to_string();
115 }
116 }
117 }
118 }
119 };
120
121 debug!("websocket tunnel ended: {}", reason);
122 let _ = peer_sink.close().await;
123 let _ = client_sink.close().await;
124}
125
126fn tungstenite_utf8(
127 t: axum::extract::ws::Utf8Bytes,
128) -> tokio_tungstenite::tungstenite::protocol::frame::Utf8Bytes {
129 tokio_tungstenite::tungstenite::protocol::frame::Utf8Bytes::from(t.as_str())
130}
131
132fn axum_to_tungstenite(m: AxumMessage) -> Option<TungsteniteMessage> {
133 Some(match m {
134 AxumMessage::Text(t) => TungsteniteMessage::Text(tungstenite_utf8(t)),
135 AxumMessage::Binary(b) => TungsteniteMessage::Binary(b),
136 AxumMessage::Ping(p) => TungsteniteMessage::Ping(p),
137 AxumMessage::Pong(p) => TungsteniteMessage::Pong(p),
138 AxumMessage::Close(c) => TungsteniteMessage::Close(c.map(|f| {
139 tokio_tungstenite::tungstenite::protocol::frame::CloseFrame {
140 code: f.code.into(),
141 reason: tungstenite_utf8(f.reason),
142 }
143 })),
144 })
145}
146
147fn tungstenite_to_axum(m: TungsteniteMessage) -> Option<AxumMessage> {
148 match m {
149 TungsteniteMessage::Text(t) => Some(AxumMessage::Text(axum::extract::ws::Utf8Bytes::from(
150 t.as_str(),
151 ))),
152 TungsteniteMessage::Binary(b) => Some(AxumMessage::Binary(b)),
153 TungsteniteMessage::Ping(p) => Some(AxumMessage::Ping(p)),
154 TungsteniteMessage::Pong(p) => Some(AxumMessage::Pong(p)),
155 TungsteniteMessage::Close(c) => Some(AxumMessage::Close(c.map(|f| {
156 axum::extract::ws::CloseFrame {
157 code: f.code.into(),
158 reason: axum::extract::ws::Utf8Bytes::from(f.reason.as_str()),
159 }
160 }))),
161 TungsteniteMessage::Frame(_) => None,
162 }
163}
164
165#[cfg(test)]
166mod tests {
167 use super::*;
168
169 #[test]
170 fn test_peer_ws_endpoint_normalization() {
171 assert_eq!(
172 Config::peer_ws_endpoint("10.0.0.2:8080").as_deref(),
173 Some("ws://10.0.0.2:8080")
174 );
175 assert_eq!(
176 Config::peer_ws_endpoint("ws://10.0.0.2:8080").as_deref(),
177 Some("ws://10.0.0.2:8080")
178 );
179 assert_eq!(
180 Config::peer_ws_endpoint("http://10.0.0.2:8080").as_deref(),
181 Some("ws://10.0.0.2:8080")
182 );
183 assert_eq!(
184 Config::peer_ws_endpoint("https://10.0.0.2:8080").as_deref(),
185 Some("wss://10.0.0.2:8080")
186 );
187 assert_eq!(Config::peer_ws_endpoint(""), None);
188 }
189
190 #[test]
191 fn test_call_params_forward_query() {
192 let params = CallParams {
193 id: Some("s.a/b c".to_string()),
194 dump_events: Some(true),
195 ping_interval: Some(30),
196 server_side_track: Some("t.1".to_string()),
197 forward: None,
198 visited: None,
199 };
200 let q = params.to_forward_query();
201 assert!(q.contains("id=s.a%2Fb%20c"));
202 assert!(q.contains("dump=true"));
203 assert!(q.contains("ping=30"));
204 assert!(q.contains("server_side_track=t.1"));
205 assert!(q.contains("forward=true"));
206 assert!(!q.contains("visited="));
207 }
208}