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::{Response, StatusCode};
use crate::dev_proxy::endpoint::{IpcEndpoint, IpcStream};
use crate::dev_proxy::vite::is_vite_request;
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>>,
}
impl DevProxyLayer {
#[must_use]
pub fn new(endpoint: Option<IpcEndpoint>) -> Self {
Self {
endpoint: endpoint.map(Arc::new),
}
}
}
impl<S> tower::Layer<S> for DevProxyLayer {
type Service = DevProxyService<S>;
fn layer(&self, inner: S) -> Self::Service {
DevProxyService {
inner,
endpoint: self.endpoint.clone(),
}
}
}
#[derive(Clone)]
pub struct DevProxyService<Inner> {
inner: Inner,
endpoint: Option<Arc<IpcEndpoint>>,
}
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 Some(endpoint) = self.endpoint.clone() else {
let fut = self.inner.call(req);
return Box::pin(fut);
};
if !is_vite_request(&req) {
let fut = self.inner.call(req);
return Box::pin(fut);
}
let inner = self.inner.clone();
Box::pin(forward_or_delegate(endpoint, req, inner))
}
}
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()))
})
}
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))
}
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 layer_with_endpoint_stores_arc() {
let layer = DevProxyLayer::new(Some(IpcEndpoint::new(std::path::PathBuf::from(
"/tmp/arcature-test.sock",
))));
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");
}
}