use crate::axum::body::Body;
use crate::axum::extract::Request;
#[must_use]
pub(crate) fn is_vite_request(req: &Request<Body>) -> bool {
let path = req.uri().path();
if path.starts_with("/@")
|| path.starts_with("/src/")
|| path.starts_with("/node_modules/.vite/")
{
return true;
}
if is_vite_ws_upgrade(req) {
return true;
}
false
}
fn is_vite_ws_upgrade(req: &Request<Body>) -> bool {
let headers = req.headers();
let connection_upgrade = headers
.get("connection")
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.to_ascii_lowercase().contains("upgrade"));
let upgrade_websocket = headers
.get("upgrade")
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.eq_ignore_ascii_case("websocket"));
let vite_protocol = headers
.get("sec-websocket-protocol")
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.contains("vite-hmr") || v.contains("vite-ping"));
connection_upgrade && upgrade_websocket && vite_protocol
}
#[cfg(test)]
mod tests {
use super::*;
use crate::axum::http::{HeaderMap, HeaderName, HeaderValue, Method, Request, Uri};
use std::str::FromStr as _;
fn request(path: &str) -> Request<Body> {
Request::builder()
.method(Method::GET)
.uri(Uri::from_str(path).expect("test path should parse as URI"))
.body(Body::empty())
.expect("test request should build")
}
fn ws_request(protocol: &str) -> Request<Body> {
let mut headers = HeaderMap::new();
headers.insert("connection", HeaderValue::from_static("upgrade"));
headers.insert("upgrade", HeaderValue::from_static("websocket"));
headers.insert(
HeaderName::from_static("sec-websocket-protocol"),
HeaderValue::from_str(protocol).expect("valid header value"),
);
Request::builder()
.method(Method::GET)
.uri(Uri::from_static("/"))
.body(Body::empty())
.map(|mut req| {
*req.headers_mut() = headers;
req
})
.expect("test ws request should build")
}
#[test]
fn vite_internal_paths_are_forwarded() {
assert!(is_vite_request(&request("/@vite/client")));
assert!(is_vite_request(&request("/@react-refresh")));
assert!(is_vite_request(&request("/@fs/src/app.tsx")));
assert!(is_vite_request(&request("/@id/react")));
}
#[test]
fn source_modules_are_forwarded() {
assert!(is_vite_request(&request("/src/app.tsx")));
assert!(is_vite_request(&request("/src/main.ts")));
assert!(is_vite_request(&request("/src/styles/app.css")));
}
#[test]
fn optimized_deps_are_forwarded() {
assert!(is_vite_request(&request(
"/node_modules/.vite/deps/react.js"
)));
}
#[test]
fn application_paths_are_not_forwarded() {
assert!(!is_vite_request(&request("/")));
assert!(!is_vite_request(&request("/api/users")));
assert!(!is_vite_request(&request("/dashboard")));
assert!(!is_vite_request(&request("/favicon.ico")));
}
#[test]
fn vite_hmr_websocket_is_forwarded() {
assert!(is_vite_request(&ws_request("vite-hmr")));
assert!(is_vite_request(&ws_request("vite-ping")));
}
#[test]
fn non_vite_websocket_is_not_forwarded() {
let req = ws_request("custom-app-protocol");
assert!(!is_vite_request(&req));
}
#[test]
fn plain_get_to_root_is_not_forwarded() {
assert!(!is_vite_request(&request("/")));
}
#[test]
fn websocket_without_vite_protocol_is_not_forwarded() {
let mut headers = HeaderMap::new();
headers.insert("connection", HeaderValue::from_static("Upgrade"));
headers.insert("upgrade", HeaderValue::from_static("websocket"));
let req = Request::builder()
.method(Method::GET)
.uri(Uri::from_static("/"))
.body(Body::empty())
.map(|mut req| {
*req.headers_mut() = headers;
req
})
.expect("request should build");
assert!(!is_vite_request(&req));
}
}