use std::collections::{BTreeMap, HashMap};
use std::sync::{Arc, Mutex};
use anyhow::{Context, Result, anyhow, bail};
use axum::body::Body;
use axum::extract::ws::{Message as ClientMessage, WebSocket, WebSocketUpgrade};
use axum::extract::{FromRequestParts, Request, State};
use axum::http::header::{CONNECTION, CONTENT_LENGTH, COOKIE, HOST, LOCATION, SEC_WEBSOCKET_PROTOCOL, SET_COOKIE, TRANSFER_ENCODING, UPGRADE};
use axum::http::{HeaderMap, HeaderValue, StatusCode};
use axum::response::Response;
use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::Message as StandMessage;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use crate::model::{App, Server};
use crate::store;
const MAX_REQUEST_BODY: usize = 256 * 1024 * 1024;
type Jars = Arc<Mutex<HashMap<(String, String), HashMap<String, String>>>>;
type Clients = Arc<Mutex<HashMap<String, reqwest::Client>>>;
#[derive(Clone)]
struct Ctx {
app: String,
port: u16,
jars: Jars,
clients: Clients,
}
pub fn listening_ports(apps: &[App]) -> Result<BTreeMap<u16, String>> {
let mut ports = BTreeMap::new();
for app in apps {
let Some(port) = app.gateway_port else { continue };
if let Some(other) = ports.insert(port, app.name.clone()) {
bail!(
"apps '{other}' and '{}' share gateway port {port} - give one of them another port with `turnout app edit NAME --port PORT`",
app.name
);
}
}
if ports.is_empty() {
bail!("no apps with a gateway port - set one with `turnout app edit NAME --port PORT`");
}
Ok(ports)
}
pub fn run() -> Result<()> {
let apps: Vec<(String, u16)> = listening_ports(&store::load_apps()?)?.into_iter().map(|(port, app)| (app, port)).collect();
let runtime = tokio::runtime::Runtime::new()?;
runtime.block_on(async move {
let jars: Jars = Arc::default();
let clients: Clients = Arc::default();
for (app, port) in apps {
let listener = tokio::net::TcpListener::bind(("127.0.0.1", port))
.await
.with_context(|| format!("cannot listen on 127.0.0.1:{port} - the port is taken (a running gateway? stop it with `turnout gateway stop`)"))?;
println!("{app}: listening on http://localhost:{port}");
let ctx = Ctx {
app,
port,
jars: jars.clone(),
clients: clients.clone(),
};
let router = axum::Router::new().fallback(proxy).with_state(ctx);
tokio::spawn(async move { axum::serve(listener, router).await });
}
tokio::signal::ctrl_c().await?;
println!("Gateway stopped.");
Ok(())
})
}
async fn proxy(State(ctx): State<Ctx>, req: Request) -> Response {
let result = if is_websocket_upgrade(req.headers()) {
websocket_proxy(ctx, req).await
} else {
forward(ctx, req).await
};
result.unwrap_or_else(|err| {
Response::builder()
.status(StatusCode::BAD_GATEWAY)
.header("content-type", "text/plain; charset=utf-8")
.body(Body::from(format!("turnout gateway: {err:#}\n")))
.expect("static response")
})
}
fn resolve_binding(ctx: &Ctx) -> Result<Server> {
let state = store::load_state()?;
let server_name = state
.bindings
.get(&ctx.app)
.cloned()
.ok_or_else(|| anyhow!("app '{0}' is not bound to a server - run `turnout use {0} SERVER`", ctx.app))?;
store::load_servers()?
.into_iter()
.find(|s| s.name == server_name)
.ok_or_else(|| anyhow!("server '{server_name}' is gone from the catalog"))
}
async fn forward(ctx: Ctx, req: Request) -> Result<Response> {
let server = resolve_binding(&ctx)?;
let server_name = server.name.clone();
let client = client_for(&ctx.clients, &server)?;
let (parts, body) = req.into_parts();
let path_query = parts.uri.path_and_query().map(|p| p.as_str()).unwrap_or("/");
let target = format!("{}{}", server.url.trim_end_matches('/'), path_query);
let mut headers = parts.headers;
strip_hop_headers(&mut headers);
headers.remove(HOST);
headers.remove(COOKIE);
let jar_key = (ctx.app.clone(), server_name.clone());
if let Some(cookie) = stored_cookie_header(&ctx.jars, &jar_key)? {
headers.insert(COOKIE, cookie);
}
let body = axum::body::to_bytes(body, MAX_REQUEST_BODY).await.context("request body too large")?;
let upstream = client
.request(parts.method, &target)
.headers(headers)
.body(body)
.send()
.await
.map_err(|err| anyhow!("stand '{server_name}' is unreachable: {err}"))?;
let status = upstream.status();
let mut resp_headers = upstream.headers().clone();
capture_cookies(&ctx.jars, &jar_key, &resp_headers);
resp_headers.remove(SET_COOKIE);
strip_hop_headers(&mut resp_headers);
rewrite_location(&mut resp_headers, &server.url, ctx.port)?;
let mut response = Response::builder().status(status);
if let Some(headers_mut) = response.headers_mut() {
*headers_mut = resp_headers;
}
Ok(response.body(Body::from_stream(upstream.bytes_stream()))?)
}
fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
headers
.get(UPGRADE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("websocket"))
}
async fn websocket_proxy(ctx: Ctx, req: Request) -> Result<Response> {
let (mut parts, _body) = req.into_parts();
let mut upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &())
.await
.map_err(|err| anyhow!("invalid websocket upgrade: {err}"))?;
let server = resolve_binding(&ctx)?;
let server_name = server.name.clone();
let path_query = parts.uri.path_and_query().map(|p| p.as_str()).unwrap_or("/");
let ws_base = if let Some(rest) = server.url.strip_prefix("https://") {
format!("wss://{rest}")
} else if let Some(rest) = server.url.strip_prefix("http://") {
format!("ws://{rest}")
} else {
bail!("server URL '{}' is neither http nor https", server.url);
};
let target = format!("{}{path_query}", ws_base.trim_end_matches('/'));
let mut request = target.into_client_request().context("cannot build the upstream websocket request")?;
let jar_key = (ctx.app.clone(), server_name.clone());
if let Some(cookie) = stored_cookie_header(&ctx.jars, &jar_key)? {
request.headers_mut().insert(COOKIE, cookie);
}
if let Some(protocol) = parts.headers.get(SEC_WEBSOCKET_PROTOCOL) {
request.headers_mut().insert(SEC_WEBSOCKET_PROTOCOL, protocol.clone());
if let Ok(protocols) = protocol.to_str() {
upgrade = upgrade.protocols(protocols.split(',').map(|p| p.trim().to_string()).collect::<Vec<_>>());
}
}
let connector = if server.accept_invalid_certs {
Some(tokio_tungstenite::Connector::Rustls(Arc::new(accept_any_tls_config())))
} else {
None
};
Ok(upgrade.on_upgrade(move |client| async move {
match tokio_tungstenite::connect_async_tls_with_config(request, None, false, connector).await {
Ok((upstream, _response)) => pump(client, upstream).await,
Err(err) => eprintln!("turnout gateway: websocket to '{server_name}' failed: {err}"),
}
}))
}
fn accept_any_tls_config() -> rustls::ClientConfig {
rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(AcceptAnyServerCert))
.with_no_client_auth()
}
#[derive(Debug)]
struct AcceptAnyServerCert;
impl rustls::client::danger::ServerCertVerifier for AcceptAnyServerCert {
fn verify_server_cert(
&self,
_end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::aws_lc_rs::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
async fn pump<S>(client: WebSocket, upstream: tokio_tungstenite::WebSocketStream<S>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
let (mut client_tx, mut client_rx) = client.split();
let (mut upstream_tx, mut upstream_rx) = upstream.split();
let to_stand = async {
while let Some(Ok(message)) = client_rx.next().await {
if upstream_tx.send(client_to_stand(message)).await.is_err() {
break;
}
}
let _ = upstream_tx.close().await;
};
let to_client = async {
while let Some(Ok(message)) = upstream_rx.next().await {
let Some(message) = stand_to_client(message) else { continue };
if client_tx.send(message).await.is_err() {
break;
}
}
let _ = client_tx.close().await;
};
tokio::join!(to_stand, to_client);
}
fn client_to_stand(message: ClientMessage) -> StandMessage {
match message {
ClientMessage::Text(text) => StandMessage::text(text.to_string()),
ClientMessage::Binary(data) => StandMessage::binary(data.to_vec()),
ClientMessage::Ping(data) => StandMessage::Ping(data.to_vec().into()),
ClientMessage::Pong(data) => StandMessage::Pong(data.to_vec().into()),
ClientMessage::Close(frame) => StandMessage::Close(frame.map(|f| tokio_tungstenite::tungstenite::protocol::CloseFrame {
code: f.code.into(),
reason: f.reason.as_str().into(),
})),
}
}
fn stand_to_client(message: StandMessage) -> Option<ClientMessage> {
match message {
StandMessage::Text(text) => Some(ClientMessage::Text(text.to_string().into())),
StandMessage::Binary(data) => Some(ClientMessage::Binary(data.to_vec().into())),
StandMessage::Ping(data) => Some(ClientMessage::Ping(data.to_vec().into())),
StandMessage::Pong(data) => Some(ClientMessage::Pong(data.to_vec().into())),
StandMessage::Close(frame) => Some(ClientMessage::Close(frame.map(|f| axum::extract::ws::CloseFrame {
code: f.code.into(),
reason: f.reason.as_str().into(),
}))),
StandMessage::Frame(_) => None,
}
}
fn client_for(clients: &Clients, server: &Server) -> Result<reqwest::Client> {
let mut cache = clients.lock().expect("client cache poisoned");
if let Some(client) = cache.get(&server.name) {
return Ok(client.clone());
}
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.danger_accept_invalid_certs(server.accept_invalid_certs)
.build()
.context("cannot build the HTTP client")?;
cache.insert(server.name.clone(), client.clone());
Ok(client)
}
fn stored_cookie_header(jars: &Jars, key: &(String, String)) -> Result<Option<HeaderValue>> {
jar_cookie_header(jars, key)
.map(|cookie| HeaderValue::from_str(&cookie).context("stored cookie is not a valid header value"))
.transpose()
}
fn jar_cookie_header(jars: &Jars, key: &(String, String)) -> Option<String> {
let jars = jars.lock().expect("cookie jar poisoned");
let jar = jars.get(key)?;
if jar.is_empty() {
return None;
}
Some(jar.iter().map(|(name, value)| format!("{name}={value}")).collect::<Vec<_>>().join("; "))
}
fn capture_cookies(jars: &Jars, key: &(String, String), headers: &HeaderMap) {
let mut jars = jars.lock().expect("cookie jar poisoned");
for value in headers.get_all(SET_COOKIE) {
let Ok(text) = value.to_str() else { continue };
let pair = text.split(';').next().unwrap_or_default();
if let Some((name, value)) = pair.split_once('=') {
jars.entry(key.clone()).or_default().insert(name.trim().to_string(), value.trim().to_string());
}
}
}
fn rewrite_location(headers: &mut HeaderMap, server_url: &str, port: u16) -> Result<()> {
let Some(location) = headers.get(LOCATION).and_then(|v| v.to_str().ok()).map(str::to_string) else {
return Ok(());
};
let base = server_url.trim_end_matches('/');
if let Some(rest) = location.strip_prefix(base) {
let rewritten = format!("http://localhost:{port}{rest}");
headers.insert(
LOCATION,
HeaderValue::from_str(&rewritten).context("rewritten location is not a valid header value")?,
);
}
Ok(())
}
fn strip_hop_headers(headers: &mut HeaderMap) {
for name in [CONNECTION, TRANSFER_ENCODING, UPGRADE, CONTENT_LENGTH] {
headers.remove(&name);
}
headers.remove("keep-alive");
headers.remove("proxy-authenticate");
headers.remove("proxy-authorization");
headers.remove("te");
headers.remove("trailer");
}
#[cfg(test)]
mod tests {
use super::*;
fn app(name: &str, port: Option<u16>) -> App {
App {
name: name.to_string(),
path: String::new(),
commands: BTreeMap::new(),
dist_dir: None,
gateway_port: port,
gateway_env: None,
env_file: None,
servers: Vec::new(),
}
}
#[test]
fn listening_ports_maps_each_port_to_its_app() {
let ports = listening_ports(&[app("web", Some(7001)), app("api", None), app("admin", Some(7002))]).unwrap();
assert_eq!(ports.get(&7001).map(String::as_str), Some("web"));
assert_eq!(ports.get(&7002).map(String::as_str), Some("admin"));
assert_eq!(ports.len(), 2, "an app without a port got a listener");
}
#[test]
fn listening_ports_refuses_a_port_two_apps_share() {
let error = listening_ports(&[app("web", Some(7001)), app("admin", Some(7001))]).unwrap_err().to_string();
assert!(error.contains("'web' and 'admin' share gateway port 7001"), "{error}");
assert!(error.contains("turnout app edit NAME --port PORT"), "{error}");
}
#[test]
fn listening_ports_needs_at_least_one_port() {
let error = listening_ports(&[app("web", None)]).unwrap_err().to_string();
assert!(error.contains("no apps with a gateway port"), "{error}");
}
}