mirror of
https://codeberg.org/icewind/haze.git
synced 2026-10-01 08:44:09 +02:00
improve websocket proxying
This commit is contained in:
parent
b977cd9dfa
commit
ad999702aa
6 changed files with 164 additions and 43 deletions
72
src/proxy.rs
72
src/proxy.rs
|
|
@ -5,26 +5,34 @@ use axum::http::header::HOST;
|
|||
use axum::http::HeaderValue;
|
||||
use axum::{
|
||||
body::Body,
|
||||
extract::{Request, State},
|
||||
extract::Request,
|
||||
response::{IntoResponse, Response},
|
||||
Router,
|
||||
};
|
||||
use bollard::Docker;
|
||||
use futures_util::StreamExt;
|
||||
use hyper::body::Incoming;
|
||||
use hyper::server::conn::http1;
|
||||
use hyper::service::service_fn;
|
||||
use hyper::StatusCode;
|
||||
use hyper_util::rt::TokioIo;
|
||||
use hyper_util::{client::legacy::connect::HttpConnector, rt::TokioExecutor};
|
||||
use miette::{miette, IntoDiagnostic};
|
||||
use std::collections::HashMap;
|
||||
use std::convert::Infallible;
|
||||
use std::fs::{create_dir_all, set_permissions};
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::PathBuf;
|
||||
use std::pin::pin;
|
||||
use std::str::FromStr;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
use tokio::io::{AsyncRead, AsyncWrite};
|
||||
use tokio::net::UnixListener;
|
||||
use tokio::signal::ctrl_c;
|
||||
use tokio::spawn;
|
||||
use tokio::time::sleep;
|
||||
use tokio_stream::wrappers::{TcpListenerStream, UnixListenerStream};
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
struct ActiveInstances {
|
||||
|
|
@ -163,20 +171,26 @@ async fn serve(instances: ActiveInstances, listen: String, base_address: String)
|
|||
ctrl_c().await.ok();
|
||||
};
|
||||
|
||||
let app = Router::new().fallback(handler).with_state(AppState {
|
||||
let state = AppState {
|
||||
instances: instances.clone(),
|
||||
base_address: base_address.clone(),
|
||||
proxy_client: Arc::new(proxy_client),
|
||||
});
|
||||
};
|
||||
|
||||
if !listen.starts_with('/') {
|
||||
let addr: SocketAddr = listen.parse().into_diagnostic()?;
|
||||
let listener = tokio::net::TcpListener::bind(addr).await.unwrap();
|
||||
println!("listening on {}", listener.local_addr().unwrap());
|
||||
axum::serve(listener, app)
|
||||
.with_graceful_shutdown(cancel)
|
||||
.await
|
||||
.unwrap();
|
||||
let mut connections = pin!(TcpListenerStream::new(listener).take_until(cancel));
|
||||
|
||||
while let Some(stream) = connections.next().await {
|
||||
match stream {
|
||||
Ok(stream) => handle_connection(state.clone(), stream),
|
||||
Err(error) => {
|
||||
error!(%error, "connection failed");
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let listen: PathBuf = listen.into();
|
||||
if let Some(parent) = listen.parent() {
|
||||
|
|
@ -187,18 +201,42 @@ async fn serve(instances: ActiveInstances, listen: String, base_address: String)
|
|||
}
|
||||
let _ = tokio::fs::remove_file(&listen).await;
|
||||
|
||||
let uds = UnixListener::bind(&listen).unwrap();
|
||||
let listener = UnixListener::bind(&listen).unwrap();
|
||||
println!("listening on {}", listen.display());
|
||||
set_permissions(&listen, PermissionsExt::from_mode(0o666)).into_diagnostic()?;
|
||||
|
||||
axum::serve(uds, app)
|
||||
.with_graceful_shutdown(cancel)
|
||||
.await
|
||||
.unwrap();
|
||||
let mut connections = pin!(UnixListenerStream::new(listener).take_until(cancel));
|
||||
|
||||
while let Some(stream) = connections.next().await {
|
||||
match stream {
|
||||
Ok(stream) => handle_connection(state.clone(), stream),
|
||||
Err(error) => {
|
||||
error!(%error, "connection failed");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn handle_connection<I: AsyncRead + AsyncWrite + Unpin + Send + 'static>(
|
||||
state: AppState,
|
||||
stream: I,
|
||||
) {
|
||||
let io = TokioIo::new(stream);
|
||||
// Spawn a tokio task to serve multiple connections concurrently
|
||||
tokio::task::spawn(async move {
|
||||
if let Err(err) = http1::Builder::new()
|
||||
.serve_connection(io, service_fn(move |req| handler(state.clone(), req)))
|
||||
.with_upgrades()
|
||||
.await
|
||||
{
|
||||
eprintln!("Error serving connection: {:?}", err);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
async fn get_remote(
|
||||
host: Option<&HeaderValue>,
|
||||
instances: &ActiveInstances,
|
||||
|
|
@ -232,9 +270,9 @@ async fn get_remote(
|
|||
}
|
||||
}
|
||||
|
||||
type Client = hyper_util::client::legacy::Client<HttpConnector, Body>;
|
||||
type Client = hyper_util::client::legacy::Client<HttpConnector, Incoming>;
|
||||
|
||||
async fn handler(State(state): State<AppState>, mut req: Request) -> Result<Response, StatusCode> {
|
||||
async fn handler(state: AppState, mut req: Request<Incoming>) -> Result<Response, Infallible> {
|
||||
let host = req.headers().get(HOST).cloned();
|
||||
let remote = match get_remote(host.as_ref(), &state.instances, &state.base_address).await {
|
||||
Ok(remote) => remote,
|
||||
|
|
@ -259,13 +297,13 @@ async fn handler(State(state): State<AppState>, mut req: Request) -> Result<Resp
|
|||
IpAddr::V4(Ipv4Addr::UNSPECIFIED),
|
||||
&uri,
|
||||
req,
|
||||
&state.proxy_client,
|
||||
state.proxy_client.as_ref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => Ok(response.map(Body::new)),
|
||||
Err(error) => {
|
||||
error!(%error, "error while proxying request");
|
||||
error!(?error, "error while proxying request");
|
||||
Ok(StatusCode::BAD_REQUEST.into_response())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue