use std::collections::HashMap;
use std::sync::Arc;
use tokio::net::TcpListener;
use tokio::sync::RwLock;
use tracing::{error, info, warn};
use crate::acme::{self, AcmeManager};
use crate::{
FallbackConfig, RouteTarget, SharedCertResolver, SharedWasmTriggers, WasmInvoker,
serve_loop_with_fallback, tls,
};
pub async fn run_proxy_with_acme(
route_table: Arc<RwLock<HashMap<String, Vec<RouteTarget>>>>,
wasm_triggers: SharedWasmTriggers,
wasm_invoker: Option<WasmInvoker>,
acme_manager: AcmeManager,
domains: Vec<String>,
) -> anyhow::Result<SharedCertResolver> {
run_proxy_with_acme_and_fallback(
route_table,
wasm_triggers,
wasm_invoker,
acme_manager,
domains,
None,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn run_proxy_with_acme_and_fallback(
route_table: Arc<RwLock<HashMap<String, Vec<RouteTarget>>>>,
wasm_triggers: SharedWasmTriggers,
wasm_invoker: Option<WasmInvoker>,
acme_manager: AcmeManager,
domains: Vec<String>,
fallback: Option<FallbackConfig>,
) -> anyhow::Result<SharedCertResolver> {
run_acme_proxy_on(
(80, 443),
route_table,
wasm_triggers,
wasm_invoker,
acme_manager,
domains,
fallback,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn run_acme_proxy_on(
(http_port, https_port): (u16, u16),
route_table: Arc<RwLock<HashMap<String, Vec<RouteTarget>>>>,
wasm_triggers: SharedWasmTriggers,
wasm_invoker: Option<WasmInvoker>,
acme_manager: AcmeManager,
domains: Vec<String>,
fallback: Option<FallbackConfig>,
) -> anyhow::Result<SharedCertResolver> {
let http_listener = TcpListener::bind(("0.0.0.0", http_port))
.await
.map_err(|e| {
anyhow::anyhow!(
"cannot listen on port {http_port} ({e}): ACME HTTP-01 validation and the \
HTTP-to-HTTPS redirect need it, so no certificate can be issued"
)
})?;
let https_listener = TcpListener::bind(("0.0.0.0", https_port))
.await
.map_err(|e| {
anyhow::anyhow!(
"cannot listen on port {https_port} ({e}): no HTTPS traffic can be served"
)
})?;
info!("Reverse proxy listening on 0.0.0.0:{http_port} (HTTP) and 0.0.0.0:{https_port} (HTTPS)");
let resolver = Arc::new(acme::DynCertResolver::with_fallback()?);
let acme_mgr = acme_manager.clone();
let routes_clone = route_table.clone();
let triggers_clone = wasm_triggers.clone();
let invoker_clone = wasm_invoker.clone();
let fallback_http = fallback.clone();
let fallback_tls = fallback.clone();
let http_handle = tokio::spawn({
let acme = acme_mgr.clone();
let routes = routes_clone.clone();
let triggers = triggers_clone.clone();
let invoker = invoker_clone.clone();
async move {
if let Err(e) = serve_loop_with_fallback(
http_listener,
routes,
triggers,
invoker,
None,
Some(acme),
fallback_http,
)
.await
{
error!("HTTP listener failed: {e}");
}
}
});
let resolver_clone = resolver.clone();
let https_handle = tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
const PER_DOMAIN_PROVISION_TIMEOUT: std::time::Duration =
std::time::Duration::from_secs(60);
for domain in &domains {
acme_mgr.add_domain(domain).await;
let fut = acme_mgr.ensure_cert_for_resolver(domain, &resolver_clone);
match tokio::time::timeout(PER_DOMAIN_PROVISION_TIMEOUT, fut).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
error!(domain = %domain, error = %e, "Failed to provision cert");
}
Err(_) => {
warn!(
domain = %domain,
timeout_secs = PER_DOMAIN_PROVISION_TIMEOUT.as_secs(),
"Cert provisioning timed out — skipping (HTTPS will start without this cert; reconciler may retry on demand)"
);
}
}
}
let config = tls::with_h2_alpn(
rustls::ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(resolver_clone),
);
let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config));
info!(
"Starting HTTPS with SNI resolver ({} domains)",
domains.len()
);
let routes = routes_clone;
let triggers = triggers_clone;
let invoker = invoker_clone;
if let Err(e) = serve_loop_with_fallback(
https_listener,
routes,
triggers,
invoker,
Some(acceptor),
Some(acme_mgr),
fallback_tls,
)
.await
{
error!("HTTPS listener failed: {e}");
}
});
tokio::spawn(async move {
tokio::select! {
_ = http_handle => warn!("HTTP listener exited"),
_ = https_handle => warn!("HTTPS listener exited"),
}
});
Ok(resolver)
}
#[cfg(test)]
mod tests {
use super::*;
async fn start_on(ports: (u16, u16)) -> anyhow::Result<SharedCertResolver> {
let cache = tempfile::tempdir().unwrap();
run_acme_proxy_on(
ports,
Arc::new(RwLock::new(HashMap::new())),
Arc::new(RwLock::new(Vec::new())),
None,
AcmeManager::new("ops@example.com", cache.path()),
Vec::new(),
None,
)
.await
}
fn free_port() -> u16 {
std::net::TcpListener::bind("127.0.0.1:0")
.unwrap()
.local_addr()
.unwrap()
.port()
}
#[tokio::test]
async fn a_taken_http_port_is_an_error_naming_the_consequence() {
let held = std::net::TcpListener::bind("0.0.0.0:0").unwrap();
let port = held.local_addr().unwrap().port();
let Err(err) = start_on((port, free_port())).await else {
panic!("must fail");
};
let msg = format!("{err:#}");
assert!(msg.contains(&format!("port {port}")), "{msg}");
assert!(msg.contains("no certificate can be issued"), "{msg}");
}
#[tokio::test]
async fn a_taken_https_port_is_an_error_too() {
let held = std::net::TcpListener::bind("0.0.0.0:0").unwrap();
let port = held.local_addr().unwrap().port();
assert!(start_on((free_port(), port)).await.is_err());
}
#[tokio::test]
async fn free_ports_start() {
assert!(start_on((free_port(), free_port())).await.is_ok());
}
}