use std::net::SocketAddr;
use std::sync::Arc;
use axum::Router;
use axum::body::Body;
use axum::extract::{Request, State};
use axum::http::{HeaderMap, HeaderValue, StatusCode, Uri};
use axum::response::{IntoResponse, Response};
use hyper::header::{COOKIE, HOST};
const PITCHFORK_HEADER: &str = "x-pitchfork";
const PROXY_HOPS_HEADER: &str = "x-pitchfork-hops";
const MAX_PROXY_HOPS: u64 = 5;
const HANDSHAKE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
const MAX_HOST_CERTS: usize = 256;
const REFUSAL_LOG_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
static REFUSED_HANDSHAKE: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
#[cfg(feature = "proxy-tls")]
static REFUSED_SNI: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
pub(crate) const SHUTDOWN_DRAIN_BUDGET: std::time::Duration = std::time::Duration::from_secs(10);
const MAX_TUNNELS: usize = 256;
static REFUSED_TUNNEL: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
static ABANDONED_HANDSHAKE: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
static ABANDONED_TUNNEL: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
const MAX_PENDING_HANDSHAKES: usize = 512;
const HOP_BY_HOP_HEADERS: &[&str] = &[
"connection",
"keep-alive",
"proxy-connection",
"transfer-encoding",
"upgrade",
];
use hyper_util::client::legacy::Client;
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::rt::TokioExecutor;
use tokio::net::{TcpListener, TcpStream};
use crate::daemon_id::DaemonId;
use crate::pitchfork_toml::ProxyTlsMode;
use crate::proxy::activity::{ACTIVITY, ActivityGuard, GuardedBody};
use crate::settings::settings;
use crate::supervisor::SUPERVISOR;
const SLUG_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(2);
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ProxyTlsRoute {
pub mode: ProxyTlsMode,
pub port: Option<u16>,
}
#[derive(Clone, Debug)]
pub struct CachedSlugEntry {
pub slug: String,
pub namespace: Option<String>,
pub daemon_name: String,
pub dir: std::path::PathBuf,
pub worktrees: Vec<crate::proxy::worktree::WorktreeEntry>,
pub rejected_worktree_prefixes: std::collections::HashSet<String>,
pub tls: ProxyTlsRoute,
pub worktree_tls: std::collections::HashMap<String, ProxyTlsRoute>,
}
impl CachedSlugEntry {
fn known_route(
&self,
dir: &std::path::Path,
namespace: Option<&str>,
daemon_name: &str,
) -> Option<ProxyTlsRoute> {
(self.dir == dir
&& self.namespace.as_deref() == namespace
&& self.daemon_name == daemon_name)
.then_some(self.tls)
}
fn known_worktree_route(
&self,
wt: &crate::proxy::worktree::WorktreeEntry,
daemon_name: &str,
) -> Option<ProxyTlsRoute> {
if self.daemon_name != daemon_name {
return None;
}
let branch = wt.sanitized_branch.to_ascii_lowercase();
self.worktrees
.iter()
.any(|known| {
known.sanitized_branch.eq_ignore_ascii_case(&branch)
&& known.path == wt.path
&& known.namespace == wt.namespace
})
.then(|| self.worktree_tls.get(&branch).copied())
.flatten()
}
}
struct SlugCache {
entries: Arc<std::collections::HashMap<String, CachedSlugEntry>>,
expires_at: std::time::Instant,
}
static SLUG_CACHE: once_cell::sync::Lazy<std::sync::RwLock<SlugCache>> =
once_cell::sync::Lazy::new(|| {
std::sync::RwLock::new(SlugCache {
entries: Arc::new(std::collections::HashMap::new()),
expires_at: std::time::Instant::now(), })
});
fn slug_snapshot() -> Arc<std::collections::HashMap<String, CachedSlugEntry>> {
let guard = SLUG_CACHE
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
Arc::clone(&guard.entries)
}
static SLUG_REFRESH: once_cell::sync::Lazy<tokio::sync::Mutex<()>> =
once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(()));
fn fresh_slugs() -> Option<Arc<std::collections::HashMap<String, CachedSlugEntry>>> {
let cache = SLUG_CACHE
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
(std::time::Instant::now() < cache.expires_at).then(|| Arc::clone(&cache.entries))
}
fn reject_case_colliding_worktrees(
wts: Vec<crate::proxy::worktree::WorktreeEntry>,
) -> (
Vec<crate::proxy::worktree::WorktreeEntry>,
std::collections::HashSet<String>,
) {
let collisions =
crate::proxy::ascii_case_collisions(wts.iter().map(|w| w.sanitized_branch.as_str()));
if collisions.is_empty() {
return (wts, collisions);
}
let (dropped, kept): (Vec<_>, Vec<_>) = wts
.into_iter()
.partition(|w| collisions.contains(&w.sanitized_branch.to_ascii_lowercase()));
let mut folded: Vec<&String> = collisions.iter().collect();
folded.sort();
for key in folded {
let mut branches: Vec<&str> = dropped
.iter()
.filter(|w| w.sanitized_branch.eq_ignore_ascii_case(key))
.map(|w| w.branch.as_str())
.collect();
branches.sort();
log::warn!(
"Worktree slug collision: branches [{}] all route to '{key}' under \
case-insensitive host matching. None of them will be routed; \
rename a branch to disambiguate.",
branches.join(", "),
);
}
(kept, collisions)
}
pub(crate) fn read_proxy_tls_route(
dir: &std::path::Path,
namespace: Option<&str>,
daemon_name: &str,
) -> miette::Result<Option<ProxyTlsRoute>> {
let Some(id) = namespace.and_then(|ns| DaemonId::try_new(ns, daemon_name).ok()) else {
return Ok(None);
};
let pt = crate::pitchfork_toml::PitchforkToml::all_merged_from(dir)?;
Ok(pt.daemons.get(&id).map(|cfg| ProxyTlsRoute {
mode: cfg.proxy_tls.unwrap_or_default(),
port: cfg.effective_proxy_tls_port(),
}))
}
fn route_or_last_known(
read: miette::Result<Option<ProxyTlsRoute>>,
known: Option<ProxyTlsRoute>,
dir: &std::path::Path,
daemon_name: &str,
) -> Option<ProxyTlsRoute> {
read.unwrap_or_else(|e| {
crate::proxy::hostname::warn_once(&format!(
"Proxy TLS route for daemon '{daemon_name}': could not read config in {}; \
keeping its last known TLS mode until the config is fixed: {e}",
dir.display()
));
known
})
}
fn build_slug_entries(
previous: &std::collections::HashMap<String, CachedSlugEntry>,
) -> std::collections::HashMap<String, CachedSlugEntry> {
let global_slugs = crate::pitchfork_toml::PitchforkToml::read_global_slugs();
let collisions = crate::proxy::ascii_case_collisions(global_slugs.keys().map(String::as_str));
let mut folded: Vec<&String> = collisions.iter().collect();
folded.sort();
for key in folded {
let mut spellings: Vec<&str> = global_slugs
.keys()
.filter(|s| s.eq_ignore_ascii_case(key))
.map(String::as_str)
.collect();
spellings.sort();
log::warn!(
"Slug collision: [{}] differ only by case and host names are case-insensitive. \
None of them will be routed; remove or rename all but one.",
spellings.join(", "),
);
}
let mut entries: std::collections::HashMap<String, CachedSlugEntry> =
std::collections::HashMap::with_capacity(global_slugs.len());
let worktree_enabled = crate::settings::settings().general.worktree;
for (slug, entry) in &global_slugs {
let key = slug.to_ascii_lowercase();
if collisions.contains(&key) {
continue;
}
let ns = entry.resolve_namespace();
let daemon_name = entry.daemon.as_deref().unwrap_or(slug).to_string();
let (worktrees, rejected_worktree_prefixes) = if worktree_enabled {
let wts = match entry.resolve_dir() {
Some(dir) => crate::proxy::worktree::discover_worktrees(&dir),
None => vec![],
};
let wts = wts
.into_iter()
.map(|mut wt| {
wt.namespace =
crate::pitchfork_toml::PitchforkToml::namespace_for_dir(&wt.path).ok();
wt
})
.collect();
reject_case_colliding_worktrees(wts)
} else {
(vec![], std::collections::HashSet::new())
};
let dir = entry.resolve_dir().unwrap_or_default();
let prev = previous.get(&key);
let tls = route_or_last_known(
read_proxy_tls_route(&dir, ns.as_deref(), &daemon_name),
prev.and_then(|p| p.known_route(&dir, ns.as_deref(), &daemon_name)),
&dir,
&daemon_name,
)
.unwrap_or_default();
let worktree_tls = worktrees
.iter()
.filter_map(|wt| {
let branch = wt.sanitized_branch.to_ascii_lowercase();
route_or_last_known(
read_proxy_tls_route(&wt.path, wt.namespace.as_deref(), &daemon_name),
prev.and_then(|p| p.known_worktree_route(wt, &daemon_name)),
&wt.path,
&daemon_name,
)
.map(|route| (branch, route))
})
.collect();
entries.insert(
key,
CachedSlugEntry {
slug: slug.clone(),
namespace: ns,
daemon_name,
dir,
worktrees,
rejected_worktree_prefixes,
tls,
worktree_tls,
},
);
}
entries
}
pub async fn get_cached_slugs() -> Arc<std::collections::HashMap<String, CachedSlugEntry>> {
if let Some(entries) = fresh_slugs() {
return entries;
}
let _refreshing = SLUG_REFRESH.lock().await;
if let Some(entries) = fresh_slugs() {
return entries;
}
let previous = slug_snapshot();
let new_entries = Arc::new(
tokio::task::spawn_blocking(move || build_slug_entries(&previous))
.await
.unwrap_or_else(|e| {
log::warn!("Failed to refresh slug cache: {e}");
std::collections::HashMap::new()
}),
);
let mut cache = SLUG_CACHE
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
cache.entries = Arc::clone(&new_entries);
cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
new_entries
}
struct RegistryCache {
registry: Arc<crate::proxy::hostname::HostRegistry>,
expires_at: std::time::Instant,
}
static HOST_REGISTRY: once_cell::sync::Lazy<std::sync::RwLock<RegistryCache>> =
once_cell::sync::Lazy::new(|| {
std::sync::RwLock::new(RegistryCache {
registry: Arc::new(crate::proxy::hostname::HostRegistry::default()),
expires_at: std::time::Instant::now(), })
});
static REGISTRY_REFRESH: once_cell::sync::Lazy<tokio::sync::Mutex<()>> =
once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(()));
fn fresh_registry() -> Option<Arc<crate::proxy::hostname::HostRegistry>> {
let cache = HOST_REGISTRY
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
(std::time::Instant::now() < cache.expires_at).then(|| Arc::clone(&cache.registry))
}
fn registry_snapshot() -> Arc<crate::proxy::hostname::HostRegistry> {
let cache = HOST_REGISTRY
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
Arc::clone(&cache.registry)
}
pub async fn get_cached_host_registry() -> Arc<crate::proxy::hostname::HostRegistry> {
if let Some(registry) = fresh_registry() {
return registry;
}
let _refreshing = REGISTRY_REFRESH.lock().await;
if let Some(registry) = fresh_registry() {
return registry;
}
let registry = Arc::new(
tokio::task::spawn_blocking(crate::proxy::hostname::HostRegistry::build)
.await
.unwrap_or_else(|e| {
log::warn!("Failed to refresh hostname registry: {e}");
crate::proxy::hostname::HostRegistry::default()
}),
);
for err in ®istry.errors {
crate::proxy::hostname::warn_once(err);
}
let mut cache = HOST_REGISTRY
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
cache.registry = Arc::clone(®istry);
cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
registry
}
fn wildcard_slug_lookup<'a>(
subdomain: &str,
entries: &'a std::collections::HashMap<String, CachedSlugEntry>,
wildcard: bool,
) -> Option<&'a CachedSlugEntry> {
let subdomain = subdomain.to_ascii_lowercase();
entries.get(&subdomain).or_else(|| {
if !wildcard {
return None;
}
subdomain
.match_indices('.')
.map(|(i, _)| &subdomain[i + 1..])
.find_map(|candidate| entries.get(candidate))
})
}
#[derive(Debug)]
enum PrefixMatch<'a> {
Worktree(&'a crate::proxy::worktree::WorktreeEntry),
Unknown,
Ambiguous,
}
fn match_worktree_prefix<'a>(cached: &'a CachedSlugEntry, prefix: &str) -> PrefixMatch<'a> {
if let Some(wt) = cached
.worktrees
.iter()
.find(|w| w.sanitized_branch.eq_ignore_ascii_case(prefix))
{
return PrefixMatch::Worktree(wt);
}
if cached
.rejected_worktree_prefixes
.contains(&prefix.to_ascii_lowercase())
{
return PrefixMatch::Ambiguous;
}
PrefixMatch::Unknown
}
fn worktree_route(cached: &CachedSlugEntry, sanitized_branch: &str) -> ProxyTlsRoute {
cached
.worktree_tls
.get(&sanitized_branch.to_ascii_lowercase())
.copied()
.unwrap_or(ProxyTlsRoute {
mode: cached.tls.mode,
port: None,
})
}
fn strip_dot_suffix_ignore_case(s: &str, suffix: &str) -> Option<String> {
let needle_len = suffix.len() + 1;
if s.len() <= needle_len {
return None;
}
let split = s.len() - needle_len;
if !s.is_char_boundary(split) {
return None;
}
let (head, tail) = s.split_at(split);
if tail.starts_with('.') && tail[1..].eq_ignore_ascii_case(suffix) {
Some(head.to_string())
} else {
None
}
}
async fn cached_slug_lookup(subdomain: &str) -> Option<CachedSlugEntry> {
let entries = get_cached_slugs().await;
wildcard_slug_lookup(subdomain, &entries, settings().proxy.wildcard).cloned()
}
static AUTO_START_IN_PROGRESS: once_cell::sync::Lazy<
tokio::sync::Mutex<std::collections::HashSet<DaemonId>>,
> = once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(std::collections::HashSet::new()));
enum ResolveResult {
Ready(u16, Option<ActivityGuard>),
Starting { slug: String },
NotFound,
Page {
project: String,
worktree: Option<String>,
daemons: Vec<String>,
dir: Option<std::path::PathBuf>,
},
Unknown { heading: String, known: Vec<String> },
Error(String),
}
type OnErrorFn = Arc<dyn Fn(&str) + Send + Sync>;
#[derive(Clone)]
struct ProxyState {
client: Arc<Client<HttpConnector, Body>>,
tld: String,
is_tls: bool,
connect_target: Option<SocketAddr>,
contact_ip: std::net::IpAddr,
cancel: tokio_util::sync::CancellationToken,
tunnel_slots: Arc<tokio::sync::Semaphore>,
tunnels: tokio_util::task::TaskTracker,
on_error: Option<OnErrorFn>,
}
pub async fn serve(
bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
cancel: tokio_util::sync::CancellationToken,
) -> crate::Result<()> {
let s = settings();
let lan_enabled = s.proxy.lan || !s.proxy.lan_ip.is_empty();
let effective_tld = crate::proxy::effective_tld(&s).to_string();
let Some(effective_port) = u16::try_from(s.proxy.port).ok().filter(|&p| p > 0) else {
let msg = format!(
"proxy.port {} is out of valid port range (1-65535), proxy server cannot start",
s.proxy.port
);
let _ = bind_tx.send(Err(msg.clone()));
miette::bail!("{msg}");
};
let mut connector = HttpConnector::new();
connector.set_connect_timeout(Some(std::time::Duration::from_secs(10)));
let client = Client::builder(TokioExecutor::new())
.pool_idle_timeout(std::time::Duration::from_secs(30))
.build(connector);
let bind_ip: std::net::IpAddr = if lan_enabled && s.proxy.host == "127.0.0.1" {
std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)
} else {
match s.proxy.host.parse() {
Ok(ip) => ip,
Err(_) => {
log::warn!(
"proxy.host {:?} is not a valid IP address — falling back to 127.0.0.1. \
The proxy will only be reachable on the loopback interface.",
s.proxy.host
);
std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
}
}
};
let addr = SocketAddr::from((bind_ip, effective_port));
let contact_ip = local_contact_ip(bind_ip);
let tunnels = tokio_util::task::TaskTracker::new();
let state = ProxyState {
client: Arc::new(client),
tld: effective_tld.clone(),
is_tls: s.proxy.https,
connect_target: Some(SocketAddr::from((contact_ip, effective_port))),
contact_ip,
cancel: cancel.clone(),
tunnel_slots: Arc::new(tokio::sync::Semaphore::new(MAX_TUNNELS)),
tunnels: tunnels.clone(),
on_error: None,
};
let plain_state = state.clone();
let app = Router::new()
.route(crate::proxy::pac::PAC_PATH, axum::routing::any(pac_handler))
.fallback(proxy_handler)
.with_state(state);
if s.proxy.https {
serve_https_with_http_fallback(
app,
addr,
&s,
effective_port,
effective_tld,
plain_state,
bind_tx,
cancel,
)
.await
} else {
serve_http(app, addr, effective_port, plain_state, bind_tx, cancel).await
}
}
macro_rules! serve_conn_until_cancelled {
($io:expr, $svc:expr, $cancel:expr) => {{
let builder = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new());
let conn = builder.serve_connection_with_upgrades($io, $svc);
tokio::pin!(conn);
tokio::select! {
r = conn.as_mut() => r,
_ = $cancel.cancelled() => {
conn.as_mut().graceful_shutdown();
conn.await
}
}
}};
}
async fn serve_http(
app: Router,
addr: SocketAddr,
effective_port: u16,
proxy_state: ProxyState,
bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
cancel: tokio_util::sync::CancellationToken,
) -> crate::Result<()> {
let listener = match TcpListener::bind(addr).await {
Ok(l) => {
if settings().proxy.sync_hosts {
crate::proxy::hosts::sync_hosts_from_settings();
}
let _ = bind_tx.send(Ok(()));
l
}
Err(e) => {
let msg = bind_error_message(effective_port, &e);
let _ = bind_tx.send(Err(msg.clone()));
return Err(miette::miette!("{msg}"));
}
};
log::info!("Proxy server listening on http://{addr}");
if effective_port < 1024 {
log::info!(
"Note: port {effective_port} is a privileged port. \
The supervisor must be started with sudo to bind to this port."
);
}
let mut conn_tasks: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
loop {
while conn_tasks.try_join_next().is_some() {}
tokio::select! {
accept_result = listener.accept() => {
let (stream, peer_addr) = match accept_result {
Ok(conn) => conn,
Err(e) => {
log::warn!("Accept error (will retry): {e}");
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
continue;
}
};
let app = app
.clone()
.layer(axum::Extension(axum::extract::ConnectInfo(peer_addr)));
let cancel = cancel.clone();
conn_tasks.spawn(async move {
let io = hyper_util::rt::TokioIo::new(stream);
let svc = hyper_util::service::TowerToHyperService::new(app);
if let Err(e) = serve_conn_until_cancelled!(io, svc, cancel) {
log::debug!("Connection error: {e}");
}
});
}
_ = cancel.cancelled() => break,
}
}
let deadline = tokio::time::Instant::now() + SHUTDOWN_DRAIN_BUDGET;
if tokio::time::timeout_at(deadline, async {
while conn_tasks.join_next().await.is_some() {}
})
.await
.is_err()
{
log::debug!("Proxy connections still open after {SHUTDOWN_DRAIN_BUDGET:?}; aborting them");
}
drop(conn_tasks);
proxy_state.tunnels.close();
let _ = tokio::time::timeout_at(deadline, proxy_state.tunnels.wait()).await;
Ok(())
}
#[cfg(feature = "proxy-tls")]
#[allow(clippy::too_many_arguments)]
async fn serve_https_with_http_fallback(
app: Router,
addr: SocketAddr,
s: &crate::settings::Settings,
effective_port: u16,
effective_tld: String,
plain_proxy_state: ProxyState,
bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
cancel: tokio_util::sync::CancellationToken,
) -> crate::Result<()> {
use rustls::ServerConfig;
use tokio_rustls::TlsAcceptor;
let (cert_path, key_path) = resolve_tls_paths(s)?;
let _ = rustls::crypto::ring::default_provider().install_default();
let resolver: Arc<dyn rustls::server::ResolvesServerCert> = if s.proxy.tls_cert.is_empty() {
if ensure_ca(&cert_path, &key_path, || {
cert_path.exists() && key_path.exists()
})? {
log::info!("Generated local CA certificate at {}", cert_path.display());
log::info!("To trust the CA in your browser, run: pitchfork proxy trust");
}
Arc::new(SniCertResolver::new(
&cert_path,
&key_path,
effective_tld.clone(),
)?)
} else {
log::info!(
"Serving the configured certificate {} (no certificates are minted)",
cert_path.display()
);
Arc::new(StaticCertResolver::new(
&cert_path,
&key_path,
effective_tld.clone(),
)?)
};
let mut tls_config = ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(resolver);
tls_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
let acceptor = TlsAcceptor::from(Arc::new(tls_config));
let listener = match TcpListener::bind(addr).await {
Ok(l) => {
if settings().proxy.sync_hosts {
crate::proxy::hosts::sync_hosts_from_settings();
}
let _ = bind_tx.send(Ok(()));
l
}
Err(e) => {
let msg = bind_error_message(effective_port, &e);
let _ = bind_tx.send(Err(msg.clone()));
return Err(miette::miette!("{msg}"));
}
};
log::info!("Proxy server listening on https://{addr} (HTTP also accepted)");
if effective_port < 1024 {
log::info!(
"Note: port {effective_port} is a privileged port. \
The supervisor must be started with sudo to bind to this port."
);
}
let redirect_app = Router::new()
.route(
crate::proxy::pac::PAC_PATH,
axum::routing::any(plain_pac_handler),
)
.fallback(plain_fallback_handler)
.with_state(PlainState {
tld: effective_tld.clone(),
port: effective_port,
proxy: plain_proxy_state.clone(),
});
let mut conn_tasks: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
let handshake_slots = Arc::new(tokio::sync::Semaphore::new(MAX_PENDING_HANDSHAKES));
loop {
while conn_tasks.try_join_next().is_some() {}
tokio::select! {
accept_result = listener.accept() => {
let (stream, peer_addr) = match accept_result {
Ok(conn) => conn,
Err(e) => {
log::warn!("Accept error (will retry): {e}");
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
continue;
}
};
let Ok(handshake_permit) = Arc::clone(&handshake_slots).try_acquire_owned() else {
if let Some(suppressed) = REFUSED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL) {
log::warn!(
"Proxy refused a connection: {MAX_PENDING_HANDSHAKES} \
still negotiating \
({suppressed} similar refusals since the last message)"
);
}
drop(stream);
continue;
};
let acceptor = acceptor.clone();
let app = app
.clone()
.layer(axum::Extension(axum::extract::ConnectInfo(peer_addr)));
let redirect_app = redirect_app.clone();
let tld = effective_tld.clone();
let cancel = cancel.clone();
conn_tasks.spawn(async move {
let handshake_deadline = tokio::time::Instant::now() + HANDSHAKE_TIMEOUT;
let mut peek_buf = [0u8; 1];
match tokio::time::timeout_at(handshake_deadline, stream.peek(&mut peek_buf)).await {
Ok(Ok(0)) | Ok(Err(_)) => return,
Err(_) => {
if let Some(suppressed) =
ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
{
log::debug!(
"Connection sent nothing within the handshake timeout \
({suppressed} similar since the last message)"
);
}
return;
}
Ok(Ok(_)) => {}
}
if peek_buf[0] == 0x16 {
let sni_budget = SNI_PEEK_TIMEOUT.min(
handshake_deadline.saturating_duration_since(tokio::time::Instant::now()),
);
match peek_sni_host(&stream, sni_budget).await {
SniProbe::Host(host) => {
let Ok(mode) = tokio::time::timeout_at(
handshake_deadline,
resolve_tls_mode(&host, &tld),
)
.await
else {
if let Some(suppressed) =
ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
{
log::debug!(
"Routing lookup for '{host}' did not finish within the \
handshake timeout ({suppressed} similar since the \
last message)"
);
}
return;
};
if mode.is_passthrough() {
drop(handshake_permit);
tokio::select! {
_ = serve_passthrough(stream, &host, &tld) => {}
_ = cancel.cancelled() => {}
}
return;
}
}
SniProbe::NoHost => {}
SniProbe::Undetermined => {
log::debug!(
"Could not read the ClientHello of a TLS connection; \
handing it to the TLS acceptor, which refuses to terminate \
a passthrough hostname."
);
let _ = tokio::time::timeout_at(handshake_deadline, async {
let _ = get_cached_slugs().await;
let _ = get_cached_host_registry().await;
})
.await;
}
}
let accepted = match tokio::time::timeout_at(
handshake_deadline,
acceptor.accept(stream),
)
.await
{
Ok(r) => r,
Err(_) => {
if let Some(suppressed) =
ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
{
log::debug!(
"TLS handshake did not complete in time \
({suppressed} similar since the last message)"
);
}
return;
}
};
drop(handshake_permit);
match accepted {
Ok(tls_stream) => {
let io = hyper_util::rt::TokioIo::new(tls_stream);
let svc = hyper_util::service::TowerToHyperService::new(app);
if let Err(e) = serve_conn_until_cancelled!(io, svc, cancel) {
log::debug!("Connection error: {e}");
}
}
Err(e) => {
log::debug!("TLS handshake error: {e}");
}
}
} else {
drop(handshake_permit);
let io = hyper_util::rt::TokioIo::new(stream);
let svc = hyper_util::service::TowerToHyperService::new(redirect_app);
let _ = serve_conn_until_cancelled!(io, svc, cancel);
}
});
while conn_tasks.try_join_next().is_some() {}
}
_ = cancel.cancelled() => {
log::info!("Proxy server shutting down (cancel signal received)");
break;
}
}
}
let deadline = tokio::time::Instant::now() + SHUTDOWN_DRAIN_BUDGET;
let _ = tokio::time::timeout_at(deadline, async {
while conn_tasks.join_next().await.is_some() {}
})
.await;
plain_proxy_state.tunnels.close();
let _ = tokio::time::timeout_at(deadline, plain_proxy_state.tunnels.wait()).await;
Ok(())
}
#[cfg(feature = "proxy-tls")]
const SNI_PEEK_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
#[cfg(feature = "proxy-tls")]
const SNI_PEEK_MAX_BYTES: usize = 16 * 1024;
#[cfg(feature = "proxy-tls")]
#[derive(Debug, PartialEq, Eq)]
enum SniProbe {
Host(String),
NoHost,
Undetermined,
}
#[cfg(feature = "proxy-tls")]
async fn peek_sni_host(stream: &TcpStream, timeout: std::time::Duration) -> SniProbe {
use crate::proxy::sni::{SniPeek, parse_sni};
const MIN_PAUSE: std::time::Duration = std::time::Duration::from_millis(10);
const MAX_PAUSE: std::time::Duration = std::time::Duration::from_millis(200);
let deadline = tokio::time::Instant::now() + timeout;
let mut buf = vec![0u8; 2048];
let mut last_n = 0;
let mut pause = MIN_PAUSE;
loop {
let n = match tokio::time::timeout_at(deadline, stream.peek(&mut buf)).await {
Ok(Ok(0)) => return SniProbe::NoHost,
Ok(Ok(n)) => n,
Ok(Err(e)) => {
log::debug!("Failed to peek at a TLS connection: {e}");
return SniProbe::Undetermined;
}
Err(_elapsed) => {
log::debug!("Timed out waiting for a client that sent no ClientHello");
return SniProbe::Undetermined;
}
};
match parse_sni(&buf[..n]) {
SniPeek::Found(host) => return SniProbe::Host(host),
SniPeek::Absent | SniPeek::NotTls => return SniProbe::NoHost,
SniPeek::Incomplete => {}
}
if n == buf.len() && buf.len() < SNI_PEEK_MAX_BYTES {
buf.resize((buf.len() * 2).min(SNI_PEEK_MAX_BYTES), 0);
continue;
}
if n >= SNI_PEEK_MAX_BYTES {
log::debug!("Giving up on SNI after {n} bytes without a complete ClientHello");
return SniProbe::Undetermined;
}
if tokio::time::Instant::now() >= deadline {
log::debug!("Timed out waiting for a complete ClientHello ({n} bytes read)");
return SniProbe::Undetermined;
}
match stream.ready(tokio::io::Interest::READABLE).await {
Ok(ready) if ready.is_read_closed() => {
log::debug!("Client closed after {n} bytes of an incomplete ClientHello");
return SniProbe::Undetermined;
}
Ok(_) => {}
Err(e) => {
log::debug!("Failed to poll a TLS connection: {e}");
return SniProbe::Undetermined;
}
}
pause = if n > last_n {
MIN_PAUSE
} else {
(pause * 2).min(MAX_PAUSE)
};
last_n = n;
tokio::time::sleep_until((tokio::time::Instant::now() + pause).min(deadline)).await;
}
}
#[cfg(feature = "proxy-tls")]
async fn serve_passthrough(mut stream: TcpStream, host: &str, tld: &str) {
let (port, _activity) = match resolve_passthrough_port(host, tld).await {
Ok(ready) => ready,
Err(msg) => {
log::warn!("TLS passthrough for '{host}' failed: {msg}");
return;
}
};
let addr = SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, port));
let mut backend = match connect_backend(addr).await {
Ok(b) => b,
Err(e) => {
log::warn!("TLS passthrough for '{host}': failed to connect to {addr}: {e}");
return;
}
};
log::debug!("TLS passthrough: splicing '{host}' to {addr}");
if let Err(e) = tokio::io::copy_bidirectional(&mut stream, &mut backend).await {
log::debug!("TLS passthrough for '{host}' ended: {e}");
}
}
#[cfg(feature = "proxy-tls")]
const PASSTHROUGH_CONNECT_GRACE: std::time::Duration = std::time::Duration::from_secs(2);
#[cfg(feature = "proxy-tls")]
async fn connect_backend(addr: SocketAddr) -> std::io::Result<TcpStream> {
let deadline = tokio::time::Instant::now() + PASSTHROUGH_CONNECT_GRACE;
loop {
match TcpStream::connect(addr).await {
Ok(stream) => return Ok(stream),
Err(e) => {
let retryable = e.kind() == std::io::ErrorKind::ConnectionRefused;
if !retryable || tokio::time::Instant::now() >= deadline {
return Err(e);
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
}
}
}
#[cfg(feature = "proxy-tls")]
async fn resolve_passthrough_port(
host: &str,
tld: &str,
) -> std::result::Result<(u16, Option<ActivityGuard>), String> {
let budget = settings().proxy_auto_start_timeout();
match tokio::time::timeout(budget, resolve_passthrough_port_inner(host, tld)).await {
Ok(result) => result,
Err(_elapsed) => Err(format!(
"no daemon was ready for '{host}' within proxy.auto_start_timeout ({budget:?})"
)),
}
}
#[cfg(feature = "proxy-tls")]
async fn resolve_passthrough_port_inner(
host: &str,
tld: &str,
) -> std::result::Result<(u16, Option<ActivityGuard>), String> {
loop {
match resolve_target(host, tld).await {
ResolveResult::Ready(port, activity) => return Ok((port, activity)),
ResolveResult::Starting { slug: _ } => {
tokio::time::sleep(std::time::Duration::from_millis(250)).await;
}
ResolveResult::NotFound => {
return Err(
"no running daemon with a port matched this hostname, and it could not \
be auto-started"
.to_string(),
);
}
ResolveResult::Page {
project, worktree, ..
} => {
return Err(match worktree {
Some(worktree) => {
format!("'{worktree}' of project '{project}' is a stack page, not a daemon")
}
None => format!("'{project}' is a project page, not a daemon"),
});
}
ResolveResult::Unknown { heading, .. } => return Err(heading),
ResolveResult::Error(msg) => return Err(msg),
}
}
}
#[cfg(not(feature = "proxy-tls"))]
#[allow(clippy::too_many_arguments)]
async fn serve_https_with_http_fallback(
_app: Router,
_addr: SocketAddr,
_s: &crate::settings::Settings,
_effective_port: u16,
_effective_tld: String,
_plain_proxy_state: ProxyState,
bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
_cancel: tokio_util::sync::CancellationToken,
) -> crate::Result<()> {
let msg = "HTTPS proxy support requires the `proxy-tls` feature.\n\
Rebuild pitchfork with: cargo build --features proxy-tls"
.to_string();
let _ = bind_tx.send(Err(msg.clone()));
miette::bail!("{msg}")
}
#[cfg(feature = "proxy-tls")]
fn resolve_tls_paths(
s: &crate::settings::Settings,
) -> crate::Result<(std::path::PathBuf, std::path::PathBuf)> {
if let Some(problem) = tls_pair_problem(&s.proxy.tls_cert, &s.proxy.tls_key) {
miette::bail!("{problem}");
}
let proxy_dir = crate::env::PITCHFORK_STATE_DIR.join("proxy");
let resolve = |configured: &str, default: &str| {
if configured.is_empty() {
proxy_dir.join(default)
} else {
std::path::PathBuf::from(configured)
}
};
Ok((
resolve(&s.proxy.tls_cert, "ca.pem"),
resolve(&s.proxy.tls_key, "ca-key.pem"),
))
}
pub(crate) fn tls_pair_problem(cert: &str, key: &str) -> Option<String> {
match (cert.is_empty(), key.is_empty()) {
(false, true) => Some(
"proxy.tls_cert is set but proxy.tls_key is empty; set both, or neither to use \
the generated CA"
.to_string(),
),
(true, false) => Some(
"proxy.tls_key is set but proxy.tls_cert is empty; set both, or neither to use \
the generated CA"
.to_string(),
),
_ => None,
}
}
#[cfg(feature = "proxy-tls")]
pub fn ensure_ca(
cert_path: &std::path::Path,
key_path: &std::path::Path,
usable: impl FnOnce() -> bool,
) -> crate::Result<bool> {
let _lock = xx::fslock::get(cert_path, false)
.map_err(|e| miette::miette!("Failed to lock {}: {e}", cert_path.display()))?;
if usable() {
return Ok(false);
}
generate_ca(cert_path, key_path)?;
clear_host_certs(&host_certs_dir_for(cert_path), None);
Ok(true)
}
#[cfg(feature = "proxy-tls")]
fn host_certs_dir_for(ca_cert_path: &std::path::Path) -> std::path::PathBuf {
ca_cert_path
.parent()
.unwrap_or(std::path::Path::new("."))
.join("host-certs")
}
#[cfg(feature = "proxy-tls")]
fn ca_cache_id(ca_cert_pem: &str) -> String {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
for b in ca_cert_pem.bytes() {
hash ^= u64::from(b);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
format!("ca-{hash:016x}")
}
#[cfg(feature = "proxy-tls")]
fn clear_host_certs(root: &std::path::Path, keep: Option<&str>) {
let Ok(entries) = std::fs::read_dir(root) else {
return;
};
for entry in entries.filter_map(|e| e.ok()) {
let path = entry.path();
let name = entry.file_name();
let result = if path.is_dir() {
let is_ca_dir = name.to_str().is_some_and(|n| n.starts_with("ca-"));
if !is_ca_dir || keep.is_some_and(|k| name == k) {
continue;
}
std::fs::remove_dir_all(&path)
} else if path.extension().is_some_and(|x| x == "pem") {
std::fs::remove_file(&path)
} else {
continue;
};
if let Err(e) = result {
log::debug!(
"Could not remove stale cached certs {}: {e}",
path.display()
);
}
}
}
#[cfg(feature = "proxy-tls")]
fn generate_ca(cert_path: &std::path::Path, key_path: &std::path::Path) -> crate::Result<()> {
use rcgen::{
BasicConstraints, CertificateParams, DistinguishedName, DnType, IsCa, KeyUsagePurpose,
};
if let Some(parent) = cert_path.parent() {
std::fs::create_dir_all(parent)
.map_err(|e| miette::miette!("Failed to create proxy cert directory: {e}"))?;
}
let mut params = CertificateParams::default();
let mut dn = DistinguishedName::new();
dn.push(DnType::CommonName, "Pitchfork Local CA");
dn.push(DnType::OrganizationName, "Pitchfork");
params.distinguished_name = dn;
params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
params.key_usages = vec![KeyUsagePurpose::KeyCertSign, KeyUsagePurpose::CrlSign];
let key_pair = rcgen::KeyPair::generate()
.map_err(|e| miette::miette!("Failed to generate CA key pair: {e}"))?;
let ca_cert = params
.self_signed(&key_pair)
.map_err(|e| miette::miette!("Failed to self-sign CA certificate: {e}"))?;
std::fs::write(cert_path, ca_cert.pem()).map_err(|e| {
miette::miette!(
"Failed to write CA certificate to {}: {e}",
cert_path.display()
)
})?;
{
#[cfg(unix)]
{
use std::io::Write;
use std::os::unix::fs::OpenOptionsExt;
std::fs::OpenOptions::new()
.write(true)
.create(true)
.truncate(true)
.mode(0o600)
.open(key_path)
.and_then(|mut f| f.write_all(key_pair.serialize_pem().as_bytes()))
.map_err(|e| {
miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
})?;
}
#[cfg(not(unix))]
{
std::fs::write(key_path, key_pair.serialize_pem()).map_err(|e| {
miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
})?;
log::debug!(
"CA private key written to {} (file permissions are not restricted \
on non-Unix platforms — consider restricting access manually)",
key_path.display()
);
}
}
Ok(())
}
#[cfg(feature = "proxy-tls")]
struct StaticCertResolver {
certified: Arc<rustls::sign::CertifiedKey>,
tld: String,
}
#[cfg(feature = "proxy-tls")]
impl std::fmt::Debug for StaticCertResolver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StaticCertResolver").finish_non_exhaustive()
}
}
#[cfg(feature = "proxy-tls")]
impl StaticCertResolver {
fn new(
cert_path: &std::path::Path,
key_path: &std::path::Path,
tld: String,
) -> crate::Result<Self> {
use rustls::pki_types::CertificateDer;
use rustls_pemfile::{certs, private_key};
let cert_pem = std::fs::read(cert_path).map_err(|e| {
miette::miette!("Failed to read proxy.tls_cert {}: {e}", cert_path.display())
})?;
let key_pem = std::fs::read(key_path).map_err(|e| {
miette::miette!("Failed to read proxy.tls_key {}: {e}", key_path.display())
})?;
let cert_ders: Vec<CertificateDer<'static>> = certs(&mut cert_pem.as_slice())
.collect::<Result<Vec<_>, _>>()
.map_err(|e| miette::miette!("Failed to parse {}: {e}", cert_path.display()))?;
if cert_ders.is_empty() {
miette::bail!("No certificates found in {}", cert_path.display());
}
let key_der = private_key(&mut key_pem.as_slice())
.map_err(|e| miette::miette!("Failed to parse {}: {e}", key_path.display()))?
.ok_or_else(|| miette::miette!("No private key found in {}", key_path.display()))?;
let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
.map_err(|e| miette::miette!("Failed to use the configured private key: {e}"))?;
let certified = rustls::sign::CertifiedKey::new(cert_ders, signing_key);
certified.keys_match().map_err(|e| {
miette::miette!(
"proxy.tls_key {} does not match proxy.tls_cert {}: {e}",
key_path.display(),
cert_path.display()
)
})?;
Ok(Self {
certified: Arc::new(certified),
tld,
})
}
}
#[cfg(feature = "proxy-tls")]
impl rustls::server::ResolvesServerCert for StaticCertResolver {
fn resolve(
&self,
client_hello: rustls::server::ClientHello<'_>,
) -> Option<Arc<rustls::sign::CertifiedKey>> {
if let Some(domain) = client_hello.server_name()
&& resolve_tls_mode_in(domain, &self.tld, &slug_snapshot(), ®istry_snapshot())
.is_passthrough()
{
log::warn!(
"Refusing to terminate TLS for '{domain}', which is configured for \
proxy_tls = \"passthrough\": its ClientHello could not be inspected before the \
handshake, so the stream could not be spliced to the daemon."
);
return None;
}
Some(Arc::clone(&self.certified))
}
}
#[cfg(feature = "proxy-tls")]
struct SniCertResolver {
issuer: rcgen::Issuer<'static, rcgen::KeyPair>,
tld: String,
host_certs_dir: std::path::PathBuf,
cache: std::sync::Mutex<CertCache>,
pending: std::sync::Mutex<std::collections::HashSet<String>>,
pending_cv: std::sync::Condvar,
}
#[cfg(feature = "proxy-tls")]
fn prune_host_certs(dir: &std::path::Path) {
let Ok(entries) = std::fs::read_dir(dir) else {
return;
};
let mut files: Vec<(std::time::SystemTime, std::path::PathBuf)> = entries
.filter_map(|e| e.ok())
.filter(|e| e.path().extension().is_some_and(|x| x == "pem"))
.filter_map(|e| {
let modified = e.metadata().and_then(|m| m.modified()).ok()?;
Some((modified, e.path()))
})
.collect();
if files.len() <= MAX_HOST_CERTS {
return;
}
files.sort_by_key(|(t, _)| *t);
let excess = files.len() - MAX_HOST_CERTS;
for (_, path) in files.into_iter().take(excess) {
if let Err(e) = std::fs::remove_file(&path) {
log::debug!("Could not prune cached cert {}: {e}", path.display());
}
}
}
#[cfg(feature = "proxy-tls")]
#[derive(Default)]
struct CertCache {
by_domain: std::collections::HashMap<String, Arc<rustls::sign::CertifiedKey>>,
order: std::collections::VecDeque<String>,
}
#[cfg(feature = "proxy-tls")]
impl CertCache {
fn get(&self, domain: &str) -> Option<&Arc<rustls::sign::CertifiedKey>> {
self.by_domain.get(domain)
}
fn insert(&mut self, domain: String, key: Arc<rustls::sign::CertifiedKey>) -> Option<String> {
if self.by_domain.insert(domain.clone(), key).is_none() {
self.order.push_back(domain);
}
if self.by_domain.len() <= MAX_HOST_CERTS {
return None;
}
let evicted = self.order.pop_front()?;
self.by_domain.remove(&evicted);
Some(evicted)
}
}
#[cfg(feature = "proxy-tls")]
impl std::fmt::Debug for SniCertResolver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SniCertResolver").finish_non_exhaustive()
}
}
#[cfg(feature = "proxy-tls")]
impl SniCertResolver {
fn new(
ca_cert_path: &std::path::Path,
ca_key_path: &std::path::Path,
tld: String,
) -> crate::Result<Self> {
let ca_key_pem = std::fs::read_to_string(ca_key_path)
.map_err(|e| miette::miette!("Failed to read CA key {}: {e}", ca_key_path.display()))?;
let ca_cert_pem = std::fs::read_to_string(ca_cert_path).map_err(|e| {
miette::miette!("Failed to read CA cert {}: {e}", ca_cert_path.display())
})?;
if !ca_cert_pem.contains("BEGIN CERTIFICATE") {
miette::bail!("CA cert file does not contain a valid PEM certificate");
}
let ca_key = rcgen::KeyPair::from_pem(&ca_key_pem)
.map_err(|e| miette::miette!("Failed to parse CA key: {e}"))?;
let issuer = rcgen::Issuer::from_ca_cert_pem(&ca_cert_pem, ca_key)
.map_err(|e| miette::miette!("Failed to parse CA cert: {e}"))?;
let host_certs_root = host_certs_dir_for(ca_cert_path);
let cache_id = ca_cache_id(&ca_cert_pem);
clear_host_certs(&host_certs_root, Some(&cache_id));
let host_certs_dir = host_certs_root.join(&cache_id);
std::fs::create_dir_all(&host_certs_dir)
.map_err(|e| miette::miette!("Failed to create host-certs dir: {e}"))?;
prune_host_certs(&host_certs_dir);
Ok(Self {
issuer,
tld,
host_certs_dir,
cache: std::sync::Mutex::new(CertCache::default()),
pending: std::sync::Mutex::new(std::collections::HashSet::new()),
pending_cv: std::sync::Condvar::new(),
})
}
fn get_or_create_checked(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
if !crate::proxy::owns_name(&self.tld, domain) {
if let Some(suppressed) = REFUSED_SNI.allow(REFUSAL_LOG_INTERVAL) {
log::warn!(
"Refusing to issue a certificate for {domain:?}: \
the pitchfork CA only signs names under .{} \
({suppressed} similar refusals since the last message)",
self.tld
);
}
return None;
}
self.get_or_create(domain)
}
fn get_or_create(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
{
let cache = self.cache.lock().ok()?;
if let Some(ck) = cache.get(domain) {
return Some(Arc::clone(ck));
}
}
loop {
{
let mut pending = self.pending.lock().ok()?;
if pending.contains(domain) {
pending = self.pending_cv.wait(pending).ok()?;
drop(pending);
} else {
pending.insert(domain.to_string());
break;
}
}
{
let cache = self.cache.lock().ok()?;
if let Some(ck) = cache.get(domain) {
return Some(Arc::clone(ck));
}
} }
let result = self.get_or_create_inner(domain);
{
let mut pending = match self.pending.lock() {
Ok(g) => g,
Err(e) => e.into_inner(),
};
pending.remove(domain);
self.pending_cv.notify_all();
}
result
}
fn get_or_create_inner(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
let disk_path = self.disk_path(domain);
if disk_path.exists() {
if let Ok(ck) = self.load_from_disk(&disk_path) {
let ck = Arc::new(ck);
self.remember(domain, &ck);
return Some(ck);
}
let _ = std::fs::remove_file(&disk_path);
}
let ck = self.sign_for_domain(domain).ok()?;
let ck = Arc::new(ck);
self.remember(domain, &ck);
Some(ck)
}
fn remember(&self, domain: &str, ck: &Arc<rustls::sign::CertifiedKey>) {
let evicted = match self.cache.lock() {
Ok(mut cache) => cache.insert(domain.to_string(), Arc::clone(ck)),
Err(_) => return,
};
if let Some(evicted) = evicted {
let path = self.disk_path(&evicted);
if let Err(e) = std::fs::remove_file(&path)
&& e.kind() != std::io::ErrorKind::NotFound
{
log::debug!("Could not evict cached cert {}: {e}", path.display());
}
}
}
fn disk_path(&self, domain: &str) -> std::path::PathBuf {
self.host_certs_dir
.join(format!("{}.pem", cert_cache_file_stem(domain)))
}
fn load_from_disk(&self, path: &std::path::Path) -> crate::Result<rustls::sign::CertifiedKey> {
use rustls::pki_types::CertificateDer;
use rustls_pemfile::{certs, private_key};
let pem = std::fs::read_to_string(path)
.map_err(|e| miette::miette!("Failed to read disk cert {}: {e}", path.display()))?;
let cert_ders: Vec<CertificateDer<'static>> = certs(&mut pem.as_bytes())
.collect::<Result<Vec<_>, _>>()
.map_err(|e| miette::miette!("Failed to parse certs from {}: {e}", path.display()))?;
if cert_ders.is_empty() {
miette::bail!("No certificates found in {}", path.display());
}
{
let (_, cert) = x509_parser::parse_x509_certificate(&cert_ders[0]).map_err(|e| {
miette::miette!("Failed to parse certificate from {}: {e}", path.display())
})?;
use chrono::Utc;
let now_ts = Utc::now().timestamp();
let not_after_ts = cert.validity().not_after.timestamp();
if not_after_ts < now_ts {
miette::bail!(
"Cached certificate at {} has expired — will regenerate",
path.display()
);
}
}
let key_der = private_key(&mut pem.as_bytes())
.map_err(|e| miette::miette!("Failed to parse key from {}: {e}", path.display()))?
.ok_or_else(|| miette::miette!("No private key found in {}", path.display()))?;
let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
.map_err(|e| miette::miette!("Failed to create signing key from disk: {e}"))?;
Ok(rustls::sign::CertifiedKey::new(cert_ders, signing_key))
}
fn sign_for_domain(&self, domain: &str) -> crate::Result<rustls::sign::CertifiedKey> {
use rcgen::date_time_ymd;
use rcgen::{CertificateParams, DistinguishedName, DnType, SanType};
use rustls::pki_types::CertificateDer;
use rustls_pemfile::private_key;
let mut params = CertificateParams::default();
let mut dn = DistinguishedName::new();
dn.push(DnType::CommonName, domain);
params.distinguished_name = dn;
{
use chrono::{Datelike, Duration, Utc};
let yesterday = Utc::now() - Duration::days(1);
let expiry = Utc::now() + Duration::days(397);
params.not_before = date_time_ymd(
yesterday.year(),
yesterday.month() as u8,
yesterday.day() as u8,
);
params.not_after =
date_time_ymd(expiry.year(), expiry.month() as u8, expiry.day() as u8);
}
let mut sans =
vec![SanType::DnsName(domain.to_string().try_into().map_err(
|e| miette::miette!("Invalid domain name '{domain}': {e}"),
)?)];
if let Some(dot_pos) = domain.find('.') {
let parent = &domain[dot_pos + 1..];
if crate::proxy::is_strictly_under_tld(&self.tld, parent) {
let wildcard = format!("*.{parent}");
if let Ok(wc) = wildcard.try_into() {
sans.push(SanType::DnsName(wc));
}
}
}
params.subject_alt_names = sans;
let leaf_key = rcgen::KeyPair::generate()
.map_err(|e| miette::miette!("Failed to generate leaf key: {e}"))?;
let leaf_cert = params
.signed_by(&leaf_key, &self.issuer)
.map_err(|e| miette::miette!("Failed to sign leaf cert for '{domain}': {e}"))?;
let cert_der = CertificateDer::from(leaf_cert.der().to_vec());
let key_pem = leaf_key.serialize_pem();
let key_der = private_key(&mut key_pem.as_bytes())
.map_err(|e| miette::miette!("Failed to parse leaf key PEM: {e}"))?
.ok_or_else(|| miette::miette!("No private key found in generated PEM"))?;
let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
.map_err(|e| miette::miette!("Failed to create signing key: {e}"))?;
let disk_path = self.disk_path(domain);
let combined_pem = format!("{}{}", leaf_cert.pem(), key_pem);
{
#[cfg(unix)]
{
use std::io::Write;
use std::os::unix::fs::OpenOptionsExt;
if let Err(e) = std::fs::OpenOptions::new()
.write(true)
.create(true)
.truncate(true)
.mode(0o600)
.open(&disk_path)
.and_then(|mut f| f.write_all(combined_pem.as_bytes()))
{
log::warn!(
"Failed to persist cert for '{domain}' to {}: {e}",
disk_path.display()
);
}
}
#[cfg(not(unix))]
{
if let Err(e) = std::fs::write(&disk_path, combined_pem) {
log::warn!(
"Failed to persist cert for '{domain}' to {}: {e}",
disk_path.display()
);
} else {
log::debug!(
"Leaf cert for '{domain}' written to {} (file permissions are not \
restricted on non-Unix platforms — consider restricting access manually)",
disk_path.display()
);
}
}
}
Ok(rustls::sign::CertifiedKey::new(vec![cert_der], signing_key))
}
}
#[cfg(feature = "proxy-tls")]
impl rustls::server::ResolvesServerCert for SniCertResolver {
fn resolve(
&self,
client_hello: rustls::server::ClientHello<'_>,
) -> Option<Arc<rustls::sign::CertifiedKey>> {
let domain = client_hello.server_name()?;
if resolve_tls_mode_in(domain, &self.tld, &slug_snapshot(), ®istry_snapshot())
.is_passthrough()
{
log::warn!(
"Refusing to terminate TLS for '{domain}', which is configured for \
proxy_tls = \"passthrough\": its ClientHello could not be inspected before the \
handshake, so the stream could not be spliced to the daemon."
);
return None;
}
self.get_or_create_checked(domain)
}
}
#[cfg(feature = "proxy-tls")]
pub(crate) fn ca_pair_problem(cert: &std::path::Path, key: &std::path::Path) -> Option<String> {
use rcgen::PublicKeyData;
let cert_pem = match std::fs::read_to_string(cert) {
Ok(p) => p,
Err(e) => return Some(format!("cannot read {}: {e}", cert.display())),
};
let key_pem = match std::fs::read_to_string(key) {
Ok(p) => p,
Err(e) => return Some(format!("cannot read {}: {e}", key.display())),
};
let key_pair = match rcgen::KeyPair::from_pem(&key_pem) {
Ok(k) => k,
Err(e) => return Some(format!("cannot parse {}: {e}", key.display())),
};
let Some(Ok(der)) = rustls_pemfile::certs(&mut cert_pem.as_bytes()).next() else {
return Some(format!("no certificate in {}", cert.display()));
};
let parsed = match x509_parser::parse_x509_certificate(&der) {
Ok((_, c)) => c,
Err(e) => return Some(format!("cannot parse {}: {e}", cert.display())),
};
if parsed.public_key().subject_public_key.data.as_ref() != key_pair.der_bytes() {
return Some(format!(
"{} is not the key for {}",
key.display(),
cert.display()
));
}
if let Err(e) = rcgen::Issuer::from_ca_cert_pem(&cert_pem, key_pair) {
return Some(format!("{} cannot sign certificates: {e}", cert.display()));
}
None
}
#[cfg(feature = "proxy-tls")]
fn cert_cache_file_stem(domain: &str) -> String {
let mut out = String::with_capacity(domain.len());
for b in domain.bytes() {
if b.is_ascii_alphanumeric() || b == b'-' || b == b'.' {
out.push(char::from(b));
} else {
out.push_str(&format!("%{b:02X}"));
}
}
out
}
#[derive(Clone)]
struct PlainState {
tld: String,
port: u16,
proxy: ProxyState,
}
async fn plain_fallback_handler(State(state): State<PlainState>, req: Request) -> Response {
if req.method() == axum::http::Method::CONNECT {
let raw_host = get_request_host(&req).unwrap_or_default();
return connect_handler(&state.proxy, req, &raw_host).await;
}
redirect_to_https_handler(req).await
}
fn reject_non_read(method: &axum::http::Method) -> Option<Response> {
if matches!(*method, axum::http::Method::GET | axum::http::Method::HEAD) {
return None;
}
let mut res = error_response(
StatusCode::METHOD_NOT_ALLOWED,
"the proxy auto-config file is read-only\n",
);
res.headers_mut().insert(
axum::http::header::ALLOW,
HeaderValue::from_static("GET, HEAD"),
);
Some(res)
}
fn pac_response(tld: &str, host: &str, port: u16) -> Response {
match crate::proxy::pac::generate(tld, host, port) {
Ok(body) => (
StatusCode::OK,
[(
axum::http::header::CONTENT_TYPE,
"application/x-ns-proxy-autoconfig",
)],
body,
)
.into_response(),
Err(e) => error_response(StatusCode::INTERNAL_SERVER_ERROR, &e.to_string()),
}
}
async fn pac_handler(State(state): State<ProxyState>, req: Request) -> Response {
let host = get_request_host(&req).unwrap_or_default();
let bare = host.split(':').next().unwrap_or("");
if !bare.is_empty() && crate::proxy::is_strictly_under_tld(&state.tld, bare) {
return proxy_handler(State(state), req).await;
}
if let Some(deny) = reject_non_read(req.method()) {
return deny;
}
let port = req
.uri()
.authority()
.and_then(|a| a.port_u16())
.or_else(|| host.rsplit(':').next().and_then(|p| p.parse().ok()))
.unwrap_or(if state.is_tls { 443 } else { 80 });
pac_response(&state.tld, &url_host(state.contact_ip), port)
}
async fn plain_pac_handler(State(state): State<PlainState>, req: Request) -> Response {
let host = get_request_host(&req).unwrap_or_default();
let bare = host.split(':').next().unwrap_or("");
if !bare.is_empty() && crate::proxy::is_strictly_under_tld(&state.tld, bare) {
return redirect_to_https_handler(req).await;
}
if let Some(deny) = reject_non_read(req.method()) {
return deny;
}
pac_response(&state.tld, &url_host(state.proxy.contact_ip), state.port)
}
fn local_contact_ip(bind_ip: std::net::IpAddr) -> std::net::IpAddr {
match bind_ip {
std::net::IpAddr::V4(ip) if ip.is_unspecified() => {
std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
}
std::net::IpAddr::V6(ip) if ip.is_unspecified() => {
std::net::IpAddr::V6(std::net::Ipv6Addr::LOCALHOST)
}
ip => ip,
}
}
fn url_host(ip: std::net::IpAddr) -> String {
match ip {
std::net::IpAddr::V6(ip) => format!("[{ip}]"),
std::net::IpAddr::V4(ip) => ip.to_string(),
}
}
fn connect_port(authority: &str) -> Option<u16> {
let (host, port) = authority.rsplit_once(':')?;
if host.contains(':') && !host.ends_with(']') {
return None;
}
port.parse().ok()
}
async fn connect_handler(state: &ProxyState, req: Request, raw_host: &str) -> Response {
let authority = req
.uri()
.authority()
.map(|a| a.as_str().to_string())
.unwrap_or_else(|| raw_host.to_string());
let host = authority
.rsplit_once(':')
.map(|(h, _)| h)
.unwrap_or(&authority)
.trim_start_matches('[')
.trim_end_matches(']')
.to_string();
if !crate::proxy::owns_name(&state.tld, &host) {
return error_response(
StatusCode::FORBIDDEN,
&format!(
"pitchfork only tunnels CONNECT for names under .{} — refusing {host}",
state.tld
),
);
}
if !state.is_tls && connect_port(&authority) == Some(443) {
return error_response(
StatusCode::BAD_GATEWAY,
&format!(
"proxy.https is false, so pitchfork cannot serve https://{host}; \
use http:// or enable proxy.https"
),
);
}
let Some(target) = state.connect_target else {
return error_response(
StatusCode::SERVICE_UNAVAILABLE,
"The proxy listener address is unknown, so CONNECT cannot be tunnelled",
);
};
if state.cancel.is_cancelled() {
return error_response(
StatusCode::SERVICE_UNAVAILABLE,
"The proxy is shutting down",
);
}
let Ok(permit) = Arc::clone(&state.tunnel_slots).try_acquire_owned() else {
if let Some(suppressed) = REFUSED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
log::warn!(
"Proxy refused a CONNECT tunnel: {MAX_TUNNELS} already open \
({suppressed} similar refusals since the last message)"
);
}
return error_response(
StatusCode::SERVICE_UNAVAILABLE,
"Too many CONNECT tunnels are open",
);
};
let cancel = state.cancel.clone();
state.tunnels.spawn(async move {
let _permit = permit;
macro_rules! setup_step {
($what:literal, $fut:expr) => {
tokio::select! {
r = tokio::time::timeout(HANDSHAKE_TIMEOUT, $fut) => match r {
Ok(Ok(v)) => v,
Ok(Err(e)) => {
if let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
log::debug!(
concat!(
"CONNECT {} for {target} failed: {e}",
" ({n} similar since the last message)"
),
$what,
target = target,
e = e,
n = n
);
}
return;
}
Err(_) => {
if let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
log::debug!(
concat!(
"CONNECT {} for {target} did not finish within",
" {timeout:?}",
" ({n} similar since the last message)"
),
$what,
target = target,
timeout = HANDSHAKE_TIMEOUT,
n = n
);
}
return;
}
},
_ = cancel.cancelled() => {
log::debug!("CONNECT {} for {target} abandoned by shutdown", $what);
return;
}
}
};
}
let upgraded = setup_step!("upgrade", hyper::upgrade::on(req));
let mut client = hyper_util::rt::TokioIo::new(upgraded);
let mut server = setup_step!("connect", tokio::net::TcpStream::connect(target));
tokio::select! {
r = tokio::io::copy_bidirectional(&mut client, &mut server) => {
if let Err(e) = r
&& let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL)
{
log::debug!(
"CONNECT tunnel to {target} ended: {e} ({n} similar since the last message)"
);
}
}
_ = cancel.cancelled() => {
log::debug!("CONNECT tunnel to {target} closed by shutdown");
}
}
});
Response::builder()
.status(StatusCode::OK)
.header(PITCHFORK_HEADER, "1")
.body(Body::empty())
.unwrap_or_else(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())
}
fn get_request_host(req: &Request) -> Option<String> {
let authority = req
.uri()
.authority()
.map(|a| a.as_str().to_string())
.filter(|s| !s.is_empty());
authority.or_else(|| {
req.headers()
.get(HOST)
.and_then(|h| h.to_str().ok())
.map(str::to_string)
})
}
fn join_cookie_fields(headers: &mut HeaderMap) {
let fields: Vec<&[u8]> = headers
.get_all(COOKIE)
.iter()
.map(HeaderValue::as_bytes)
.collect();
if fields.len() < 2 {
return;
}
let joined = HeaderValue::from_bytes(&fields.join(b"; ".as_slice()))
.expect("valid header values joined with \"; \" form a valid header value");
headers.insert(COOKIE, joined);
}
fn inject_forwarded_headers(req: &mut Request, is_tls: bool, host_header: &str) {
let remote_addr = req
.extensions()
.get::<axum::extract::ConnectInfo<SocketAddr>>()
.map(|ci| ci.0.ip().to_string())
.unwrap_or_else(|| "127.0.0.1".to_string());
let proto = if is_tls { "https" } else { "http" };
let default_port = if is_tls { "443" } else { "80" };
let forwarded_for = remote_addr.clone();
let forwarded_proto = proto.to_string();
let forwarded_host = host_header.to_string();
let forwarded_port = host_header
.rsplit_once(':')
.map(|(_, port)| port.to_string())
.unwrap_or_else(|| default_port.to_string());
for name in [
"x-forwarded-for",
"x-forwarded-proto",
"x-forwarded-host",
"x-forwarded-port",
"forwarded",
] {
if let Ok(header_name) = axum::http::HeaderName::from_bytes(name.as_bytes()) {
req.headers_mut().remove(&header_name);
}
}
let headers = [
("x-forwarded-for", forwarded_for),
("x-forwarded-proto", forwarded_proto),
("x-forwarded-host", forwarded_host),
("x-forwarded-port", forwarded_port),
];
for (name, value) in headers {
if let Ok(v) = HeaderValue::from_str(&value) {
let header_name = axum::http::HeaderName::from_static(name);
req.headers_mut().insert(header_name, v);
}
}
}
async fn proxy_handler(State(state): State<ProxyState>, mut req: Request) -> Response {
let Some(raw_host) = get_request_host(&req) else {
return error_response(StatusCode::BAD_REQUEST, "Missing Host header");
};
if req.method() == axum::http::Method::CONNECT {
return connect_handler(&state, req, &raw_host).await;
}
let host = if raw_host.starts_with('[') {
raw_host
.split("]:")
.next()
.unwrap_or(&raw_host)
.trim_start_matches('[')
.trim_end_matches(']')
.to_string()
} else {
raw_host.split(':').next().unwrap_or(&raw_host).to_string()
};
let host = host.trim_end_matches('.').to_string();
let is_from_pitchfork = req.headers().contains_key(PROXY_HOPS_HEADER);
let hops: u64 = if is_from_pitchfork {
req.headers()
.get(PROXY_HOPS_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse().ok())
.unwrap_or(0)
} else {
0
};
if hops >= MAX_PROXY_HOPS {
return error_response(
StatusCode::LOOP_DETECTED,
&format!(
"Loop detected for '{host}': request has passed through the proxy {hops} times.\n\
This usually means a backend is proxying back through pitchfork without rewriting \n\
the Host header. If you use Vite/webpack proxy, set changeOrigin: true."
),
);
}
let local_client = is_local_client(&req);
let target_port = if let Some(subdomain) = strip_tld(&host, &state.tld) {
if subdomain.eq_ignore_ascii_case("pitchfork") {
crate::web::port()
} else {
None
}
} else {
None
};
let (target_port, activity) = if let Some(port) = target_port {
(port, None)
} else {
if resolve_tls_mode(&host, &state.tld).await.is_passthrough() {
return error_response(
StatusCode::BAD_GATEWAY,
&passthrough_unroutable_message(&host, state.is_tls),
);
}
match resolve_target(&host, &state.tld).await {
ResolveResult::Ready(port, activity) => (port, activity),
ResolveResult::Starting { slug } => {
return starting_html_response(&slug, &raw_host);
}
ResolveResult::Page {
project,
worktree,
daemons,
dir,
} => {
if !local_client {
return unknown_host_response(&host, "Not found", &[]);
}
if let Some(base) = crate::web::url()
&& let Some(resolved) = dir.clone()
&& let Some(path) = tokio::task::spawn_blocking(move || {
crate::web::routes::api::projects::page_path_for_dir(&resolved)
})
.await
.ok()
.flatten()
{
return page_redirect_response(&base, &path);
}
return page_placeholder_response(
&project,
worktree.as_deref(),
&daemons,
&state.tld,
&host_port_suffix(&raw_host),
crate::web::url().as_deref(),
);
}
ResolveResult::Unknown { heading, known } => {
return unknown_host_response(
&host,
if local_client { &heading } else { "Not found" },
if local_client { &known } else { &[] },
);
}
ResolveResult::NotFound => {
return error_response(
StatusCode::BAD_GATEWAY,
&format!(
"No daemon found for host '{host}'.\n\
A daemon is reachable once it configures a `port` and its project is \
known to pitchfork; run `pitchfork proxy status` to see the hostnames \
it serves.\n\
Expected format: <daemon>.<project>.{tld}",
tld = state.tld
),
);
}
ResolveResult::Error(msg) => {
if local_client {
return error_response(StatusCode::BAD_GATEWAY, &msg);
}
log::warn!("Refused '{host}' for a non-local client: {msg}");
return error_response(
StatusCode::BAD_GATEWAY,
&format!("'{host}' is not available."),
);
}
}
};
let path_and_query = req
.uri()
.path_and_query()
.map(|pq| pq.as_str())
.unwrap_or("/");
let forward_uri = match Uri::builder()
.scheme("http")
.authority(format!("localhost:{target_port}"))
.path_and_query(path_and_query)
.build()
{
Ok(uri) => uri,
Err(e) => {
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
&format!("Failed to build forward URI: {e}"),
);
}
};
*req.uri_mut() = forward_uri;
req.headers_mut().insert(
HOST,
HeaderValue::from_str(&format!("localhost:{target_port}"))
.unwrap_or_else(|_| HeaderValue::from_static("localhost")),
);
inject_forwarded_headers(&mut req, state.is_tls, &raw_host);
if let Ok(v) = HeaderValue::from_str(&(hops + 1).to_string()) {
req.headers_mut()
.insert(axum::http::HeaderName::from_static(PROXY_HOPS_HEADER), v);
}
let pseudo_headers: Vec<_> = req
.headers()
.keys()
.filter(|k| k.as_str().starts_with(':'))
.cloned()
.collect();
for key in pseudo_headers {
req.headers_mut().remove(&key);
}
join_cookie_fields(req.headers_mut());
*req.version_mut() = axum::http::Version::HTTP_11;
let client_upgrade = hyper::upgrade::on(&mut req);
let result = match tokio::time::timeout(
std::time::Duration::from_secs(120),
state.client.request(req),
)
.await
{
Ok(r) => r,
Err(_elapsed) => {
let msg = format!(
"Request to daemon on port {target_port} timed out after 120 s.\n\
The daemon accepted the connection but did not respond in time."
);
log::warn!("{msg}");
if let Some(ref on_error) = state.on_error {
on_error(&msg);
}
return error_response(StatusCode::GATEWAY_TIMEOUT, &msg);
}
};
match result {
Ok(mut resp) => {
let backend_upgrade = hyper::upgrade::on(&mut resp);
let (mut parts, body) = resp.into_parts();
parts.headers.insert(
axum::http::HeaderName::from_static(PITCHFORK_HEADER),
HeaderValue::from_static("1"),
);
parts.headers.remove(PROXY_HOPS_HEADER);
if state.is_tls && parts.status != StatusCode::SWITCHING_PROTOCOLS {
for h in HOP_BY_HOP_HEADERS {
if let Ok(name) = axum::http::HeaderName::from_bytes(h.as_bytes()) {
parts.headers.remove(&name);
}
}
}
if parts.status == StatusCode::SWITCHING_PROTOCOLS {
let cancel = state.cancel.clone();
state.tunnels.spawn(async move {
let _activity = activity;
let splice = async move {
if let (Ok(client_upgraded), Ok(backend_upgraded)) =
(client_upgrade.await, backend_upgrade.await)
{
let mut client_io = hyper_util::rt::TokioIo::new(client_upgraded);
let mut backend_io = hyper_util::rt::TokioIo::new(backend_upgraded);
let _ = tokio::io::copy_bidirectional(&mut client_io, &mut backend_io)
.await;
}
};
tokio::select! {
_ = splice => {}
_ = cancel.cancelled() => {}
}
});
return Response::from_parts(parts, Body::empty());
}
Response::from_parts(parts, Body::new(GuardedBody::new(body, activity)))
}
Err(e) => {
let msg = format!(
"Failed to connect to daemon on port {target_port}: {e}\n\
The daemon may have stopped or is not yet ready."
);
if let Some(ref on_error) = state.on_error {
on_error(&msg);
} else {
log::warn!("{msg}");
}
error_response(StatusCode::BAD_GATEWAY, &msg)
}
}
}
fn passthrough_unroutable_message(host: &str, is_tls: bool) -> String {
if is_tls {
format!(
"'{host}' uses proxy_tls = \"passthrough\", which routes on the host name \
in the TLS ClientHello.\n\
This connection's TLS handshake named a different host, so it was \
terminated here and the request inside it cannot be spliced to the \
daemon.\n\
Make the connection itself name '{host}' rather than overriding the Host \
header of a connection opened to something else."
)
} else {
format!(
"'{host}' uses proxy_tls = \"passthrough\", but the proxy is serving \
plain HTTP.\n\
Passthrough splices a TLS stream to the daemon, so it requires \
settings.proxy.https = true.\n\
Enable HTTPS on the proxy, or set proxy_tls = \"terminate\" on the daemon."
)
}
}
async fn resolve_target(host: &str, tld: &str) -> ResolveResult {
let deadline = auto_start_deadline();
let ctx = match resolve_route_context(host, tld).await {
Ok(ctx) => ctx,
Err(result) => return result,
};
let daemons = {
let state_file = SUPERVISOR.state_file.lock().await;
state_file.daemons.clone()
};
let daemon_name = &ctx.cached.daemon_name;
let running_matches: Vec<(&DaemonId, &crate::daemon::Daemon)> = daemons
.iter()
.filter(|(id, d)| {
id.name() == daemon_name
&& d.status.is_running()
&& match &ctx.expected_namespace {
Some(ns) => id.namespace() == ns,
None => true,
}
})
.collect();
match running_matches.as_slice() {
[] => {
try_auto_start(
&ctx.cached.slug,
&ctx.cached,
ctx.worktree_dir.as_deref(),
ctx.expected_namespace.as_deref(),
&ctx.route,
deadline,
)
.await
}
[(id, d), ..] => match select_daemon_port(&ctx.route, d) {
Some(port) => match begin_running(id, deadline).await {
Running::Yes(activity) => ResolveResult::Ready(port, Some(activity)),
Running::Stopping => ResolveResult::Starting {
slug: ctx.cached.slug.clone(),
},
Running::No => {
try_auto_start(
&ctx.cached.slug,
&ctx.cached,
ctx.worktree_dir.as_deref(),
ctx.expected_namespace.as_deref(),
&ctx.route,
deadline,
)
.await
}
},
None => ResolveResult::NotFound,
},
}
}
enum Running {
Yes(ActivityGuard),
Stopping,
No,
}
fn auto_start_deadline() -> tokio::time::Instant {
tokio::time::Instant::now() + settings().proxy_auto_start_timeout()
}
async fn wait_out_idle_stops(ids: &[DaemonId], deadline: tokio::time::Instant) -> bool {
while ids.iter().any(|id| ACTIVITY.is_idle_stopping(id)) {
if tokio::time::Instant::now() >= deadline {
return false;
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
true
}
async fn begin_running(id: &DaemonId, deadline: tokio::time::Instant) -> Running {
let activity = match ACTIVITY.begin(id) {
Some(activity) => activity,
None => {
if !wait_out_idle_stops(std::slice::from_ref(id), deadline).await {
return Running::Stopping;
}
match ACTIVITY.begin(id) {
Some(activity) => activity,
None => return Running::Stopping,
}
}
};
let running = {
let state_file = SUPERVISOR.state_file.lock().await;
state_file
.daemons
.get(id)
.is_some_and(|d| d.status.is_running())
};
if running {
Running::Yes(activity)
} else {
Running::No
}
}
struct RouteContext {
cached: CachedSlugEntry,
expected_namespace: Option<String>,
worktree_dir: Option<std::path::PathBuf>,
route: ProxyTlsRoute,
}
async fn resolve_route_context(host: &str, tld: &str) -> Result<RouteContext, ResolveResult> {
let Some(subdomain) = strip_tld(host, tld) else {
return Err(ResolveResult::NotFound);
};
let cached = cached_slug_lookup(&subdomain).await.filter(|cached| {
if crate::proxy::hostname::hostname_fits(&cached.slug) {
return true;
}
crate::proxy::hostname::warn_once(&format!(
"Slug '{}' plus the configured proxy.tld is over the DNS length limit, \
so it is not routed.",
cached.slug
));
false
});
let Some(cached) = cached else {
return Err(resolve_registry_target(&subdomain).await);
};
let (expected_namespace, worktree_dir, route) = if !subdomain.eq_ignore_ascii_case(&cached.slug)
{
let prefix = strip_dot_suffix_ignore_case(&subdomain, &cached.slug);
match prefix {
Some(ref p) => match match_worktree_prefix(&cached, p) {
PrefixMatch::Worktree(wt) => {
let ns = wt.namespace.clone().or_else(|| {
log::warn!(
"Worktree '{}' has no cached namespace; \
falling back to parent slug namespace.",
wt.path.display()
);
cached.namespace.clone()
});
let route = worktree_route(&cached, &wt.sanitized_branch);
(ns, Some(wt.path.clone()), route)
}
PrefixMatch::Ambiguous => {
return Err(ResolveResult::Error(format!(
"'{host}' is ambiguous: more than one branch or workspace of '{slug}' \
sanitizes to the prefix '{p}', and host names are case-insensitive.\n\
Rename one of them so the prefixes differ by more than case, then \
reload.\n\
The supervisor log lists the colliding branches.",
slug = cached.slug,
)));
}
PrefixMatch::Unknown => (cached.namespace.clone(), None, cached.tls),
},
None => (cached.namespace.clone(), None, cached.tls),
}
} else {
(cached.namespace.clone(), None, cached.tls)
};
Ok(RouteContext {
cached,
expected_namespace,
worktree_dir,
route,
})
}
pub(crate) async fn resolve_tls_mode(host: &str, tld: &str) -> ProxyTlsMode {
let entries = get_cached_slugs().await;
let registry = get_cached_host_registry().await;
resolve_tls_mode_in(host, tld, &entries, ®istry)
}
fn resolve_tls_mode_in(
host: &str,
tld: &str,
entries: &std::collections::HashMap<String, CachedSlugEntry>,
registry: &crate::proxy::hostname::HostRegistry,
) -> ProxyTlsMode {
let Some(subdomain) = strip_tld(host, tld) else {
return ProxyTlsMode::Terminate;
};
let wildcard = settings().proxy.wildcard;
if let Some(cached) = wildcard_slug_lookup(&subdomain, entries, wildcard)
&& crate::proxy::hostname::hostname_fits(&cached.slug)
{
if !subdomain.eq_ignore_ascii_case(&cached.slug)
&& let Some(prefix) = strip_dot_suffix_ignore_case(&subdomain, &cached.slug)
&& let PrefixMatch::Worktree(wt) = match_worktree_prefix(cached, &prefix)
{
return worktree_route(cached, &wt.sanitized_branch).mode;
}
return cached.tls.mode;
}
match registry.resolve(&subdomain, wildcard) {
crate::proxy::hostname::HostTarget::Daemon { proxy_tls, .. } => {
proxy_tls.unwrap_or_default()
}
_ => ProxyTlsMode::Terminate,
}
}
fn select_daemon_port(route: &ProxyTlsRoute, daemon: &crate::daemon::Daemon) -> Option<u16> {
let Some(want) = route.port else {
let detected = daemon.active_port.filter(|&p| p != 0);
let first_declared = daemon.resolved_port.iter().copied().find(|&p| p != 0);
return if route.mode.is_passthrough() {
first_declared.or(detected)
} else {
detected.or(first_declared)
};
};
let configured = daemon
.port
.as_ref()
.map(|p| p.expect.as_slice())
.unwrap_or(&[]);
if let Some(idx) = configured.iter().position(|&p| p == want)
&& let Some(&resolved) = daemon.resolved_port.get(idx)
{
return (resolved != 0).then_some(resolved);
}
if daemon.resolved_port.contains(&want) {
return Some(want);
}
if daemon.resolved_port.is_empty() {
return Some(want);
}
crate::proxy::hostname::warn_once(&format!(
"Daemon {} has proxy_tls_port {want}, which is not among its resolved ports {:?}; \
refusing to route rather than forwarding to a port it never bound. \
Restart the daemon if its ports changed.",
daemon.id, daemon.resolved_port,
));
None
}
struct AutoStartGuard {
daemon_id: DaemonId,
}
impl Drop for AutoStartGuard {
fn drop(&mut self) {
let daemon_id = self.daemon_id.clone();
tokio::spawn(async move {
AUTO_START_IN_PROGRESS.lock().await.remove(&daemon_id);
});
}
}
static STARTUP_LOCKS: once_cell::sync::Lazy<
std::sync::Mutex<std::collections::HashMap<DaemonId, std::sync::Weak<tokio::sync::Mutex<()>>>>,
> = once_cell::sync::Lazy::new(Default::default);
async fn lock_startup_graph(mut ids: Vec<DaemonId>) -> Vec<tokio::sync::OwnedMutexGuard<()>> {
ids.sort();
ids.dedup();
let locks: Vec<Arc<tokio::sync::Mutex<()>>> = {
let mut map = STARTUP_LOCKS.lock().unwrap_or_else(|e| e.into_inner());
map.retain(|_, lock| lock.strong_count() > 0);
ids.into_iter()
.map(|id| {
let entry = map.entry(id).or_default();
entry.upgrade().unwrap_or_else(|| {
let lock = Arc::new(tokio::sync::Mutex::new(()));
*entry = Arc::downgrade(&lock);
lock
})
})
.collect()
};
let mut guards = Vec::with_capacity(locks.len());
for lock in locks {
guards.push(lock.lock_owned().await);
}
guards
}
async fn try_auto_start(
slug: &str,
cached: &CachedSlugEntry,
worktree_dir: Option<&std::path::Path>,
expected_namespace: Option<&str>,
route: &ProxyTlsRoute,
deadline: tokio::time::Instant,
) -> ResolveResult {
let s = settings();
if !s.proxy.auto_start {
return ResolveResult::NotFound;
}
let ns = expected_namespace
.map(|s| s.to_string())
.or_else(|| cached.namespace.clone())
.unwrap_or_else(|| "global".to_string());
let daemon_id = match DaemonId::try_new(&ns, &cached.daemon_name) {
Ok(id) => id,
Err(_) => return ResolveResult::NotFound,
};
{
let mut in_progress = AUTO_START_IN_PROGRESS.lock().await;
if !in_progress.insert(daemon_id.clone()) {
return ResolveResult::Starting {
slug: slug.to_string(),
};
}
}
let guard = Arc::new(AutoStartGuard {
daemon_id: daemon_id.clone(),
});
let timeout = s.proxy_auto_start_timeout();
let timed_out = || {
log::warn!("Auto-start: total timeout ({timeout:?}) exceeded for daemon {daemon_id}");
ResolveResult::Error(format!(
"Auto-start for '{daemon_id}' timed out after {timeout:?}.\n\
The daemon and its dependencies did not all become ready, with the daemon \
bound to a port, within the configured proxy_auto_start_timeout.\n\
Startup continues in the background; reload to check again.\n\
Increase the timeout or check the logs of the daemon and its dependencies \
for slow startup."
))
};
log::info!("Auto-start: starting daemon {daemon_id} for slug '{slug}'");
let start = tokio::spawn({
let guard = guard.clone();
let daemon_id = daemon_id.clone();
let config_dir = worktree_dir.unwrap_or(&cached.dir).to_path_buf();
async move {
let _guard = guard;
start_with_dependencies(&daemon_id, &config_dir).await
}
});
match tokio::time::timeout_at(deadline, start).await {
Ok(Ok(Ok(()))) => {}
Ok(Ok(Err(result))) => return result,
Ok(Err(e)) => {
log::warn!("Auto-start: start task for {daemon_id} failed: {e}");
return ResolveResult::Error(format!("Failed to start daemon '{daemon_id}': {e}"));
}
Err(_elapsed) => return timed_out(),
}
let result =
tokio::time::timeout_at(deadline, wait_for_active_port(slug, &daemon_id, route)).await;
drop(guard);
result.unwrap_or_else(|_elapsed| timed_out())
}
async fn start_with_dependencies(
daemon_id: &DaemonId,
config_dir: &std::path::Path,
) -> std::result::Result<(), ResolveResult> {
let loaded = {
let dir = config_dir.to_path_buf();
tokio::task::spawn_blocking(move || {
crate::pitchfork_toml::PitchforkToml::all_merged_all_namespaces_from(&dir)
})
.await
};
let pt = match loaded {
Ok(Ok(pt)) => pt,
Ok(Err(e)) => {
log::warn!(
"Auto-start: failed to load config from {}: {e}",
config_dir.display()
);
return Err(ResolveResult::NotFound);
}
Err(e) => {
log::warn!("Auto-start: config loading task failed for {daemon_id}: {e}");
return Err(ResolveResult::Error(format!(
"Failed to load configuration: {e}"
)));
}
};
if !pt.daemons.contains_key(daemon_id) {
log::debug!(
"Auto-start: daemon {daemon_id} not found in config at {}",
config_dir.display()
);
return Err(ResolveResult::NotFound);
}
if SUPERVISOR
.state_file
.lock()
.await
.disabled
.contains(daemon_id)
{
return Err(ResolveResult::Error(format!(
"Daemon '{daemon_id}' is disabled, so it is not started.\n\
Enable it with: pitchfork enable {daemon_id}"
)));
}
let graph: Vec<DaemonId> =
match crate::deps::resolve_dependencies(std::slice::from_ref(daemon_id), &pt.daemons) {
Ok(order) => order.levels.into_iter().flatten().collect(),
Err(e) => {
log::warn!("Auto-start: cannot resolve dependencies of {daemon_id}: {e}");
return Err(ResolveResult::Error(format!(
"Cannot start '{daemon_id}': {e}"
)));
}
};
let _activity = match ACTIVITY.begin_all(&graph) {
Some(activity) => activity,
None => {
wait_out_idle_stops(&graph, auto_start_deadline()).await;
match ACTIVITY.begin_all(&graph) {
Some(activity) => activity,
None => {
return Err(ResolveResult::Error(format!(
"A dependency of '{daemon_id}' is still being stopped for inactivity.\n\
Reload to start it again."
)));
}
}
}
};
let proxy_idle = proxy_idle_timeouts(daemon_id, &graph, &pt).await;
let _locks = lock_startup_graph(graph).await;
let ipc = match crate::ipc::client::IpcClient::connect(false).await {
Ok(ipc) => Arc::new(ipc),
Err(e) => {
log::warn!("Auto-start: failed to connect to the supervisor: {e}");
return Err(ResolveResult::Error(format!(
"Failed to start daemon '{daemon_id}': {e}"
)));
}
};
let opts = crate::ipc::batch::StartOptions {
quiet: true,
proxy_idle: Some(proxy_idle),
..Default::default()
};
let result = match ipc
.start_daemons_with_config(std::slice::from_ref(daemon_id), opts, pt)
.await
{
Ok(result) => result,
Err(e) => {
log::warn!("Auto-start: failed to start {daemon_id}: {e}");
return Err(ResolveResult::Error(format!(
"Failed to start daemon '{daemon_id}': {e}"
)));
}
};
if !result.any_failed {
return Ok(());
}
let message = match result.failed.first() {
Some((id, reason)) if id == daemon_id => {
format!("Daemon '{daemon_id}' failed to start: {reason}\nCheck its logs for errors.")
}
Some((dep, reason)) => format!(
"Daemon '{daemon_id}' was not started because its dependency '{dep}' failed: \
{reason}\n\
Check the logs of '{dep}' for errors."
),
None => format!(
"Daemon '{daemon_id}' or one of its dependencies failed to start.\n\
Check the supervisor log for errors."
),
};
log::warn!("Auto-start: {message}");
Err(ResolveResult::Error(message))
}
async fn wait_for_active_port(
slug: &str,
daemon_id: &DaemonId,
route: &ProxyTlsRoute,
) -> ResolveResult {
let poll_interval = std::time::Duration::from_millis(250);
loop {
let daemons = {
let sf = SUPERVISOR.state_file.lock().await;
sf.daemons.clone()
};
if let Some(d) = daemons.get(daemon_id) {
if d.status.is_running() {
if let Some(port) = select_daemon_port(route, d) {
return match ACTIVITY.begin(daemon_id) {
Some(activity) => {
log::info!("Auto-start: daemon {daemon_id} is ready on port {port}");
ResolveResult::Ready(port, Some(activity))
}
None => ResolveResult::Starting {
slug: slug.to_string(),
},
};
}
} else {
log::warn!(
"Auto-start: daemon {daemon_id} is no longer running (status: {})",
d.status
);
return ResolveResult::Error(format!(
"Daemon '{daemon_id}' started but exited unexpectedly.\n\
Check its logs for errors."
));
}
} else {
log::warn!("Auto-start: daemon {daemon_id} not found in state file after start");
return ResolveResult::Error(format!(
"Daemon '{daemon_id}' started but disappeared from the state file.\n\
Check its logs for errors."
));
}
tokio::time::sleep(poll_interval).await;
}
}
async fn proxy_idle_timeouts(
daemon_id: &DaemonId,
closure: &[DaemonId],
pt: &crate::pitchfork_toml::PitchforkToml,
) -> std::collections::HashMap<DaemonId, u64> {
let own = |id: &DaemonId| pt.daemons.get(id).and_then(|d| d.proxy_idle_timeout);
let target = match own(daemon_id) {
Some(configured) => configured.duration(),
None => {
let project_dir = crate::ipc::batch::resolve_config_base_dir(
pt.daemons.get(daemon_id).and_then(|d| d.path.as_deref()),
);
tokio::task::spawn_blocking(move || {
crate::settings::Settings::load_from_dir(&project_dir).proxy_idle_timeout()
})
.await
.ok()
.flatten()
}
};
let millis = |d: std::time::Duration| u64::try_from(d.as_millis()).unwrap_or(u64::MAX);
closure
.iter()
.filter_map(|id| {
let grace = if id == daemon_id {
target
} else {
own(id).map_or(target, |configured| configured.duration())
};
grace.map(|g| (id.clone(), millis(g)))
})
.collect()
}
async fn resolve_registry_target(subdomain: &str) -> ResolveResult {
let registry = get_cached_host_registry().await;
if !crate::proxy::hostname::hostname_fits(subdomain) {
return ResolveResult::Unknown {
heading: "Host name too long".to_string(),
known: registry.project_labels(),
};
}
match registry.resolve(subdomain, settings().proxy.wildcard) {
crate::proxy::hostname::HostTarget::Daemon {
ref dir,
ref namespace,
ref daemon,
proxy_tls,
proxy_tls_port,
..
} => {
let route = ProxyTlsRoute {
mode: proxy_tls.unwrap_or_default(),
port: proxy_tls_port,
};
let per_checkout = registry.shares_daemon_id(namespace, daemon);
resolve_registry_daemon(subdomain, dir, namespace, daemon, per_checkout, &route).await
}
crate::proxy::hostname::HostTarget::ProjectPage { project } => {
let entry = registry.projects.get(&project);
ResolveResult::Page {
daemons: entry.map(|p| p.primary.labels()).unwrap_or_default(),
dir: entry.map(|p| p.primary.dir.clone()),
project,
worktree: None,
}
}
crate::proxy::hostname::HostTarget::WorktreePage { project, worktree } => {
let checkout = registry
.projects
.get(&project)
.and_then(|p| p.worktrees.get(&worktree));
ResolveResult::Page {
daemons: checkout.map(|c| c.labels()).unwrap_or_default(),
dir: checkout.map(|c| c.dir.clone()),
project,
worktree: Some(worktree),
}
}
crate::proxy::hostname::HostTarget::UnknownProject { known } => ResolveResult::Unknown {
heading: "Unknown project".to_string(),
known,
},
crate::proxy::hostname::HostTarget::UnknownDaemon {
project,
worktree,
known,
} => ResolveResult::Unknown {
heading: match worktree {
Some(wt) => format!("Unknown daemon in '{wt}' of project '{project}'"),
None => format!("Unknown daemon in project '{project}'"),
},
known,
},
}
}
async fn resolve_registry_daemon(
host: &str,
dir: &std::path::Path,
namespace: &str,
daemon: &str,
per_checkout: bool,
route: &ProxyTlsRoute,
) -> ResolveResult {
let deadline = auto_start_deadline();
let daemons = {
let state_file = SUPERVISOR.state_file.lock().await;
state_file.daemons.clone()
};
let mut matches: Vec<crate::daemon::Daemon> = daemons
.iter()
.filter(|(id, d)| {
id.name() == daemon && id.namespace() == namespace && d.status.is_running()
})
.map(|(_, d)| d.clone())
.collect();
matches = sort_by_checkout(matches, dir).await;
if let Some(d) = matches.first() {
if per_checkout && !runs_in_checkout(d.clone(), dir).await {
return ResolveResult::Error(format!(
"'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
running from {}.\n\
These checkouts share the namespace '{namespace}', so pitchfork cannot run \
both copies at once.\n\
Give each checkout its own top-level `namespace`, or stop the other one first.",
dir.display(),
d.dir
.as_deref()
.map(|p| p.display().to_string())
.unwrap_or_else(|| "an unknown directory".to_string()),
));
}
let Some(port) = select_daemon_port(route, d) else {
return ResolveResult::NotFound;
};
match begin_running(&d.id, deadline).await {
Running::Yes(activity) => return ResolveResult::Ready(port, Some(activity)),
Running::Stopping => {
return ResolveResult::Starting {
slug: host.to_string(),
};
}
Running::No => {}
}
}
let cached = CachedSlugEntry {
slug: host.to_string(),
namespace: Some(namespace.to_string()),
daemon_name: daemon.to_string(),
dir: dir.to_path_buf(),
worktrees: vec![],
rejected_worktree_prefixes: std::collections::HashSet::new(),
tls: *route,
worktree_tls: std::collections::HashMap::new(),
};
let result = try_auto_start(host, &cached, None, Some(namespace), route, deadline).await;
if per_checkout && let ResolveResult::Ready(..) = result {
let started = {
let state_file = SUPERVISOR.state_file.lock().await;
state_file
.daemons
.iter()
.find(|(id, _)| id.name() == daemon && id.namespace() == namespace)
.map(|(_, d)| d.clone())
};
if let Some(d) = started
&& !runs_in_checkout(d.clone(), dir).await
{
return ResolveResult::Error(format!(
"'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
running from {}.\n\
These checkouts share the namespace '{namespace}', so pitchfork cannot run \
both copies at once.\n\
Give each checkout its own top-level `namespace`, or stop the other one first.",
dir.display(),
d.dir
.as_deref()
.map(|p| p.display().to_string())
.unwrap_or_else(|| "an unknown directory".to_string()),
));
}
}
result
}
fn is_local_client(req: &Request) -> bool {
req.extensions()
.get::<axum::extract::ConnectInfo<SocketAddr>>()
.is_some_and(|ci| ci.0.ip().is_loopback())
}
fn daemon_runs_in(daemon: &crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
daemon
.dir
.as_deref()
.is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
}
async fn runs_in_checkout(daemon: crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
let checkout = checkout.to_path_buf();
tokio::task::spawn_blocking(move || daemon_runs_in(&daemon, &checkout))
.await
.unwrap_or(false)
}
async fn sort_by_checkout(
daemons: Vec<crate::daemon::Daemon>,
checkout: &std::path::Path,
) -> Vec<crate::daemon::Daemon> {
if daemons.len() < 2 {
return daemons;
}
let dirs: Vec<Option<std::path::PathBuf>> = daemons.iter().map(|d| d.dir.clone()).collect();
let checkout = checkout.to_path_buf();
let here = tokio::task::spawn_blocking(move || {
dirs.iter()
.map(|dir| {
dir.as_deref()
.is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
})
.collect::<Vec<bool>>()
})
.await;
match here {
Ok(here) => {
let mut ordered: Vec<(bool, crate::daemon::Daemon)> =
here.into_iter().zip(daemons).collect();
ordered.sort_by_key(|(here, _)| !here);
ordered.into_iter().map(|(_, d)| d).collect()
}
Err(e) => {
log::warn!("Checkout attribution task failed: {e}");
daemons
}
}
}
fn escape_html(s: &str) -> String {
s.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
.replace('\'', "'")
}
fn html_page(status: StatusCode, title: &str, body: String) -> Response {
let html = format!(
r##"<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>{title} — pitchfork</title>
<style>
* {{ margin: 0; padding: 0; box-sizing: border-box; }}
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background: #0f1117;
color: #e1e4e8;
display: flex;
align-items: center;
justify-content: center;
min-height: 100vh;
}}
.container {{ max-width: 640px; padding: 2rem; }}
h1 {{ font-size: 1.5rem; font-weight: 600; margin-bottom: 0.75rem; }}
p {{ color: #8b949e; font-size: 0.9rem; margin-bottom: 0.75rem; }}
ul {{ list-style: none; margin: 0.5rem 0 1rem; }}
li {{ margin: 0.25rem 0; }}
code, a {{
color: #58a6ff;
font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
text-decoration: none;
}}
</style>
</head>
<body>
<div class="container">{body}</div>
</body>
</html>"##
);
Response::builder()
.status(status)
.header("content-type", "text/html; charset=utf-8")
.body(Body::from(html))
.unwrap_or_else(|_| (status, title.to_string()).into_response())
}
fn host_port_suffix(raw_host: &str) -> String {
let port = if raw_host.starts_with('[') {
raw_host.split_once("]:").map(|(_, port)| port)
} else {
raw_host.rsplit_once(':').map(|(_, port)| port)
};
port.filter(|p| p.chars().all(|c| c.is_ascii_digit()) && !p.is_empty())
.map(|p| format!(":{p}"))
.unwrap_or_default()
}
fn page_redirect_response(base: &str, path: &str) -> Response {
let target = format!("{base}{path}");
Response::builder()
.status(StatusCode::FOUND)
.header(axum::http::header::LOCATION, &target)
.header(axum::http::header::CACHE_CONTROL, "no-store")
.body(axum::body::Body::from(format!(
"This page is at {target}\n"
)))
.unwrap_or_else(|_| {
html_page(
StatusCode::INTERNAL_SERVER_ERROR,
"pitchfork",
String::new(),
)
})
}
fn page_placeholder_response(
project: &str,
worktree: Option<&str>,
daemons: &[String],
tld: &str,
port_suffix: &str,
web_url: Option<&str>,
) -> Response {
let heading = match worktree {
Some(wt) => format!("{} · {}", escape_html(project), escape_html(wt)),
None => escape_html(project),
};
let suffix = match worktree {
Some(wt) => format!(
"{}.{}.{}",
escape_html(wt),
escape_html(project),
escape_html(tld)
),
None => format!("{}.{}", escape_html(project), escape_html(tld)),
};
let list = if daemons.is_empty() {
"<p>No daemon in this checkout has a port configured.</p>".to_string()
} else {
let items: String = daemons
.iter()
.map(|d| {
let d = escape_html(d);
format!("<li><a href=\"//{d}.{suffix}{port_suffix}\">{d}.{suffix}</a></li>")
})
.collect();
format!("<p>Daemons here:</p><ul>{items}</ul>")
};
let page = if worktree.is_some() {
"stack"
} else {
"project"
};
let explanation = match web_url {
Some(url) => format!(
"<p>This address is reserved for the {page} page, which the web UI serves. No registered project covers this checkout, so it has no page yet: add its directory under <code>[namespaces]</code> in your user config, or run <code>pitchfork proxy add</code> from it. <a href=\"{url}/projects\">Open the project list</a>.</p>",
url = escape_html(url),
),
None => format!(
"<p>This address is reserved for the {page} page, which the web UI serves. Enable it with <code>[settings.web] auto_start = true</code> to open this address.</p>"
),
};
let body = format!("<h1>{heading}</h1>{explanation}{list}");
html_page(StatusCode::OK, "pitchfork", body)
}
fn unknown_host_response(host: &str, heading: &str, known: &[String]) -> Response {
let list = if known.is_empty() {
"<p>Nothing is registered under this name yet.</p>".to_string()
} else {
let items: String = known
.iter()
.map(|k| format!("<li><code>{}</code></li>", escape_html(k)))
.collect();
format!("<p>Known names:</p><ul>{items}</ul>")
};
let body = format!(
"<h1>{heading}</h1><p>No route for <code>{host}</code>.</p>{list}",
heading = escape_html(heading),
host = escape_html(host),
);
html_page(StatusCode::NOT_FOUND, "Not found", body)
}
fn strip_tld(host: &str, tld: &str) -> Option<String> {
strip_dot_suffix_ignore_case(host.trim_end_matches('.'), tld)
}
fn bind_error_message(port: u16, err: &std::io::Error) -> String {
if port < 1024 {
format!(
"Failed to bind proxy server to port {port}: {err}\n\
Hint: ports below 1024 require elevated privileges. Run \
`pitchfork proxy setup`, which grants the bind capability on Linux, \
or set an unprivileged proxy.port and let setup redirect {port} to it."
)
} else {
format!(
"Failed to bind proxy server to port {port}: {err}\n\
Hint: another process may already be using this port."
)
}
}
fn starting_html_response(slug: &str, raw_host: &str) -> Response {
let escaped_slug = slug
.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
.replace('\'', "'");
let escaped_host = raw_host
.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
.replace('\'', "'");
let html = format!(
r##"<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<meta http-equiv="refresh" content="2">
<title>Starting {escaped_slug}… — pitchfork</title>
<style>
* {{ margin: 0; padding: 0; box-sizing: border-box; }}
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background: #0f1117;
color: #e1e4e8;
display: flex;
align-items: center;
justify-content: center;
min-height: 100vh;
}}
.container {{
text-align: center;
max-width: 480px;
padding: 2rem;
}}
.spinner {{
width: 48px;
height: 48px;
border: 4px solid rgba(255, 255, 255, 0.1);
border-top-color: #58a6ff;
border-radius: 50%;
animation: spin 0.8s linear infinite;
margin: 0 auto 1.5rem;
}}
@keyframes spin {{
to {{ transform: rotate(360deg); }}
}}
h1 {{
font-size: 1.5rem;
font-weight: 600;
margin-bottom: 0.5rem;
}}
.slug {{
color: #58a6ff;
font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
}}
.host {{
color: #8b949e;
font-size: 0.875rem;
margin-top: 0.25rem;
}}
.hint {{
color: #8b949e;
font-size: 0.8rem;
margin-top: 1.5rem;
}}
</style>
</head>
<body>
<div class="container">
<div class="spinner"></div>
<h1>Starting <span class="slug">{escaped_slug}</span>…</h1>
<p class="host">{escaped_host}</p>
<p class="hint">This page will refresh automatically when the daemon is ready.</p>
</div>
</body>
</html>"##
);
Response::builder()
.status(StatusCode::SERVICE_UNAVAILABLE)
.header("content-type", "text/html; charset=utf-8")
.header("retry-after", "2")
.body(Body::from(html))
.unwrap_or_else(|_| (StatusCode::SERVICE_UNAVAILABLE, "Starting…").into_response())
}
async fn redirect_to_https_handler(req: Request) -> Response {
if req.headers().contains_key("upgrade") {
log::warn!("Dropping plain-HTTP WebSocket upgrade attempt — use wss:// instead of ws://");
return (
StatusCode::BAD_REQUEST,
"WebSocket over plain HTTP is not supported on the HTTPS port. Use wss:// instead.",
)
.into_response();
}
let raw_host = get_request_host(&req);
let Some(raw_host) = raw_host else {
return (StatusCode::BAD_REQUEST, "Missing Host header").into_response();
};
let hostname = if raw_host.starts_with('[') {
raw_host
.split_once("]:")
.map(|(host, _)| host)
.unwrap_or(&raw_host)
.trim_start_matches('[')
.trim_end_matches(']')
} else {
let mut parts = raw_host.rsplitn(2, ':');
let last = parts.next().unwrap_or(&raw_host);
parts.next().unwrap_or(last)
};
let path = req
.uri()
.path_and_query()
.map(|pq| pq.as_str())
.unwrap_or("/");
let https_port = match u16::try_from(settings().proxy.port).ok().filter(|&p| p > 0) {
Some(443) | None => String::new(),
Some(port) => format!(":{port}"),
};
let host_for_url = if raw_host.starts_with('[') {
format!("[{hostname}]")
} else {
hostname.to_string()
};
let location = format!("https://{host_for_url}{https_port}{path}");
(
StatusCode::FOUND,
[
(axum::http::header::LOCATION, location),
(
axum::http::HeaderName::from_static(PITCHFORK_HEADER),
"1".to_string(),
),
],
)
.into_response()
}
fn error_response(status: StatusCode, message: &str) -> Response {
(
status,
[(
axum::http::HeaderName::from_static(PITCHFORK_HEADER),
HeaderValue::from_static("1"),
)],
message.to_string(),
)
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn a_request_forwards_after_an_idle_stop_is_called_off() {
let id = DaemonId::new("calledoff", "api");
SUPERVISOR.state_file.lock().await.daemons.insert(
id.clone(),
crate::daemon::Daemon {
id: id.clone(),
status: crate::daemon_status::DaemonStatus::Running,
..Default::default()
},
);
assert!(ACTIVITY.claim_idle_stop(&id, std::time::Duration::ZERO));
let releaser = {
let id = id.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
ACTIVITY.release_idle_stop(&id);
})
};
assert!(matches!(
begin_running(&id, auto_start_deadline()).await,
Running::Yes(_)
));
releaser.await.unwrap();
SUPERVISOR.state_file.lock().await.daemons.remove(&id);
}
#[tokio::test]
async fn a_request_waits_out_an_idle_stop() {
let id = DaemonId::new("waitproj", "api");
assert!(ACTIVITY.claim_idle_stop(&id, std::time::Duration::ZERO));
let releaser = {
let id = id.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
ACTIVITY.release_idle_stop(&id);
})
};
assert!(matches!(
begin_running(&id, auto_start_deadline()).await,
Running::No
));
releaser.await.unwrap();
assert!(ACTIVITY.begin(&id).is_some());
}
#[tokio::test]
async fn a_get_only_route_never_reaches_the_fallback() {
async fn routed() -> &'static str {
"routed"
}
async fn fell_through() -> &'static str {
"fell-through"
}
async fn post_to(app: Router) -> String {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let _ = axum::serve(listener, app).await;
});
let mut sock = tokio::net::TcpStream::connect(addr).await.unwrap();
sock.write_all(
b"POST /proxy.pac HTTP/1.1\r\nHost: example\r\nConnection: close\r\n\r\n",
)
.await
.unwrap();
let mut raw = Vec::new();
sock.read_to_end(&mut raw).await.unwrap();
server.abort();
let text = String::from_utf8_lossy(&raw).into_owned();
let status = text.split_whitespace().nth(1).unwrap_or("").to_string();
if status == "405" {
return status;
}
text.rsplit("\r\n").next().unwrap_or("").to_string()
}
let get_only = Router::new()
.route("/proxy.pac", axum::routing::get(routed))
.fallback(fell_through);
assert_eq!(
post_to(get_only).await,
"405",
"axum reached the fallback on a method mismatch, so `any` is unnecessary"
);
let any_method = Router::new()
.route("/proxy.pac", axum::routing::any(routed))
.fallback(fell_through);
assert_eq!(
post_to(any_method).await,
"routed",
"`any` did not deliver the POST to the handler"
);
}
#[cfg(feature = "proxy-tls")]
#[test]
fn a_ca_pair_is_checked_not_just_found() {
let dir = tempfile::tempdir().unwrap();
let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
generate_ca(&cert, &key).unwrap();
assert_eq!(ca_pair_problem(&cert, &key), None);
let (other_cert, other_key) = (dir.path().join("b.pem"), dir.path().join("b-key.pem"));
generate_ca(&other_cert, &other_key).unwrap();
assert!(ca_pair_problem(&cert, &other_key).is_some());
std::fs::write(&cert, "-----BEGIN CERTIFICATE-----\nAAAA\n").unwrap();
assert!(ca_pair_problem(&cert, &key).is_some());
}
#[cfg(feature = "proxy-tls")]
#[test]
fn a_new_ca_clears_leaves_signed_by_the_old_one() {
let dir = tempfile::tempdir().unwrap();
let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
let host_certs = host_certs_dir_for(&cert);
std::fs::create_dir_all(&host_certs).unwrap();
std::fs::write(host_certs.join("api.localhost.pem"), "old leaf").unwrap();
std::fs::write(host_certs.join("notes.txt"), "keep me").unwrap();
assert!(!ensure_ca(&cert, &key, || true).unwrap());
assert!(host_certs.join("api.localhost.pem").exists());
assert!(ensure_ca(&cert, &key, || false).unwrap());
assert!(!host_certs.join("api.localhost.pem").exists());
assert!(host_certs.join("notes.txt").exists());
}
#[cfg(feature = "proxy-tls")]
#[test]
fn leaves_are_cached_per_ca_so_a_replaced_ca_is_never_served() {
let dir = tempfile::tempdir().unwrap();
let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
let _ = rustls::crypto::ring::default_provider().install_default();
ensure_ca(&cert, &key, || false).unwrap();
let first = SniCertResolver::new(&cert, &key, "localhost".into()).unwrap();
assert!(first.get_or_create_checked("api.localhost").is_some());
let old_dir = first.host_certs_dir.clone();
assert!(std::fs::read_dir(&old_dir).unwrap().count() > 0);
ensure_ca(&cert, &key, || false).unwrap();
std::fs::create_dir_all(&old_dir).unwrap();
std::fs::write(old_dir.join("late.localhost.pem"), "old leaf").unwrap();
let second = SniCertResolver::new(&cert, &key, "localhost".into()).unwrap();
assert_ne!(second.host_certs_dir, old_dir);
assert!(!old_dir.exists());
}
#[cfg(feature = "proxy-tls")]
#[test]
fn cert_cache_file_names_do_not_collide() {
assert_ne!(
cert_cache_file_stem("a_b.localhost"),
cert_cache_file_stem("a.b.localhost")
);
assert_eq!(cert_cache_file_stem("api.localhost"), "api.localhost");
assert_eq!(cert_cache_file_stem("a_b.localhost"), "a%5Fb.localhost");
assert_eq!(
cert_cache_file_stem("*.proj.localhost"),
"%2A.proj.localhost"
);
assert!(!cert_cache_file_stem("../x").contains('/'));
}
#[test]
fn connect_port_reads_the_authority() {
assert_eq!(connect_port("api.localhost:443"), Some(443));
assert_eq!(connect_port("api.localhost:80"), Some(80));
assert_eq!(connect_port("[::1]:443"), Some(443));
assert_eq!(connect_port("api.localhost"), None);
assert_eq!(connect_port("::1"), None);
}
#[test]
fn a_half_configured_certificate_pair_is_refused() {
assert!(tls_pair_problem("", "/k.pem").is_some());
assert!(tls_pair_problem("/c.pem", "").is_some());
assert!(tls_pair_problem("", "").is_none());
assert!(tls_pair_problem("/c.pem", "/k.pem").is_none());
}
#[cfg(feature = "proxy-tls")]
fn test_resolver(tld: &str) -> (SniCertResolver, tempfile::TempDir) {
let dir = tempfile::tempdir().unwrap();
let cert = dir.path().join("ca.pem");
let key = dir.path().join("ca-key.pem");
generate_ca(&cert, &key).unwrap();
let _ = rustls::crypto::ring::default_provider().install_default();
(
SniCertResolver::new(&cert, &key, tld.to_string()).unwrap(),
dir,
)
}
#[cfg(feature = "proxy-tls")]
fn sans_for(resolver: &SniCertResolver, domain: &str) -> Vec<String> {
let ck = resolver.get_or_create(domain).expect("a certificate");
let (_, cert) = x509_parser::parse_x509_certificate(&ck.cert[0]).unwrap();
cert.subject_alternative_name()
.unwrap()
.map(|ext| {
ext.value
.general_names
.iter()
.filter_map(|n| match n {
x509_parser::extensions::GeneralName::DNSName(d) => Some(d.to_string()),
_ => None,
})
.collect()
})
.unwrap_or_default()
}
#[cfg(feature = "proxy-tls")]
#[test]
fn the_ca_refuses_to_sign_for_a_name_outside_the_tld() {
use rustls::server::ResolvesServerCert;
let (resolver, _dir) = test_resolver("localhost");
let cache = resolver.host_certs_dir.clone();
for foreign in [
"login.microsoftonline.com",
"example.com",
"notlocalhost",
"localhost.evil.com",
] {
assert!(
resolver.get_or_create_checked(foreign).is_none(),
"expected {foreign:?} to be refused"
);
}
let cached: Vec<String> = std::fs::read_dir(&cache)
.map(|rd| {
rd.filter_map(|e| e.ok())
.map(|e| e.file_name().to_string_lossy().into_owned())
.collect()
})
.unwrap_or_default();
assert!(
cached.is_empty(),
"unexpected cached certificates: {cached:?}"
);
for ours in ["localhost", "api.localhost", "core.wt.proj.localhost"] {
assert!(
resolver.get_or_create_checked(ours).is_some(),
"expected {ours:?} to be issued"
);
}
let _ = &resolver as &dyn ResolvesServerCert;
}
#[cfg(feature = "proxy-tls")]
#[test]
fn the_certificate_cache_is_bounded_and_evicts_the_oldest() {
let (resolver, _dir) = test_resolver("localhost");
let host_certs = resolver.host_certs_dir.clone();
for i in 0..MAX_HOST_CERTS + 8 {
assert!(
resolver
.get_or_create_checked(&format!("h{i}.localhost"))
.is_some()
);
}
let cached = resolver.cache.lock().unwrap();
assert_eq!(cached.by_domain.len(), MAX_HOST_CERTS);
assert_eq!(cached.order.len(), MAX_HOST_CERTS);
assert!(cached.get("h0.localhost").is_none());
assert!(
cached
.get(&format!("h{}.localhost", MAX_HOST_CERTS + 7))
.is_some()
);
drop(cached);
let on_disk = std::fs::read_dir(&host_certs).unwrap().count();
assert!(
on_disk <= MAX_HOST_CERTS,
"{on_disk} files cached, expected at most {MAX_HOST_CERTS}"
);
}
#[cfg(feature = "proxy-tls")]
#[test]
fn the_disk_cache_is_pruned_at_startup() {
let dir = tempfile::tempdir().unwrap();
let host_certs = dir.path().join("host-certs");
std::fs::create_dir_all(&host_certs).unwrap();
for i in 0..MAX_HOST_CERTS + 20 {
std::fs::write(host_certs.join(format!("old{i}.pem")), "stale").unwrap();
}
std::fs::write(host_certs.join("notes.txt"), "keep me").unwrap();
prune_host_certs(&host_certs);
let pems = std::fs::read_dir(&host_certs)
.unwrap()
.filter_map(|e| e.ok())
.filter(|e| e.path().extension().is_some_and(|x| x == "pem"))
.count();
assert_eq!(pems, MAX_HOST_CERTS);
assert!(host_certs.join("notes.txt").exists());
let small = dir.path().join("small");
std::fs::create_dir_all(&small).unwrap();
std::fs::write(small.join("a.pem"), "x").unwrap();
prune_host_certs(&small);
assert!(small.join("a.pem").exists());
}
#[cfg(feature = "proxy-tls")]
#[test]
fn a_minted_certificate_never_wildcards_the_whole_tld() {
let (resolver, _dir) = test_resolver("localhost");
let sans = sans_for(&resolver, "api.localhost");
assert!(sans.contains(&"api.localhost".to_string()));
assert!(
!sans.iter().any(|s| s.starts_with('*')),
"unexpected wildcard in {sans:?}"
);
let sans = sans_for(&resolver, "core.wt.proj.localhost");
assert!(sans.contains(&"*.wt.proj.localhost".to_string()));
let (resolver, _dir) = test_resolver("dev.internal");
let sans = sans_for(&resolver, "api.dev.internal");
assert!(
!sans.iter().any(|s| s.starts_with('*')),
"unexpected wildcard in {sans:?}"
);
}
#[tokio::test]
async fn test_page_placeholder_explains_which_step_is_missing() {
async fn body_of(response: Response) -> String {
let (_, body) = response.into_parts();
let bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap();
String::from_utf8(bytes.to_vec()).unwrap()
}
let disabled = page_placeholder_response("shop", None, &[], "localhost", "", None);
assert_eq!(disabled.status(), StatusCode::OK);
let disabled = body_of(disabled).await;
assert!(disabled.contains("auto_start"), "{disabled}");
assert!(!disabled.contains("[namespaces]"));
let running = page_placeholder_response(
"shop",
Some("feature-a"),
&[],
"localhost",
"",
Some("http://127.0.0.1:3120"),
);
let running = body_of(running).await;
assert!(running.contains("[namespaces]"), "{running}");
assert!(running.contains("http://127.0.0.1:3120/projects"));
assert!(!running.contains("auto_start"));
}
#[test]
fn test_page_redirect_targets_the_web_ui() {
let response = page_redirect_response("http://127.0.0.1:3120", "/projects/shop/feature-a");
assert_eq!(response.status(), StatusCode::FOUND);
assert_eq!(
response
.headers()
.get(axum::http::header::LOCATION)
.unwrap(),
"http://127.0.0.1:3120/projects/shop/feature-a"
);
let project = page_redirect_response("http://127.0.0.1:3120/ps", "/projects/shop");
assert_eq!(
project.headers().get(axum::http::header::LOCATION).unwrap(),
"http://127.0.0.1:3120/ps/projects/shop"
);
}
#[tokio::test]
async fn overlapping_startup_graphs_wait_for_each_other() {
use std::time::Duration;
let id = |name: &str| DaemonId::try_new("startup-lock-test", name).unwrap();
let first = lock_startup_graph(vec![id("shared"), id("app-a")]).await;
let second = tokio::spawn(lock_startup_graph(vec![id("shared"), id("app-b")]));
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(
!second.is_finished(),
"a graph sharing a daemon must wait for the graph starting it"
);
let disjoint = tokio::time::timeout(
Duration::from_secs(1),
lock_startup_graph(vec![id("other")]),
)
.await;
assert!(disjoint.is_ok());
drop(first);
let second = tokio::time::timeout(Duration::from_secs(1), second)
.await
.expect("the waiting graph proceeds once the first releases")
.unwrap();
assert_eq!(second.len(), 2);
drop(second);
drop(disjoint);
drop(lock_startup_graph(vec![id("unrelated")]).await);
let map = STARTUP_LOCKS.lock().unwrap();
assert!(
!map.keys()
.any(|k| k.namespace() == "startup-lock-test" && k.name() != "unrelated"),
"released startup locks must be pruned"
);
}
#[test]
fn test_strip_tld() {
assert_eq!(
strip_tld("api.myproject.localhost", "localhost"),
Some("api.myproject".to_string())
);
assert_eq!(
strip_tld("API.MyProject.LOCALHOST", "localhost"),
Some("API.MyProject".to_string())
);
assert_eq!(
strip_tld("api.localhost", "LOCALHOST"),
Some("api".to_string())
);
assert_eq!(
strip_tld("api.localhost", "localhost"),
Some("api".to_string())
);
assert_eq!(strip_tld("localhost", "localhost"), None);
assert_eq!(
strip_tld("API.LocalHost", "localhost"),
Some("API".to_string())
);
assert_eq!(
strip_tld("api.localhost.", "localhost"),
Some("api".to_string())
);
assert_eq!(
strip_tld("API.MyProject.LOCALHOST.", "localhost"),
Some("API.MyProject".to_string())
);
assert_eq!(strip_tld("localhost.", "localhost"), None);
assert_eq!(
strip_tld("api.myproject.test", "test"),
Some("api.myproject".to_string())
);
assert_eq!(strip_tld("other.com", "localhost"), None);
}
fn make_entry(name: &str) -> CachedSlugEntry {
CachedSlugEntry {
slug: name.to_string(),
namespace: None,
daemon_name: name.to_string(),
dir: std::path::PathBuf::from(format!("/tmp/{name}")),
worktrees: vec![],
rejected_worktree_prefixes: std::collections::HashSet::new(),
tls: ProxyTlsRoute::default(),
worktree_tls: std::collections::HashMap::new(),
}
}
fn make_daemon(
configured: &[u16],
resolved: &[u16],
active: Option<u16>,
) -> crate::daemon::Daemon {
crate::daemon::Daemon {
id: DaemonId::try_new("proj", "api").unwrap(),
port: crate::config_types::PortConfig::from_parts(
configured.to_vec(),
crate::config_types::PortBump(0),
),
resolved_port: resolved.to_vec(),
active_port: active,
..crate::daemon::Daemon::default()
}
}
#[test]
fn test_select_daemon_port_prefers_active_port() {
let route = ProxyTlsRoute::default();
let d = make_daemon(&[8443, 9443], &[8443, 9443], Some(8443));
assert_eq!(select_daemon_port(&route, &d), Some(8443));
}
#[test]
fn test_select_daemon_port_skips_port_zero() {
for mode in [ProxyTlsMode::Passthrough, ProxyTlsMode::Terminate] {
let route = ProxyTlsRoute { mode, port: None };
let mixed = make_daemon(&[0, 8443], &[0, 8443], None);
assert_eq!(select_daemon_port(&route, &mixed), Some(8443));
let detected_placeholder = make_daemon(&[0, 8443], &[0, 8443], Some(0));
assert_eq!(
select_daemon_port(&route, &detected_placeholder),
Some(8443)
);
let unresolved = make_daemon(&[0], &[0], Some(0));
assert_eq!(select_daemon_port(&route, &unresolved), None);
}
}
#[test]
fn test_select_daemon_port_passthrough_prefers_declared_first_port() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: None,
};
let d = make_daemon(&[8443, 9080], &[8443, 9080], Some(9080));
assert_eq!(select_daemon_port(&route, &d), Some(8443));
let detected_only = make_daemon(&[], &[], Some(9080));
assert_eq!(select_daemon_port(&route, &detected_only), Some(9080));
}
#[test]
fn test_select_daemon_port_falls_back_to_first_resolved() {
let route = ProxyTlsRoute::default();
let d = make_daemon(&[8443, 9443], &[8443, 9443], None);
assert_eq!(select_daemon_port(&route, &d), Some(8443));
let none = make_daemon(&[], &[], None);
assert_eq!(select_daemon_port(&route, &none), None);
}
#[test]
fn test_select_daemon_port_honors_configured_port() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(9443),
};
let d = make_daemon(&[8443, 9443], &[8443, 9443], Some(8443));
assert_eq!(select_daemon_port(&route, &d), Some(9443));
}
#[test]
fn test_select_daemon_port_follows_auto_bump() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(9443),
};
let d = make_daemon(&[8443, 9443], &[8444, 9444], Some(8444));
assert_eq!(select_daemon_port(&route, &d), Some(9444));
}
#[test]
fn test_select_daemon_port_skips_a_configured_port_resolved_to_zero() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(9443),
};
let pending = make_daemon(&[8443, 9443], &[8443, 0], Some(8443));
assert_eq!(select_daemon_port(&route, &pending), None);
let ready = make_daemon(&[8443, 9443], &[8443, 9444], Some(8443));
assert_eq!(select_daemon_port(&route, &ready), Some(9444));
}
#[test]
fn test_select_daemon_port_refuses_a_port_the_daemon_never_bound() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(9443),
};
let stale = make_daemon(&[8443], &[8443], Some(8443));
assert_eq!(select_daemon_port(&route, &stale), None);
let bare = make_daemon(&[], &[], None);
assert_eq!(select_daemon_port(&route, &bare), Some(9443));
}
#[test]
fn test_select_daemon_port_honors_configured_port_when_terminating() {
let route = ProxyTlsRoute {
mode: ProxyTlsMode::Terminate,
port: Some(9080),
};
let d = make_daemon(&[8080, 9080], &[8080, 9080], Some(8080));
assert_eq!(select_daemon_port(&route, &d), Some(9080));
let bumped = make_daemon(&[8080, 9080], &[8081, 9081], Some(8081));
assert_eq!(select_daemon_port(&route, &bumped), Some(9081));
}
#[test]
fn test_read_proxy_tls_route_absent_without_config() {
let dir = tempfile::tempdir().unwrap();
assert_eq!(
read_proxy_tls_route(dir.path(), Some("proj"), "api").unwrap(),
None
);
std::fs::write(
dir.path().join("pitchfork.toml"),
"[daemons.other]\nrun = \"serve\"\n",
)
.unwrap();
assert_eq!(
read_proxy_tls_route(dir.path(), Some("proj"), "api").unwrap(),
None,
"a config without this daemon says nothing about it"
);
assert_eq!(read_proxy_tls_route(dir.path(), None, "api").unwrap(), None);
}
#[test]
fn test_unreadable_config_keeps_the_last_known_route() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(
dir.path().join("pitchfork.toml"),
"[daemons.api]\nrun = \"serve\"\nport = 8443\nproxy_tls = \"passthrough\"\n\n\
[daemons.broken]\nrun = \"serve\"\nproxy_tls = \"passthrough\"\n",
)
.unwrap();
let read = read_proxy_tls_route(dir.path(), Some("proj"), "api");
assert!(
read.is_err(),
"an invalid sibling makes the config unreadable"
);
let known = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(8443),
};
assert_eq!(
route_or_last_known(read, Some(known), dir.path(), "api"),
Some(known)
);
assert_eq!(
route_or_last_known(Ok(None), Some(known), dir.path(), "api"),
None
);
}
#[test]
fn test_known_route_requires_the_same_target() {
let known = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(8443),
};
let mut entry = make_entry("api");
entry.namespace = Some("proj".to_string());
entry.tls = known;
let dir = entry.dir.clone();
assert_eq!(entry.known_route(&dir, Some("proj"), "api"), Some(known));
assert_eq!(entry.known_route(&dir, Some("proj"), "web"), None);
assert_eq!(entry.known_route(&dir, Some("other"), "api"), None);
assert_eq!(
entry.known_route(std::path::Path::new("/elsewhere"), Some("proj"), "api"),
None
);
let wt = make_worktree("feature/x", "feature-x");
entry.worktrees = vec![wt.clone()];
entry.worktree_tls.insert("feature-x".to_string(), known);
assert_eq!(entry.known_worktree_route(&wt, "api"), Some(known));
assert_eq!(entry.known_worktree_route(&wt, "web"), None);
let moved = crate::proxy::worktree::WorktreeEntry {
path: std::path::PathBuf::from("/elsewhere/feature-x"),
..wt
};
assert_eq!(entry.known_worktree_route(&moved, "api"), None);
}
#[test]
fn test_worktree_route_inherits_when_unknown() {
let mut entry = make_entry("spliced");
entry.tls = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(8443),
};
entry.worktrees = vec![
make_worktree("feature/known", "feature-known"),
make_worktree("feature/unknown", "feature-unknown"),
];
entry.worktree_tls.insert(
"feature-known".to_string(),
ProxyTlsRoute {
mode: ProxyTlsMode::Terminate,
port: None,
},
);
let mut entries = std::collections::HashMap::new();
entries.insert("spliced".to_string(), entry);
let mode =
|host: &str| resolve_tls_mode_in(host, "localhost", &entries, &Default::default());
assert_eq!(
mode("feature-known.spliced.localhost"),
ProxyTlsMode::Terminate,
"an explicit worktree setting wins"
);
assert_eq!(
mode("feature-unknown.spliced.localhost"),
ProxyTlsMode::Passthrough,
"a worktree with nothing recorded inherits the slug"
);
let cached = entries.get("spliced").unwrap();
assert_eq!(
worktree_route(cached, "feature-unknown"),
ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: None,
}
);
assert_eq!(
worktree_route(cached, "feature-known"),
ProxyTlsRoute {
mode: ProxyTlsMode::Terminate,
port: None,
}
);
}
#[test]
fn test_worktree_route_lookup() {
let mut entry = make_entry("myapp");
entry.tls = ProxyTlsRoute {
mode: ProxyTlsMode::Terminate,
port: None,
};
entry.worktrees = vec![
make_worktree("feature/b", "feature-b"),
make_worktree("feature/c", "feature-c"),
];
entry.worktree_tls.insert(
"feature-b".to_string(),
ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(9443),
},
);
let lookup = |prefix: &str| match match_worktree_prefix(&entry, prefix) {
PrefixMatch::Worktree(wt) => entry
.worktree_tls
.get(&wt.sanitized_branch.to_ascii_lowercase())
.copied()
.unwrap_or(entry.tls),
_ => entry.tls,
};
assert_eq!(lookup("feature-b").mode, ProxyTlsMode::Passthrough);
assert_eq!(lookup("feature-b").port, Some(9443));
assert_eq!(lookup("feature-c").mode, ProxyTlsMode::Terminate);
assert_eq!(lookup("tenant").mode, ProxyTlsMode::Terminate);
}
#[test]
fn test_wildcard_slug_lookup_exact_match() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
let result = wildcard_slug_lookup("myapp", &entries, true);
assert!(result.is_some());
assert_eq!(result.unwrap().daemon_name, "myapp");
}
#[test]
fn test_wildcard_slug_lookup_subdomain_fallback() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
assert!(result.is_some());
assert_eq!(result.unwrap().daemon_name, "myapp");
}
#[test]
fn test_wildcard_slug_lookup_nested_fallback() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
let result = wildcard_slug_lookup("a.b.myapp", &entries, true);
assert!(result.is_some());
assert_eq!(result.unwrap().daemon_name, "myapp");
}
#[test]
fn test_wildcard_slug_lookup_no_match() {
let entries = std::collections::HashMap::new();
let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
assert!(result.is_none());
}
#[test]
fn test_wildcard_slug_lookup_disabled() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
let result = wildcard_slug_lookup("tenant.myapp", &entries, false);
assert!(result.is_none());
let result = wildcard_slug_lookup("myapp", &entries, false);
assert!(result.is_some());
}
#[test]
fn test_wildcard_slug_lookup_exact_beats_wildcard() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
let mut tenant_entry = make_entry("tenant-daemon");
tenant_entry.slug = "tenant.myapp".to_string();
entries.insert("tenant.myapp".to_string(), tenant_entry);
let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
assert!(result.is_some());
assert_eq!(result.unwrap().daemon_name, "tenant-daemon");
}
#[test]
fn test_wildcard_slug_lookup_ignores_case() {
let mut entries = std::collections::HashMap::new();
entries.insert("myapp".to_string(), make_entry("myapp"));
for host in ["MyApp", "MYAPP", "myapp"] {
let result = wildcard_slug_lookup(host, &entries, true);
assert!(result.is_some(), "exact lookup failed for {host}");
assert_eq!(result.unwrap().daemon_name, "myapp");
}
for host in ["Tenant.MyApp", "tenant.MYAPP", "A.B.MyApp"] {
let result = wildcard_slug_lookup(host, &entries, true);
assert!(result.is_some(), "wildcard lookup failed for {host}");
assert_eq!(result.unwrap().daemon_name, "myapp");
}
}
#[test]
fn test_wildcard_slug_lookup_case_insensitive_registration() {
let mut entries = std::collections::HashMap::new();
let mut entry = make_entry("upper");
entry.slug = "MyApp".to_string();
entries.insert("myapp".to_string(), entry);
for host in ["myapp", "MyApp", "tenant.MYAPP"] {
let result = wildcard_slug_lookup(host, &entries, true);
assert!(result.is_some(), "lookup failed for {host}");
assert_eq!(result.unwrap().daemon_name, "upper");
}
}
fn make_worktree(branch: &str, sanitized: &str) -> crate::proxy::worktree::WorktreeEntry {
crate::proxy::worktree::WorktreeEntry {
path: std::path::PathBuf::from(format!("/tmp/{sanitized}")),
branch: branch.to_string(),
sanitized_branch: sanitized.to_string(),
namespace: Some(sanitized.to_string()),
}
}
#[test]
fn test_reject_case_colliding_worktrees_drops_both_sides() {
let wts = vec![
make_worktree("Feature-A", "Feature-A"),
make_worktree("feature-a", "feature-a"),
make_worktree("main", "main"),
];
let (kept, rejected) = reject_case_colliding_worktrees(wts);
assert_eq!(kept.len(), 1);
assert_eq!(kept[0].sanitized_branch, "main");
assert!(rejected.contains("feature-a"));
}
#[test]
fn test_reject_case_colliding_worktrees_keeps_unambiguous() {
let wts = vec![
make_worktree("main", "main"),
make_worktree("feature/a", "feature-a"),
];
let (kept, rejected) = reject_case_colliding_worktrees(wts);
assert_eq!(kept.len(), 2);
assert!(rejected.is_empty());
}
#[test]
fn test_reject_case_colliding_worktrees_drops_sanitize_duplicates() {
let wts = vec![
make_worktree("feature/a", "feature-a"),
make_worktree("feature.a", "feature-a"),
];
let (kept, rejected) = reject_case_colliding_worktrees(wts);
assert!(kept.is_empty());
assert!(rejected.contains("feature-a"));
}
#[test]
fn test_match_worktree_prefix() {
let mut entry = make_entry("myapp");
entry.worktrees = vec![make_worktree("feature/b", "feature-b")];
entry
.rejected_worktree_prefixes
.insert("feature-a".to_string());
assert!(matches!(
match_worktree_prefix(&entry, "feature-b"),
PrefixMatch::Worktree(_)
));
assert!(matches!(
match_worktree_prefix(&entry, "Feature-B"),
PrefixMatch::Worktree(_)
));
assert!(matches!(
match_worktree_prefix(&entry, "feature-a"),
PrefixMatch::Ambiguous
));
assert!(matches!(
match_worktree_prefix(&entry, "FEATURE-A"),
PrefixMatch::Ambiguous
));
assert!(matches!(
match_worktree_prefix(&entry, "tenant"),
PrefixMatch::Unknown
));
}
#[test]
fn test_passthrough_unroutable_message() {
let no_tls = passthrough_unroutable_message("api.localhost", false);
assert!(no_tls.contains("api.localhost"), "{no_tls}");
assert!(no_tls.contains("settings.proxy.https = true"), "{no_tls}");
let mismatch = passthrough_unroutable_message("api.localhost", true);
assert!(mismatch.contains("named a different host"), "{mismatch}");
assert!(
!mismatch.contains("settings.proxy.https"),
"a host mismatch is not an HTTPS configuration problem: {mismatch}"
);
}
#[tokio::test]
async fn test_slug_snapshot_is_the_cached_table() {
let from_async = get_cached_slugs().await;
let from_sync = slug_snapshot();
assert!(
Arc::ptr_eq(&from_async, &from_sync),
"the synchronous read must see the same table routing does"
);
}
#[tokio::test]
async fn test_concurrent_refresh_returns_the_published_table() {
let (first, second) = tokio::join!(get_cached_slugs(), get_cached_slugs());
assert!(
Arc::ptr_eq(&first, &second),
"overlapping refreshes must agree on one table"
);
assert!(
Arc::ptr_eq(&first, &slug_snapshot()),
"and it must be the published one"
);
}
#[test]
fn test_resolve_tls_mode_in() {
let mut entries = std::collections::HashMap::new();
let mut spliced = make_entry("spliced");
spliced.tls = ProxyTlsRoute {
mode: ProxyTlsMode::Passthrough,
port: Some(8443),
};
spliced.worktrees = vec![
make_worktree("feature/b", "feature-b"),
make_worktree("feature/c", "feature-c"),
];
spliced.worktree_tls.insert(
"feature-b".to_string(),
ProxyTlsRoute {
mode: ProxyTlsMode::Terminate,
port: None,
},
);
entries.insert("spliced".to_string(), spliced);
entries.insert("plain".to_string(), make_entry("plain"));
let mode =
|host: &str| resolve_tls_mode_in(host, "localhost", &entries, &Default::default());
assert_eq!(mode("spliced.localhost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("SPLICED.localhost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("Spliced.LocalHost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("spliced.localhost."), ProxyTlsMode::Passthrough);
assert_eq!(
mode("FEATURE-C.Spliced.LOCALHOST."),
ProxyTlsMode::Passthrough
);
assert_eq!(mode("tenant.spliced.localhost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("feature-b.spliced.localhost"), ProxyTlsMode::Terminate);
assert_eq!(
mode("feature-c.spliced.localhost"),
ProxyTlsMode::Passthrough
);
assert_eq!(mode("plain.localhost"), ProxyTlsMode::Terminate);
assert_eq!(mode("unknown.localhost"), ProxyTlsMode::Terminate);
assert_eq!(mode("localhost"), ProxyTlsMode::Terminate);
assert_eq!(mode("spliced.example.com"), ProxyTlsMode::Terminate);
}
#[test]
fn test_resolve_tls_mode_in_uses_the_hostname_registry() {
let dir = tempfile::tempdir().unwrap();
let project = dir.path().join("autoproj");
std::fs::create_dir_all(&project).unwrap();
std::fs::write(
project.join("pitchfork.toml"),
"[daemons.secure]\nrun = \"serve\"\nport = 8443\nproxy_tls = \"passthrough\"\n\
[daemons.plain]\nrun = \"serve\"\nport = 8080\n",
)
.unwrap();
let registry =
crate::proxy::hostname::HostRegistry::from_dirs(std::slice::from_ref(&project));
let slugs = std::collections::HashMap::new();
let mode = |host: &str| resolve_tls_mode_in(host, "localhost", &slugs, ®istry);
assert_eq!(mode("secure.autoproj.localhost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("plain.autoproj.localhost"), ProxyTlsMode::Terminate);
std::fs::write(project.join("pitchfork.toml"), "this is not toml = [\n").unwrap();
assert_eq!(mode("secure.autoproj.localhost"), ProxyTlsMode::Passthrough);
assert_eq!(mode("autoproj.localhost"), ProxyTlsMode::Terminate);
assert_eq!(mode("nothing.autoproj.localhost"), ProxyTlsMode::Terminate);
assert_eq!(mode("unknown.localhost"), ProxyTlsMode::Terminate);
}
#[test]
fn test_resolve_tls_mode_in_empty_table() {
let entries = std::collections::HashMap::new();
assert_eq!(
resolve_tls_mode_in(
"spliced.localhost",
"localhost",
&entries,
&Default::default()
),
ProxyTlsMode::Terminate
);
}
#[test]
fn test_strip_dot_suffix_ignore_case() {
assert_eq!(
strip_dot_suffix_ignore_case("feature-a.myapp", "myapp"),
Some("feature-a".to_string())
);
assert_eq!(
strip_dot_suffix_ignore_case("Feature-A.MyApp", "myapp"),
Some("Feature-A".to_string())
);
assert_eq!(
strip_dot_suffix_ignore_case("feature-a.myapp", "MYAPP"),
Some("feature-a".to_string())
);
assert_eq!(strip_dot_suffix_ignore_case("xmyapp", "myapp"), None);
assert_eq!(strip_dot_suffix_ignore_case(".myapp", "myapp"), None);
assert_eq!(strip_dot_suffix_ignore_case("myapp", "myapp"), None);
assert_eq!(
strip_dot_suffix_ignore_case("feature-a.other", "myapp"),
None
);
assert_eq!(
strip_dot_suffix_ignore_case("café.myapp", "myapp"),
Some("café".to_string())
);
assert_eq!(strip_dot_suffix_ignore_case("café", "afé"), None);
}
#[cfg(feature = "proxy-tls")]
fn client_hello_wire(host: &str) -> Vec<u8> {
let mut entry = vec![0u8];
entry.extend_from_slice(&(host.len() as u16).to_be_bytes());
entry.extend_from_slice(host.as_bytes());
let mut sni = (entry.len() as u16).to_be_bytes().to_vec();
sni.extend_from_slice(&entry);
let mut ext = vec![0x00, 0x00];
ext.extend_from_slice(&(sni.len() as u16).to_be_bytes());
ext.extend_from_slice(&sni);
let mut body = vec![0x03, 0x03];
body.extend_from_slice(&[0x22; 32]);
body.push(0);
body.extend_from_slice(&[0x00, 0x02, 0x13, 0x01]);
body.extend_from_slice(&[0x01, 0x00]);
body.extend_from_slice(&(ext.len() as u16).to_be_bytes());
body.extend_from_slice(&ext);
let mut msg = vec![0x01];
let len = body.len() as u32;
msg.extend_from_slice(&[(len >> 16) as u8, (len >> 8) as u8, len as u8]);
msg.extend_from_slice(&body);
let mut record = vec![0x16, 0x03, 0x01];
record.extend_from_slice(&(msg.len() as u16).to_be_bytes());
record.extend_from_slice(&msg);
record
}
#[cfg(feature = "proxy-tls")]
async fn probe_over_socket(
writes: Vec<Vec<u8>>,
gap: std::time::Duration,
timeout: std::time::Duration,
) -> (SniProbe, Vec<u8>) {
probe_over_socket_then(writes, gap, timeout, false).await
}
#[cfg(feature = "proxy-tls")]
async fn probe_over_socket_then(
writes: Vec<Vec<u8>>,
gap: std::time::Duration,
timeout: std::time::Duration,
close_after: bool,
) -> (SniProbe, Vec<u8>) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
.await
.unwrap();
let addr = listener.local_addr().unwrap();
let total: usize = writes.iter().map(Vec::len).sum();
let client = tokio::spawn(async move {
let mut sock = TcpStream::connect(addr).await.unwrap();
for chunk in writes {
sock.write_all(&chunk).await.unwrap();
sock.flush().await.unwrap();
tokio::time::sleep(gap).await;
}
if close_after {
sock.shutdown().await.unwrap();
}
tokio::time::sleep(std::time::Duration::from_secs(2)).await;
});
let (stream, _) = listener.accept().await.unwrap();
let probe = peek_sni_host(&stream, timeout).await;
let mut replayed = vec![0u8; total];
let mut stream = stream;
let read = tokio::time::timeout(
std::time::Duration::from_secs(2),
stream.read_exact(&mut replayed),
)
.await;
let replayed = match read {
Ok(Ok(_)) => replayed,
_ => vec![],
};
client.abort();
(probe, replayed)
}
#[cfg(feature = "proxy-tls")]
#[tokio::test]
async fn test_peek_sni_host_reads_a_hello_split_across_writes() {
let wire = client_hello_wire("api.localhost");
let writes: Vec<Vec<u8>> = wire.chunks(3).map(<[u8]>::to_vec).collect();
let (probe, replayed) = probe_over_socket(
writes,
std::time::Duration::from_millis(5),
std::time::Duration::from_secs(5),
)
.await;
assert_eq!(probe, SniProbe::Host("api.localhost".to_string()));
assert_eq!(replayed, wire, "peeked bytes must still be readable");
}
#[cfg(feature = "proxy-tls")]
#[tokio::test]
async fn test_peek_sni_host_undetermined_when_a_hello_stalls() {
let wire = client_hello_wire("api.localhost");
let truncated = wire[..wire.len() / 2].to_vec();
let (probe, _) = probe_over_socket(
vec![truncated],
std::time::Duration::ZERO,
std::time::Duration::from_millis(150),
)
.await;
assert_eq!(probe, SniProbe::Undetermined);
}
#[cfg(feature = "proxy-tls")]
#[tokio::test]
async fn test_peek_sni_host_notices_a_client_that_closes_mid_hello() {
let wire = client_hello_wire("api.localhost");
let truncated = wire[..wire.len() / 2].to_vec();
let started = std::time::Instant::now();
let (probe, _) = probe_over_socket_then(
vec![truncated],
std::time::Duration::ZERO,
std::time::Duration::from_secs(5),
true,
)
.await;
assert_eq!(probe, SniProbe::Undetermined);
assert!(
started.elapsed() < std::time::Duration::from_secs(1),
"the probe must end when the client closes, not at its deadline"
);
}
#[cfg(feature = "proxy-tls")]
#[tokio::test]
async fn test_peek_sni_host_gives_up_on_a_silent_peer() {
let started = std::time::Instant::now();
let (probe, _) = probe_over_socket(
vec![],
std::time::Duration::ZERO,
std::time::Duration::from_millis(150),
)
.await;
assert_eq!(probe, SniProbe::Undetermined);
assert!(
started.elapsed() < std::time::Duration::from_secs(1),
"the probe must end at its deadline, not wait on the peer"
);
}
#[cfg(feature = "proxy-tls")]
#[tokio::test]
async fn test_peek_sni_host_reports_no_host_for_non_tls() {
let (probe, _) = probe_over_socket(
vec![b"GET / HTTP/1.1\r\n\r\n".to_vec()],
std::time::Duration::ZERO,
std::time::Duration::from_secs(5),
)
.await;
assert_eq!(probe, SniProbe::NoHost);
}
#[cfg(feature = "proxy-tls")]
#[test]
fn test_generate_ca() {
let dir = tempfile::tempdir().unwrap();
let cert_path = dir.path().join("ca.pem");
let key_path = dir.path().join("ca-key.pem");
generate_ca(&cert_path, &key_path).unwrap();
assert!(cert_path.exists(), "ca.pem should be created");
assert!(key_path.exists(), "ca-key.pem should be created");
let cert_pem = std::fs::read_to_string(&cert_path).unwrap();
let key_pem = std::fs::read_to_string(&key_path).unwrap();
assert!(cert_pem.contains("BEGIN CERTIFICATE"), "should be PEM cert");
assert!(
key_pem.contains("BEGIN") && key_pem.contains("PRIVATE KEY"),
"should be PEM key"
);
}
fn cookie_fields(headers: &HeaderMap) -> Vec<&[u8]> {
headers
.get_all(COOKIE)
.iter()
.map(HeaderValue::as_bytes)
.collect()
}
#[test]
fn test_join_cookie_fields_joins_with_semicolon_space() {
let mut headers = HeaderMap::new();
headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
headers.append(COOKIE, HeaderValue::from_static("consent=ads,stats"));
headers.append(COOKIE, HeaderValue::from_static("theme=dark"));
join_cookie_fields(&mut headers);
assert_eq!(
cookie_fields(&headers),
vec![&b"_session=abc123; consent=ads,stats; theme=dark"[..]]
);
}
#[test]
fn test_join_cookie_fields_joins_bytes_outside_ascii() {
let mut headers = HeaderMap::new();
headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
headers.append(
COOKIE,
HeaderValue::from_bytes(b"name=Jos\xc3\xa9").unwrap(),
);
join_cookie_fields(&mut headers);
assert_eq!(
cookie_fields(&headers),
vec![&b"_session=abc123; name=Jos\xc3\xa9"[..]]
);
}
#[test]
fn test_join_cookie_fields_without_cookies() {
let mut headers = HeaderMap::new();
headers.insert(HOST, HeaderValue::from_static("app.localhost"));
join_cookie_fields(&mut headers);
assert!(headers.get(COOKIE).is_none());
}
#[test]
fn test_host_port_suffix() {
assert_eq!(host_port_suffix("api.myproj.localhost:8088"), ":8088");
assert_eq!(host_port_suffix("api.myproj.localhost"), "");
assert_eq!(host_port_suffix("[::1]:8088"), ":8088");
assert_eq!(host_port_suffix("[::1]"), "");
assert_eq!(host_port_suffix("host:notaport"), "");
}
#[test]
fn test_daemon_runs_in() {
let temp = tempfile::tempdir().unwrap();
let repo = temp.path().join("my-repo");
std::fs::create_dir_all(repo.join(".git/worktrees/feature")).unwrap();
std::fs::create_dir_all(repo.join("sub")).unwrap();
let nested = repo.join(".worktrees/feature");
std::fs::create_dir_all(&nested).unwrap();
std::fs::write(
nested.join(".git"),
format!(
"gitdir: {}\n",
repo.join(".git/worktrees/feature").display()
),
)
.unwrap();
let root = |p: &std::path::Path| crate::proxy::hostname::checkout_root_of(p);
let repo_root = root(&repo);
let nested_root = root(&nested);
let mut daemon = crate::daemon::Daemon {
dir: Some(repo.join("sub")),
..Default::default()
};
assert!(daemon_runs_in(&daemon, &repo_root));
assert!(!daemon_runs_in(&daemon, &nested_root));
daemon.dir = Some(nested.clone());
assert!(daemon_runs_in(&daemon, &nested_root));
assert!(!daemon_runs_in(&daemon, &repo_root));
daemon.dir = Some(temp.path().join("elsewhere"));
assert!(!daemon_runs_in(&daemon, &repo_root));
daemon.dir = None;
assert!(!daemon_runs_in(&daemon, &repo_root));
}
#[test]
fn test_is_local_client() {
let build = |info: Option<SocketAddr>| {
let mut req = Request::new(Body::empty());
if let Some(addr) = info {
req.extensions_mut()
.insert(axum::extract::ConnectInfo(addr));
}
req
};
assert!(is_local_client(&build(Some(
"127.0.0.1:5000".parse().unwrap()
))));
assert!(is_local_client(&build(Some("[::1]:5000".parse().unwrap()))));
assert!(!is_local_client(&build(Some(
"192.168.1.42:5000".parse().unwrap()
))));
assert!(!is_local_client(&build(None)));
}
}