1use std::sync::Arc;
23
24use bytes::Bytes;
25use http::HeaderMap;
26use http_body_util::Full;
27use hyper::upgrade::OnUpgrade;
28use hyper::{Request, Response};
29use hyper_util::rt::TokioIo;
30use tokio::net::TcpStream;
31use tracing::{debug, warn};
32
33#[derive(Debug, thiserror::Error)]
37pub enum WsError {
38 #[error("failed to connect to upstream {0}: {1}")]
39 Connect(String, String),
40 #[error("upstream websocket handshake failed: {0}")]
41 Handshake(String),
42 #[error("upstream rejected the websocket upgrade (status {0})")]
43 UpstreamRejected(u16),
44 #[error("upstream upgrade failed: {0}")]
45 Upgrade(String),
46 #[error("upstream TLS setup failed: {0}")]
47 Tls(String),
48}
49
50pub fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
54 let connection_upgrade = headers
55 .get(http::header::CONNECTION)
56 .and_then(|v| v.to_str().ok())
57 .map(|v| {
58 v.split(',')
59 .any(|token| token.trim().eq_ignore_ascii_case("upgrade"))
60 })
61 .unwrap_or(false);
62
63 let upgrade_websocket = headers
64 .get(http::header::UPGRADE)
65 .and_then(|v| v.to_str().ok())
66 .map(|v| v.eq_ignore_ascii_case("websocket"))
67 .unwrap_or(false);
68
69 connection_upgrade && upgrade_websocket
70}
71
72pub fn is_h2_websocket_connect(method: &http::Method, extensions: &http::Extensions) -> bool {
78 method == http::Method::CONNECT
79 && extensions
80 .get::<hyper::ext::Protocol>()
81 .is_some_and(|p| p.as_str().eq_ignore_ascii_case("websocket"))
82}
83
84fn synthesize_ws_key() -> String {
88 use base64::{engine::general_purpose::STANDARD, Engine};
89 use ring::rand::{SecureRandom, SystemRandom};
90 let mut buf = [0u8; 16];
91 SystemRandom::new()
92 .fill(&mut buf)
93 .expect("system RNG unavailable");
94 STANDARD.encode(buf)
95}
96
97const FORWARD_HEADERS: &[&str] = &[
100 "upgrade",
101 "connection",
102 "sec-websocket-key",
103 "sec-websocket-version",
104 "sec-websocket-protocol",
105 "sec-websocket-extensions",
106];
107
108const ECHO_HEADERS: &[&str] = &[
110 "upgrade",
111 "connection",
112 "sec-websocket-accept",
113 "sec-websocket-protocol",
114 "sec-websocket-extensions",
115];
116
117#[allow(clippy::too_many_arguments)]
142pub async fn proxy_upgrade(
143 host: String,
144 port: u16,
145 path: String,
146 tls: bool,
147 verify: bool,
148 tls_identity: Option<Arc<crate::outbound::tls::UpstreamTls>>,
149 fwd_headers: &std::collections::HashMap<String, Vec<String>>,
150 client_on_upgrade: OnUpgrade,
151 client_is_h2: bool,
152) -> Result<Response<Full<Bytes>>, WsError> {
153 let tcp = TcpStream::connect((host.as_str(), port))
158 .await
159 .map_err(|e| WsError::Connect(format!("{}:{}", host, port), e.to_string()))?;
160
161 let mut sender = if tls {
162 let connector = crate::outbound::client_tls_connector(verify, tls_identity.as_ref())
163 .map_err(WsError::Tls)?;
164 let server_name = rustls::pki_types::ServerName::try_from(host.clone())
165 .map_err(|e| WsError::Tls(format!("invalid server name '{}': {}", host, e)))?;
166 let tls_stream = connector
167 .connect(server_name, tcp)
168 .await
169 .map_err(|e| WsError::Tls(e.to_string()))?;
170 let (sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(tls_stream))
171 .await
172 .map_err(|e| WsError::Handshake(e.to_string()))?;
173 tokio::spawn(async move {
174 if let Err(e) = conn.with_upgrades().await {
175 debug!("upstream wss connection closed: {}", e);
176 }
177 });
178 sender
179 } else {
180 let (sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(tcp))
181 .await
182 .map_err(|e| WsError::Handshake(e.to_string()))?;
183 tokio::spawn(async move {
184 if let Err(e) = conn.with_upgrades().await {
185 debug!("upstream ws connection closed: {}", e);
186 }
187 });
188 sender
189 };
190
191 let mut builder = Request::builder()
196 .method(http::Method::GET)
197 .uri(&path)
198 .header(http::header::HOST, format!("{}:{}", host, port));
199 for name in FORWARD_HEADERS {
200 if client_is_h2
205 && matches!(
206 *name,
207 "upgrade" | "connection" | "sec-websocket-key" | "sec-websocket-version"
208 )
209 {
210 continue;
211 }
212 if let Some(values) = fwd_headers.get(*name) {
213 for value in values {
214 builder = builder.header(*name, value);
215 }
216 }
217 }
218 if client_is_h2 {
219 builder = builder
220 .header("upgrade", "websocket")
221 .header("connection", "upgrade")
222 .header("sec-websocket-key", synthesize_ws_key());
223 if !fwd_headers.contains_key("sec-websocket-version") {
224 builder = builder.header("sec-websocket-version", "13");
225 }
226 }
227 let req = builder
228 .body(Full::new(Bytes::new()))
229 .map_err(|e| WsError::Handshake(e.to_string()))?;
230
231 let resp = sender
233 .send_request(req)
234 .await
235 .map_err(|e| WsError::Handshake(e.to_string()))?;
236 if resp.status() != http::StatusCode::SWITCHING_PROTOCOLS {
237 return Err(WsError::UpstreamRejected(resp.status().as_u16()));
238 }
239
240 let client_response = if client_is_h2 {
246 let mut b = Response::builder().status(http::StatusCode::OK);
247 if let Some(value) = resp.headers().get("sec-websocket-protocol") {
249 b = b.header("sec-websocket-protocol", value);
250 }
251 b
252 } else {
253 let mut b = Response::builder().status(http::StatusCode::SWITCHING_PROTOCOLS);
254 for name in ECHO_HEADERS {
255 if let Some(value) = resp.headers().get(*name) {
256 b = b.header(*name, value);
257 }
258 }
259 b
260 };
261
262 let upstream_upgraded = hyper::upgrade::on(resp)
263 .await
264 .map_err(|e| WsError::Upgrade(e.to_string()))?;
265
266 let client_response = client_response
267 .body(Full::new(Bytes::new()))
268 .map_err(|e| WsError::Handshake(e.to_string()))?;
269
270 tokio::spawn(async move {
272 match client_on_upgrade.await {
273 Ok(client_upgraded) => {
274 let mut client_io = TokioIo::new(client_upgraded);
275 let mut upstream_io = TokioIo::new(upstream_upgraded);
276 if let Err(e) =
277 tokio::io::copy_bidirectional(&mut client_io, &mut upstream_io).await
278 {
279 debug!("websocket relay closed: {}", e);
280 }
281 }
282 Err(e) => warn!("client websocket upgrade failed: {}", e),
283 }
284 });
285
286 Ok(client_response)
287}
288
289pub fn bad_gateway_502() -> Response<Full<Bytes>> {
292 Response::builder()
293 .status(http::StatusCode::BAD_GATEWAY)
294 .header(http::header::CONTENT_TYPE, "application/json")
295 .body(Full::new(Bytes::from_static(
296 br#"{"error":"bad_gateway","message":"upstream websocket handshake failed"}"#,
297 )))
298 .unwrap()
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304
305 fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
306 let mut h = HeaderMap::new();
307 for (k, v) in pairs {
308 h.append(
309 http::HeaderName::from_bytes(k.as_bytes()).unwrap(),
310 http::HeaderValue::from_str(v).unwrap(),
311 );
312 }
313 h
314 }
315
316 #[test]
317 fn test_detect_plain_upgrade() {
318 assert!(is_websocket_upgrade(&headers(&[
319 ("connection", "Upgrade"),
320 ("upgrade", "websocket"),
321 ])));
322 }
323
324 #[test]
325 fn test_detect_multi_token_connection() {
326 assert!(is_websocket_upgrade(&headers(&[
328 ("connection", "keep-alive, Upgrade"),
329 ("upgrade", "WebSocket"),
330 ])));
331 }
332
333 #[test]
334 fn test_detect_requires_both_headers() {
335 assert!(!is_websocket_upgrade(&headers(&[(
336 "connection",
337 "Upgrade"
338 )])));
339 assert!(!is_websocket_upgrade(&headers(&[("upgrade", "websocket")])));
340 assert!(!is_websocket_upgrade(&headers(&[])));
341 }
342
343 #[test]
344 fn test_detect_rejects_non_websocket_upgrade() {
345 assert!(!is_websocket_upgrade(&headers(&[
347 ("connection", "Upgrade"),
348 ("upgrade", "h2c"),
349 ])));
350 }
351
352 #[test]
353 fn test_bad_gateway_502_shape() {
354 let resp = bad_gateway_502();
355 assert_eq!(resp.status(), 502);
356 assert_eq!(
357 resp.headers().get("content-type").unwrap(),
358 "application/json"
359 );
360 }
361
362 fn req_with(method: http::Method, protocol: Option<&'static str>) -> http::Request<()> {
363 let mut req = http::Request::builder().method(method).body(()).unwrap();
364 if let Some(p) = protocol {
365 req.extensions_mut()
366 .insert(hyper::ext::Protocol::from_static(p));
367 }
368 req
369 }
370
371 #[test]
372 fn test_detect_h2_extended_connect() {
373 let req = req_with(http::Method::CONNECT, Some("websocket"));
374 assert!(is_h2_websocket_connect(req.method(), req.extensions()));
375 }
376
377 #[test]
378 fn test_detect_h2_rejects_non_connect_and_non_websocket() {
379 let get = req_with(http::Method::GET, Some("websocket"));
381 assert!(!is_h2_websocket_connect(get.method(), get.extensions()));
382 let plain = req_with(http::Method::CONNECT, None);
384 assert!(!is_h2_websocket_connect(plain.method(), plain.extensions()));
385 let h2c = req_with(http::Method::CONNECT, Some("h2c"));
387 assert!(!is_h2_websocket_connect(h2c.method(), h2c.extensions()));
388 }
389
390 #[test]
391 fn test_synthesize_ws_key_is_16_bytes() {
392 use base64::{engine::general_purpose::STANDARD, Engine};
393 let key = synthesize_ws_key();
394 assert_eq!(STANDARD.decode(key).unwrap().len(), 16);
395 }
396}