use std::convert::Infallible;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use crate::axum::body::Body;
use crate::axum::http::{Method, Response, StatusCode};
use crate::dev_proxy::endpoint::{IpcEndpoint, IpcStream};
use crate::dev_proxy::vite::ViteRoutes;
type Request = crate::axum::extract::Request<Body>;
type BoxFuture = Pin<Box<dyn Future<Output = Result<Response<Body>, Infallible>> + Send>>;
#[derive(Clone)]
pub struct DevProxyLayer {
endpoint: Option<Arc<IpcEndpoint>>,
routes: ViteRoutes,
}
impl DevProxyLayer {
#[must_use]
pub fn new(endpoint: Option<IpcEndpoint>) -> Self {
Self::with_routes(endpoint, crate::dev_proxy::config::prefixes_from_env())
}
#[must_use]
pub(crate) fn with_routes(endpoint: Option<IpcEndpoint>, routes: ViteRoutes) -> Self {
Self {
endpoint: endpoint.map(Arc::new),
routes,
}
}
}
impl<S> tower::Layer<S> for DevProxyLayer {
type Service = DevProxyService<S>;
fn layer(&self, inner: S) -> Self::Service {
DevProxyService {
inner,
endpoint: self.endpoint.clone(),
routes: self.routes.clone(),
}
}
}
#[derive(Clone)]
pub struct DevProxyService<Inner> {
inner: Inner,
endpoint: Option<Arc<IpcEndpoint>>,
routes: ViteRoutes,
}
impl<Inner> tower::Service<Request> for DevProxyService<Inner>
where
Inner: tower::Service<Request, Response = Response<Body>, Error = Infallible>
+ Clone
+ Send
+ 'static,
Inner::Future: Send + 'static,
{
type Response = Response<Body>;
type Error = Infallible;
type Future = BoxFuture;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request) -> Self::Future {
let clone = self.inner.clone();
let inner = std::mem::replace(&mut self.inner, clone);
let Some(endpoint) = self.endpoint.clone() else {
let mut inner = inner;
return Box::pin(async move { inner.call(req).await });
};
if self.routes.matches_request(&req) {
return Box::pin(forward_or_delegate(endpoint, req, inner));
}
Box::pin(delegate_or_retry(endpoint, req, inner))
}
}
async fn delegate_or_retry<Inner>(
endpoint: Arc<IpcEndpoint>,
req: Request,
mut inner: Inner,
) -> Result<Response<Body>, Infallible>
where
Inner: tower::Service<Request, Response = Response<Body>, Error = Infallible>
+ Clone
+ Send
+ 'static,
Inner::Future: Send + 'static,
{
let replayable = matches!(*req.method(), Method::GET | Method::HEAD)
&& !req
.headers()
.contains_key(crate::axum::http::header::UPGRADE);
if !replayable {
return inner.call(req).await;
}
let replay = crate::axum::http::Request::builder()
.method(req.method().clone())
.uri(req.uri().clone())
.version(req.version())
.body(Body::empty())
.map(|mut replay| {
*replay.headers_mut() = req.headers().clone();
replay
});
let from_app = inner.call(req).await?;
if from_app.status() != StatusCode::NOT_FOUND {
return Ok(from_app);
}
let Ok(replay) = replay else {
return Ok(from_app);
};
let Ok(stream) = endpoint.connect().await else {
return Ok(from_app);
};
match forward(stream, replay).await {
Ok(from_vite) if from_vite.status() != StatusCode::NOT_FOUND => Ok(from_vite),
_ => Ok(from_app),
}
}
async fn forward_or_delegate<Inner>(
endpoint: Arc<IpcEndpoint>,
req: Request,
mut inner: Inner,
) -> Result<Response<Body>, Infallible>
where
Inner: tower::Service<Request, Response = Response<Body>, Error = Infallible>
+ Clone
+ Send
+ 'static,
Inner::Future: Send + 'static,
{
match endpoint.connect().await {
Ok(stream) => match forward(stream, req).await {
Ok(response) => Ok(response),
Err(err) => Ok(bad_gateway(err)),
},
Err(_) => inner.call(req).await,
}
}
fn bad_gateway(err: ForwardError) -> Response<Body> {
eprintln!("warning: vite ipc forward failed: {err}");
Response::builder()
.status(StatusCode::BAD_GATEWAY)
.header(
crate::axum::http::HeaderName::from_static("x-content-type-options"),
crate::axum::http::HeaderValue::from_static("nosniff"),
)
.header(
crate::axum::http::HeaderName::from_static("content-type"),
crate::axum::http::HeaderValue::from_static("text/plain; charset=utf-8"),
)
.body(Body::from(
"vite dev server closed the connection (is Vite running?)",
))
.unwrap_or_else(|_| {
Response::builder()
.status(StatusCode::BAD_GATEWAY)
.body(Body::empty())
.unwrap_or_else(|_| Response::new(Body::empty()))
})
}
pub(crate) async fn forward(
stream: IpcStream,
req: Request,
) -> Result<Response<Body>, ForwardError> {
let (mut parts, body) = req.into_parts();
let browser_upgrade = parts.extensions.remove::<hyper::upgrade::OnUpgrade>();
let upstream_req = crate::axum::http::Request::from_parts(parts, body);
let io = hyper_util::rt::tokio::WithHyperIo::new(stream);
let (mut sender, conn) = hyper::client::conn::http1::handshake(io).await?;
tokio::spawn(async move {
if let Err(err) = conn.with_upgrades().await {
eprintln!("warning: vite ipc connection ended: {err}");
}
});
let response = sender.send_request(upstream_req).await?;
let status = response.status();
if status == StatusCode::SWITCHING_PROTOCOLS {
let (mut resp_parts, _body) = response.into_parts();
let vite_upgrade = resp_parts.extensions.remove::<hyper::upgrade::OnUpgrade>();
let browser_response = Response::from_parts(resp_parts, Body::empty());
if let (Some(browser), Some(vite)) = (browser_upgrade, vite_upgrade) {
tokio::spawn(tunnel(browser, vite));
}
return Ok(browser_response);
}
Ok(map_response(response))
}
async fn tunnel(browser: hyper::upgrade::OnUpgrade, vite: hyper::upgrade::OnUpgrade) {
let browser = match browser.await {
Ok(io) => io,
Err(err) => {
eprintln!("warning: browser-side hmr upgrade failed: {err}");
return;
}
};
let vite = match vite.await {
Ok(io) => io,
Err(err) => {
eprintln!("warning: vite-side hmr upgrade failed: {err}");
return;
}
};
let mut browser = hyper_util::rt::TokioIo::new(browser);
let mut vite = hyper_util::rt::TokioIo::new(vite);
let _ = tokio::io::copy_bidirectional(&mut browser, &mut vite).await;
}
fn map_response(resp: hyper::Response<hyper::body::Incoming>) -> Response<Body> {
let (parts, body) = resp.into_parts();
Response::from_parts(parts, Body::new(body))
}
pub(crate) enum ForwardError {
Handshake(hyper::Error),
Send(hyper::Error),
}
impl From<hyper::Error> for ForwardError {
fn from(err: hyper::Error) -> Self {
if err.is_canceled() {
ForwardError::Handshake(err)
} else {
ForwardError::Send(err)
}
}
}
impl std::fmt::Display for ForwardError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ForwardError::Handshake(err) => {
write!(f, "vite ipc handshake failed: {err}")
}
ForwardError::Send(err) => {
write!(f, "vite ipc send failed: {err}")
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::axum::http::HeaderValue;
#[test]
fn layer_with_no_endpoint_is_passthrough_marker() {
let layer = DevProxyLayer::new(None);
assert!(layer.endpoint.is_none());
}
#[test]
fn a_layer_built_from_the_environment_still_knows_the_conventional_roots() {
let layer = DevProxyLayer::new(None);
assert!(layer.routes.matches_path("/resources/js/app.tsx"));
}
#[test]
fn layer_with_endpoint_stores_arc() {
let layer = DevProxyLayer::with_routes(
Some(IpcEndpoint::new(std::path::PathBuf::from(
"/tmp/arcature-test.sock",
))),
crate::dev_proxy::vite::ViteRoutes::defaults(),
);
let endpoint = layer
.endpoint
.as_ref()
.expect("endpoint should be stored when provided");
assert_eq!(
endpoint.path(),
std::path::Path::new("/tmp/arcature-test.sock")
);
}
#[test]
fn header_value_roundtrips() {
let v = HeaderValue::from_static("websocket");
assert_eq!(v.as_bytes(), b"websocket");
}
}