use crate::service::{ServiceTrait, ServiceType}; use crate::Result; use crate::{Cloud, HazeConfig}; use axum::http::header::HOST; use axum::http::HeaderValue; use axum::{ body::Body, extract::Request, response::{IntoResponse, Response}, }; 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_rustls::TlsAcceptor; use tokio_rustls::rustls::ServerConfig; use tokio_rustls::rustls::pki_types::pem::PemObject; use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer}; use tokio_stream::wrappers::{TcpListenerStream, UnixListenerStream}; use tracing::{debug, error, info}; struct ActiveInstances { known: Mutex>, last: Mutex>, docker: Docker, config: HazeConfig, } impl ActiveInstances { pub fn new(docker: Docker, config: HazeConfig) -> Self { ActiveInstances { known: Mutex::default(), last: Mutex::default(), docker, config, } } pub async fn get(&self, name: &str) -> Option { if let Some(ip) = self.known.lock().unwrap().get(name).cloned() { return Some(ip); } let addr = self.resolve_addr(name).await?; println!("{name} => {addr}"); self.known.lock().unwrap().insert(name.into(), addr); Some(addr) } async fn resolve_addr(&self, name: &str) -> Option { // instance if let Ok(cloud) = Cloud::get_by_filter(&self.docker, Some(name.into()), &self.config).await { return Some(SocketAddr::new(cloud.ip?, 80)); } // service if ServiceType::from_str(name).is_ok() { let cloud = self.last()?; let service = cloud.services().find(|service| service.name() == name)?; let ip = service .get_ips(&self.docker, &cloud.id) .await .ok()? .next()?; return Some(SocketAddr::new(ip, service.proxy_port())); } // instance-service for (pos, _) in name.match_indices('-') { let instance_name = &name[..pos]; let service_name = &name[pos + 1..]; if ServiceType::from_str(service_name).is_err() { continue; } let cloud = match Cloud::get_by_filter(&self.docker, Some(instance_name.into()), &self.config) .await { Ok(cloud) => cloud, Err(_) => continue, }; let service = match cloud.services().find(|s| s.name() == service_name) { Some(service) => service, None => continue, }; let ip = service .get_ips(&self.docker, &cloud.id) .await .ok()? .next()?; return Some(SocketAddr::new(ip, service.proxy_port())); } None } pub fn last_addr(&self) -> Option { self.last .lock() .unwrap() .as_ref() .and_then(|cloud| Some(SocketAddr::new(cloud.ip?, 80))) } pub fn last(&self) -> Option { self.last.lock().unwrap().clone() } async fn update_last(&self) { let last = Cloud::get_by_filter(&self.docker, None, &self.config) .await .ok(); let mut old = self.last.lock().unwrap(); if old.as_ref() != last.as_ref() { info!(instance = ?last, "Found new instance"); // remove cached base-service mappings self.known .lock() .unwrap() .retain(|key, _| ServiceType::from_str(key).is_err()); *old = last; } } } pub async fn proxy(docker: Docker, config: HazeConfig) -> Result<()> { if config.proxy.listen.is_empty() { return Err(miette!("Proxy not configured")); } let listen = config.proxy.listen.clone(); let acceptor = match (&config.proxy.cert, &config.proxy.key) { (None, None) => None, (Some(_), None) => return Err(miette!("`cert` is set without `key`")), (None, Some(_)) => return Err(miette!("`key` is set without `cert`")), (Some(cert), Some(key)) => Some(tls_acceptor(cert, key)?), }; let base_address = config.proxy.address.clone(); let instances = ActiveInstances::new(docker, config); serve(instances, listen, base_address, acceptor).await } /// Build a TLS acceptor from a PEM encoded certificate chain and private key on disk fn tls_acceptor(cert: &str, key: &str) -> Result { let certs = CertificateDer::pem_file_iter(cert) .map_err(|e| miette!("failed to load certificate from {cert}: {e}"))? .collect::, _>>() .map_err(|e| miette!("failed to load certificate from {cert}: {e}"))?; let key = PrivateKeyDer::from_pem_file(key) .map_err(|e| miette!("failed to load private key from {key}: {e}"))?; let mut server_config = ServerConfig::builder() .with_no_client_auth() .with_single_cert(certs, key) .into_diagnostic()?; server_config.alpn_protocols = vec![b"http/1.1".to_vec()]; Ok(TlsAcceptor::from(Arc::new(server_config))) } #[derive(Clone)] struct AppState { instances: Arc, base_address: Arc, proxy_client: Arc, } async fn serve( instances: ActiveInstances, listen: String, base_address: String, acceptor: Option, ) -> Result<()> { let instances = Arc::new(instances); let base_address = Arc::new(base_address); let last_instances = instances.clone(); let proxy_client: Client = hyper_util::client::legacy::Client::<(), ()>::builder(TokioExecutor::new()) .build(HttpConnector::new()); spawn(async move { loop { sleep(Duration::from_secs(1)).await; last_instances.update_last().await; } }); let cancel = async { ctrl_c().await.ok(); }; 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(); let scheme = if acceptor.is_some() { "https" } else { "http" }; println!("Listening on {scheme}://{}", listener.local_addr().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, acceptor.clone()).await, Err(error) => { error!(%error, "connection failed"); } } } } else { let listen: PathBuf = listen.into(); if let Some(parent) = listen.parent() && !parent.exists() { create_dir_all(parent).into_diagnostic()?; set_permissions(parent, PermissionsExt::from_mode(0o755)).into_diagnostic()?; } let _ = tokio::fs::remove_file(&listen).await; let listener = UnixListener::bind(&listen).unwrap(); println!("listening on {}", listen.display()); set_permissions(&listen, PermissionsExt::from_mode(0o666)).into_diagnostic()?; 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, None).await, Err(error) => { error!(%error, "connection failed"); } } } } Ok(()) } async fn handle_connection( state: AppState, stream: I, acceptor: Option, ) { // Spawn a tokio task to serve multiple connections concurrently match acceptor { Some(acceptor) => match acceptor.accept(stream).await { Ok(stream) => serve_connection(state, stream).await, Err(error) => error!(%error, "tls handshake failed"), }, None => serve_connection(state, stream).await, } } async fn serve_connection( 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, base_address: &str, ) -> Result { let host = match host.and_then(|host| host.to_str().ok()) { Some(host) => host, None => return Err("No or invalid hostname provided".into()), }; let ip = if host == base_address { instances .last_addr() .ok_or_else(|| String::from("No running instance known")) } else { let requested_instance = host.split('.').next().unwrap(); if requested_instance == "host-push" { Ok((IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 7867).into()) } else { instances .get(requested_instance) .await .ok_or_else(|| format!("Error {} has no known ip", requested_instance)) } }; match ip { Ok(ip) => Ok(ip), Err(e) => { eprintln!("{}", e); Err(e) } } } type Client = hyper_util::client::legacy::Client; async fn handler(state: AppState, mut req: Request) -> Result { let host = req.headers().get(HOST).cloned(); let remote = match get_remote(host.as_ref(), &state.instances, &state.base_address).await { Ok(remote) => remote, Err(e) => { return Ok(hyper::Response::builder() .status(StatusCode::BAD_REQUEST) .body(e.into()) .unwrap()) } }; let uri = format!("http://{remote}"); debug!(target = uri, "proxying request"); // fix weird duplicate host header req.headers_mut().remove(HOST); if let Some(host) = host { req.headers_mut().insert(HOST, host.clone()); } match hyper_reverse_proxy::call( IpAddr::V4(Ipv4Addr::UNSPECIFIED), &uri, req, state.proxy_client.as_ref(), ) .await { Ok(response) => Ok(response.map(Body::new)), Err(error) => { error!(?error, "error while proxying request"); Ok(StatusCode::BAD_REQUEST.into_response()) } } }