1use crate::app::AppState;
11use crate::call::active_call::CallParams;
12use crate::config::Config;
13use axum::extract::ws::{Message as AxumMessage, WebSocket};
14use futures::{SinkExt, StreamExt};
15use std::time::Duration;
16use tokio::net::TcpStream;
17use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
18use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
19use tracing::{debug, info, warn};
20
21type PeerStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
22
23const MAX_FORWARD_HOPS: usize = 16;
25
26fn normalize_endpoint(addr: &str) -> String {
28 let mut s = addr.trim();
29 for scheme in ["wss://", "ws://", "https://", "http://"] {
30 if let Some(rest) = s.strip_prefix(scheme) {
31 s = rest;
32 break;
33 }
34 }
35 s.trim_end_matches('/').to_string()
36}
37
38fn addr_already_visited(peer: &str, visited: &[String], self_addr: &str) -> bool {
39 let normalized = normalize_endpoint(peer);
40 if !self_addr.is_empty() && normalize_endpoint(self_addr) == normalized {
41 return true;
42 }
43 visited.iter().any(|v| normalize_endpoint(v) == normalized)
44}
45
46pub async fn try_forward(
51 app_state: &AppState,
52 session_id: &str,
53 params: &CallParams,
54) -> Option<PeerStream> {
55 if app_state.config.peers.is_empty() {
56 return None;
57 }
58
59 let self_addr = app_state.config.http_addr.trim().to_string();
60 let mut visited = params.visited_list();
61 if visited.len() >= MAX_FORWARD_HOPS {
62 warn!(
63 session_id,
64 visited = ?visited,
65 "peer forward exceeded max hops, giving up"
66 );
67 return None;
68 }
69 if !self_addr.is_empty() {
70 visited.push(self_addr.clone());
71 }
72 let visited_str = visited.join(",");
73 let query = params.to_forward_query(&visited_str);
74
75 for peer in &app_state.config.peers {
76 let Some(base) = Config::peer_ws_endpoint(peer) else {
77 warn!(peer, "skipping invalid peer address");
78 continue;
79 };
80 if addr_already_visited(peer, &visited, &self_addr) {
81 debug!(session_id, peer, "skipping already-visited peer");
82 continue;
83 }
84 let url = format!("{}/call?{}", base.trim_end_matches('/'), query);
85 debug!(session_id, %url, "attempting peer forward");
86 match tokio::time::timeout(
87 Duration::from_secs(3),
88 tokio_tungstenite::connect_async(&url),
89 )
90 .await
91 {
92 Ok(Ok((ws, _resp))) => {
93 info!(session_id, %url, "peer accepted forwarded websocket");
94 return Some(ws);
95 }
96 Ok(Err(e)) => {
97 warn!(session_id, %url, "peer forward connection failed: {}", e);
98 }
99 Err(_) => {
100 warn!(session_id, %url, "peer forward timed out");
101 }
102 }
103 }
104 None
105}
106
107pub async fn tunnel(client: WebSocket, peer: PeerStream) {
110 let (mut client_sink, mut client_stream) = client.split();
111 let (mut peer_sink, mut peer_stream) = peer.split();
112
113 let reason = loop {
114 tokio::select! {
115 msg = client_stream.next() => {
116 match msg {
117 Some(Ok(m)) => {
118 if let Some(t) = axum_to_tungstenite(m)
119 && let Err(e) = peer_sink.send(t).await
120 {
121 break format!("forward to peer failed: {}", e);
122 }
123 }
124 Some(Err(e)) => {
125 break format!("client websocket error: {}", e);
126 }
127 None => {
128 break "client websocket closed".to_string();
129 }
130 }
131 }
132 msg = peer_stream.next() => {
133 match msg {
134 Some(Ok(m)) => {
135 if let Some(a) = tungstenite_to_axum(m)
136 && let Err(e) = client_sink.send(a).await
137 {
138 break format!("forward to client failed: {}", e);
139 }
140 }
141 Some(Err(e)) => {
142 break format!("peer websocket error: {}", e);
143 }
144 None => {
145 break "peer websocket closed".to_string();
146 }
147 }
148 }
149 }
150 };
151
152 debug!("websocket tunnel ended: {}", reason);
153 let _ = peer_sink.close().await;
154 let _ = client_sink.close().await;
155}
156
157fn tungstenite_utf8(
158 t: axum::extract::ws::Utf8Bytes,
159) -> tokio_tungstenite::tungstenite::protocol::frame::Utf8Bytes {
160 tokio_tungstenite::tungstenite::protocol::frame::Utf8Bytes::from(t.as_str())
161}
162
163fn axum_to_tungstenite(m: AxumMessage) -> Option<TungsteniteMessage> {
164 Some(match m {
165 AxumMessage::Text(t) => TungsteniteMessage::Text(tungstenite_utf8(t)),
166 AxumMessage::Binary(b) => TungsteniteMessage::Binary(b),
167 AxumMessage::Ping(p) => TungsteniteMessage::Ping(p),
168 AxumMessage::Pong(p) => TungsteniteMessage::Pong(p),
169 AxumMessage::Close(c) => TungsteniteMessage::Close(c.map(|f| {
170 tokio_tungstenite::tungstenite::protocol::frame::CloseFrame {
171 code: f.code.into(),
172 reason: tungstenite_utf8(f.reason),
173 }
174 })),
175 })
176}
177
178fn tungstenite_to_axum(m: TungsteniteMessage) -> Option<AxumMessage> {
179 match m {
180 TungsteniteMessage::Text(t) => Some(AxumMessage::Text(
181 axum::extract::ws::Utf8Bytes::from(t.as_str()),
182 )),
183 TungsteniteMessage::Binary(b) => Some(AxumMessage::Binary(b)),
184 TungsteniteMessage::Ping(p) => Some(AxumMessage::Ping(p)),
185 TungsteniteMessage::Pong(p) => Some(AxumMessage::Pong(p)),
186 TungsteniteMessage::Close(c) => Some(AxumMessage::Close(c.map(|f| {
187 axum::extract::ws::CloseFrame {
188 code: f.code.into(),
189 reason: axum::extract::ws::Utf8Bytes::from(f.reason.as_str()),
190 }
191 }))),
192 TungsteniteMessage::Frame(_) => None,
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199
200 #[test]
201 fn test_peer_ws_endpoint_normalization() {
202 assert_eq!(
203 Config::peer_ws_endpoint("10.0.0.2:8080").as_deref(),
204 Some("ws://10.0.0.2:8080")
205 );
206 assert_eq!(
207 Config::peer_ws_endpoint("ws://10.0.0.2:8080").as_deref(),
208 Some("ws://10.0.0.2:8080")
209 );
210 assert_eq!(
211 Config::peer_ws_endpoint("http://10.0.0.2:8080").as_deref(),
212 Some("ws://10.0.0.2:8080")
213 );
214 assert_eq!(
215 Config::peer_ws_endpoint("https://10.0.0.2:8080").as_deref(),
216 Some("wss://10.0.0.2:8080")
217 );
218 assert_eq!(Config::peer_ws_endpoint(""), None);
219 }
220
221 #[test]
222 fn test_call_params_forward_query() {
223 let params = CallParams {
224 id: Some("s.a/b c".to_string()),
225 dump_events: Some(true),
226 ping_interval: Some(30),
227 server_side_track: Some("t.1".to_string()),
228 forward: None,
229 visited: None,
230 };
231 let q = params.to_forward_query("node1,node2");
232 assert!(q.contains("id=s.a%2Fb%20c"));
233 assert!(q.contains("dump=true"));
234 assert!(q.contains("ping=30"));
235 assert!(q.contains("server_side_track=t.1"));
236 assert!(q.contains("forward=true"));
237 assert!(q.contains("visited=node1%2Cnode2"));
238 }
239
240 #[test]
241 fn test_visited_list_parsing() {
242 let params = CallParams {
243 id: None,
244 dump_events: None,
245 ping_interval: None,
246 server_side_track: None,
247 forward: None,
248 visited: Some("node1, node2,node3".to_string()),
249 };
250 assert_eq!(params.visited_list(), vec!["node1", "node2", "node3"]);
251 }
252
253 #[test]
254 fn test_addr_already_visited() {
255 let visited = vec!["ws://10.0.0.2:8080".to_string()];
256 assert!(addr_already_visited("10.0.0.2:8080", &visited, "0.0.0.0:8080"));
257 assert!(addr_already_visited("http://10.0.0.2:8080/", &visited, "0.0.0.0:8080"));
258 assert!(!addr_already_visited("10.0.0.3:8080", &visited, "0.0.0.0:8080"));
259 assert!(addr_already_visited("0.0.0.0:8080", &[], "0.0.0.0:8080"));
260 }
261}