use std::collections::HashSet;
use std::{collections::HashMap, net::SocketAddr, path::Path, sync::Arc};
use crate::engine::workload::{ResolvedWorkload, WorkloadComponent};
use crate::wit::WitInterface;
use crate::wit::WitWorld;
use crate::{engine::ctx::Ctx, plugin::HostPlugin};
use anyhow::{Context, bail, ensure};
use hyper::server::conn::http1;
use tokio::net::TcpListener;
use tracing::{debug, error, info, warn};
use wasmtime::component::InstancePre;
use wasmtime::{AsContextMut, StoreContextMut};
use wasmtime_wasi_http::{
WasiHttpView,
bindings::{ProxyPre, http::types::Scheme},
body::HyperOutgoingBody,
io::TokioIo,
};
use rustls::{ServerConfig, pki_types::CertificateDer};
use rustls_pemfile::{certs, private_key};
use tokio::sync::{RwLock, mpsc};
use tokio_rustls::TlsAcceptor;
const HTTP_SERVER_ID: &str = "http-server";
#[derive(Clone, Debug)]
struct HttpWorkloadConfig {
host_header: String,
}
pub type WorkloadHandles =
Arc<RwLock<HashMap<String, (ResolvedWorkload, InstancePre<Ctx>, String)>>>;
pub struct HttpServer {
addr: SocketAddr,
workload_handles: WorkloadHandles,
workload_configs: Arc<RwLock<HashMap<String, HttpWorkloadConfig>>>,
shutdown_tx: Arc<RwLock<Option<mpsc::Sender<()>>>>,
tls_acceptor: Option<TlsAcceptor>,
}
impl HttpServer {
pub fn new(addr: SocketAddr) -> Self {
Self {
addr,
workload_handles: Arc::default(),
workload_configs: Arc::default(),
shutdown_tx: Arc::new(RwLock::new(None)),
tls_acceptor: None,
}
}
pub async fn new_with_tls(
addr: SocketAddr,
cert_path: &Path,
key_path: &Path,
ca_path: Option<&Path>,
) -> anyhow::Result<Self> {
let tls_config = load_tls_config(cert_path, key_path, ca_path).await?;
let tls_acceptor = TlsAcceptor::from(Arc::new(tls_config));
Ok(Self {
addr,
workload_handles: Arc::default(),
workload_configs: Arc::default(),
shutdown_tx: Arc::new(RwLock::new(None)),
tls_acceptor: Some(tls_acceptor),
})
}
}
#[async_trait::async_trait]
impl HostPlugin for HttpServer {
fn id(&self) -> &'static str {
HTTP_SERVER_ID
}
fn world(&self) -> WitWorld {
WitWorld {
imports: HashSet::from([WitInterface::from(
"wasi:http/incoming-handler,outgoing-handler",
)]),
..Default::default()
}
}
async fn start(&self) -> anyhow::Result<()> {
let addr = self.addr;
let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
let shutdown_tx_clone = self.shutdown_tx.clone();
let workload_handles = self.workload_handles.clone();
let tls_acceptor = self.tls_acceptor.clone();
*shutdown_tx_clone.write().await = Some(shutdown_tx);
let listener = TcpListener::bind(addr).await?;
debug!(addr = ?addr, "HTTP server listening");
tokio::spawn(async move {
if let Err(e) =
run_http_server(listener, workload_handles, &mut shutdown_rx, tls_acceptor).await
{
error!(err = ?e, addr = ?addr, "HTTP server error");
}
});
let protocol = if self.tls_acceptor.is_some() {
"HTTPS"
} else {
"HTTP"
};
debug!(addr = ?addr, protocol = protocol, "HTTP server starting");
Ok(())
}
async fn on_component_bind(
&self,
component: &mut WorkloadComponent,
interfaces: std::collections::HashSet<crate::wit::WitInterface>,
) -> anyhow::Result<()> {
let Some(http_iface) = interfaces.iter().find(|iface| {
iface.namespace == "wasi"
&& iface.package == "http"
&& iface.interfaces.contains("incoming-handler")
}) else {
bail!(
"No wasi:http/incoming-handler interface found, plugin should not be bound to this workload"
);
};
if interfaces.len() > 1 {
warn!(
interfaces = ?interfaces,
"ignoring non-wasi:http/incoming-handler interfaces",
);
} else if http_iface.interfaces.len() > 1 {
warn!(
interfaces = ?http_iface.interfaces,
"ignoring non-incoming-handler interfaces",
);
}
let host_header = http_iface
.config
.get("host")
.cloned()
.unwrap_or_else(|| "*".to_string());
let id = component.id();
debug!(host = %host_header, workload_id = id, "binding HTTP config for workload");
let config = HttpWorkloadConfig { host_header };
self.workload_configs
.write()
.await
.insert(id.to_string(), config);
Ok(())
}
async fn on_workload_resolved(
&self,
resolved_handle: &ResolvedWorkload,
component_id: &str,
) -> anyhow::Result<()> {
let config = self
.workload_configs
.read()
.await
.get(component_id)
.cloned()
.ok_or_else(|| anyhow::anyhow!("No HTTP config found for workload {component_id}"))?;
debug!(host = %config.host_header, workload_id = resolved_handle.id(), component_id, "storing resolved workload handle");
let instance_pre = resolved_handle.instantiate_pre(component_id).await?;
self.workload_handles.write().await.insert(
config.host_header,
(
resolved_handle.clone(),
instance_pre,
component_id.to_string(),
),
);
Ok(())
}
async fn on_workload_unbind(
&self,
workload: &ResolvedWorkload,
_interfaces: HashSet<WitInterface>,
) -> anyhow::Result<()> {
debug!(workload_id = workload.id(), "removing HTTP workload handle");
let mut handles_guard = self.workload_handles.write().await;
handles_guard.retain(|_, (handle, _, _)| handle.id() != workload.id());
self.workload_configs.write().await.remove(workload.id());
Ok(())
}
async fn stop(&self) -> anyhow::Result<()> {
info!(addr = ?self.addr, "HTTP server stopping");
let mut shutdown_guard = self.shutdown_tx.write().await;
if let Some(tx) = shutdown_guard.take() {
let _ = tx.send(()).await;
}
Ok(())
}
}
async fn run_http_server(
listener: TcpListener,
workload_handles: WorkloadHandles,
shutdown_rx: &mut mpsc::Receiver<()>,
tls_acceptor: Option<TlsAcceptor>,
) -> anyhow::Result<()> {
loop {
tokio::select! {
_ = shutdown_rx.recv() => {
info!("HTTP server received shutdown signal");
break;
}
result = listener.accept() => {
match result {
Ok((client, client_addr)) => {
debug!(addr = ?client_addr, "new HTTP client connection");
let handles_clone = workload_handles.clone();
let tls_acceptor_clone = tls_acceptor.clone();
tokio::spawn(async move {
let service = hyper::service::service_fn(move |req| {
let handles = handles_clone.clone();
async move {
handle_http_request(req, handles).await
}
});
let result = if let Some(acceptor) = tls_acceptor_clone {
match acceptor.accept(client).await {
Ok(tls_stream) => {
http1::Builder::new()
.keep_alive(true)
.serve_connection(TokioIo::new(tls_stream), service)
.await
}
Err(e) => {
error!(addr = ?client_addr, err = ?e, "TLS handshake failed");
return;
}
}
} else {
http1::Builder::new()
.keep_alive(true)
.serve_connection(TokioIo::new(client), service)
.await
};
if let Err(e) = result {
error!(addr = ?client_addr, err = ?e, "error serving HTTP client");
}
});
}
Err(e) => {
error!(err = ?e, "failed to accept HTTP connection");
}
}
}
}
}
Ok(())
}
async fn handle_http_request(
req: hyper::Request<hyper::body::Incoming>,
workload_handles: WorkloadHandles,
) -> Result<hyper::Response<HyperOutgoingBody>, hyper::Error> {
let method = req.method().clone();
let uri = req.uri().clone();
let host_header = req
.headers()
.get("host")
.and_then(|h| h.to_str().ok())
.unwrap_or("<no host header>")
.to_string();
debug!(
method = %method,
uri = %uri,
host = %host_header,
"HTTP request received"
);
let workload_handle = {
let handles = workload_handles.read().await;
debug!(host = %host_header, "looking up workload handle for host header");
if let Some(handle) = handles.get(&host_header) {
Some(handle.clone())
} else {
debug!("No exact match for host header, trying wildcard '*'");
handles.get("*").cloned()
}
};
let response = match workload_handle {
Some((handle, instance_pre, component_id)) => {
match invoke_component_handler(handle, instance_pre, &component_id, req).await {
Ok(resp) => resp,
Err(e) => {
error!(err = ?e, host = %host_header, "failed to invoke component");
hyper::Response::builder()
.status(500)
.body(HyperOutgoingBody::default())
.unwrap()
}
}
}
None => {
warn!(host = %host_header, "No workload bound to host header or wildcard '*'");
hyper::Response::builder()
.status(404)
.body(HyperOutgoingBody::default())
.unwrap()
}
};
Ok(response)
}
async fn invoke_component_handler(
workload_handle: ResolvedWorkload,
instance_pre: InstancePre<Ctx>,
component_id: &str,
req: hyper::Request<hyper::body::Incoming>,
) -> anyhow::Result<hyper::Response<HyperOutgoingBody>> {
let mut store = workload_handle.new_store(component_id).await?;
handle_component_request(store.as_context_mut(), instance_pre, req).await
}
pub async fn handle_component_request<'a>(
mut store: StoreContextMut<'a, Ctx>,
pre: InstancePre<Ctx>,
req: hyper::Request<hyper::body::Incoming>,
) -> anyhow::Result<hyper::Response<HyperOutgoingBody>> {
let (sender, receiver) = tokio::sync::oneshot::channel();
let req = store.data_mut().new_incoming_request(Scheme::Http, req)?;
let out = store.data_mut().new_response_outparam(sender)?;
let pre = ProxyPre::new(pre).context("failed to instantiate proxy pre")?;
let proxy = pre.instantiate_async(&mut store).await?;
proxy
.wasi_http_incoming_handler()
.call_handle(&mut store, req, out)
.await?;
match receiver.await {
Ok(Ok(resp)) => Ok(resp),
Ok(Err(e)) => Err(e.into()),
Err(e) => {
error!(err = ?e, "error receiving http response");
Err(anyhow::anyhow!(
"oneshot channel closed but no response was sent"
))
}
}
}
async fn load_tls_config(
cert_path: &Path,
key_path: &Path,
ca_path: Option<&Path>,
) -> anyhow::Result<ServerConfig> {
let cert_data = tokio::fs::read(cert_path)
.await
.with_context(|| format!("Failed to read certificate file: {}", cert_path.display()))?;
let mut cert_reader = std::io::Cursor::new(cert_data);
let cert_chain: Vec<CertificateDer<'static>> = certs(&mut cert_reader)
.collect::<Result<Vec<_>, _>>()
.with_context(|| format!("Failed to parse certificate file: {}", cert_path.display()))?;
ensure!(
!cert_chain.is_empty(),
"No certificates found in file: {}",
cert_path.display()
);
let key_data = tokio::fs::read(key_path)
.await
.with_context(|| format!("Failed to read private key file: {}", key_path.display()))?;
let mut key_reader = std::io::Cursor::new(key_data);
let key = private_key(&mut key_reader)
.with_context(|| format!("Failed to parse private key file: {}", key_path.display()))?
.ok_or_else(|| anyhow::anyhow!("No private key found in file: {}", key_path.display()))?;
let config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(cert_chain, key)
.with_context(|| "Failed to create TLS configuration")?;
if let Some(ca_path) = ca_path {
let ca_data = tokio::fs::read(ca_path)
.await
.with_context(|| format!("Failed to read CA file: {}", ca_path.display()))?;
let mut ca_reader = std::io::Cursor::new(ca_data);
let ca_certs: Vec<CertificateDer<'static>> = certs(&mut ca_reader)
.collect::<Result<Vec<_>, _>>()
.with_context(|| format!("Failed to parse CA file: {}", ca_path.display()))?;
ensure!(
!ca_certs.is_empty(),
"No CA certificates found in file: {}",
ca_path.display()
);
debug!("CA certificate loaded, but client certificate verification not yet implemented");
}
Ok(config)
}