1use std::net::SocketAddr;
11use std::sync::Arc;
12
13use axum::Router;
14use axum::body::Body;
15use axum::extract::{Request, State};
16use axum::http::{HeaderMap, HeaderValue, StatusCode, Uri};
17use axum::response::{IntoResponse, Response};
18use hyper::header::{COOKIE, HOST};
19
20const PITCHFORK_HEADER: &str = "x-pitchfork";
22
23const PROXY_HOPS_HEADER: &str = "x-pitchfork-hops";
26
27const MAX_PROXY_HOPS: u64 = 5;
29
30const HANDSHAKE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
37
38const MAX_HOST_CERTS: usize = 256;
46
47const REFUSAL_LOG_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
49
50static REFUSED_HANDSHAKE: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
52
53#[cfg(feature = "proxy-tls")]
55static REFUSED_SNI: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
56
57pub(crate) const SHUTDOWN_DRAIN_BUDGET: std::time::Duration = std::time::Duration::from_secs(10);
62
63const MAX_TUNNELS: usize = 256;
70
71static REFUSED_TUNNEL: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
73
74static ABANDONED_HANDSHAKE: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
82
83static ABANDONED_TUNNEL: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
85
86const MAX_PENDING_HANDSHAKES: usize = 512;
93
94const HOP_BY_HOP_HEADERS: &[&str] = &[
97 "connection",
98 "keep-alive",
99 "proxy-connection",
100 "transfer-encoding",
101 "upgrade",
102];
103
104use hyper_util::client::legacy::Client;
105use hyper_util::client::legacy::connect::HttpConnector;
106use hyper_util::rt::TokioExecutor;
107use tokio::net::{TcpListener, TcpStream};
108
109use crate::daemon_id::DaemonId;
110use crate::pitchfork_toml::ProxyTlsMode;
111use crate::proxy::activity::{ACTIVITY, ActivityGuard, GuardedBody};
112use crate::settings::settings;
113use crate::supervisor::SUPERVISOR;
114
115const SLUG_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(2);
128
129#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
135pub struct ProxyTlsRoute {
136 pub mode: ProxyTlsMode,
138 pub port: Option<u16>,
141}
142
143#[derive(Clone, Debug)]
145pub struct CachedSlugEntry {
146 pub slug: String,
148 pub namespace: Option<String>,
150 pub daemon_name: String,
152 pub dir: std::path::PathBuf,
154 pub worktrees: Vec<crate::proxy::worktree::WorktreeEntry>,
156 pub rejected_worktree_prefixes: std::collections::HashSet<String>,
160 pub tls: ProxyTlsRoute,
162 pub worktree_tls: std::collections::HashMap<String, ProxyTlsRoute>,
166}
167
168impl CachedSlugEntry {
169 fn known_route(
174 &self,
175 dir: &std::path::Path,
176 namespace: Option<&str>,
177 daemon_name: &str,
178 ) -> Option<ProxyTlsRoute> {
179 (self.dir == dir
180 && self.namespace.as_deref() == namespace
181 && self.daemon_name == daemon_name)
182 .then_some(self.tls)
183 }
184
185 fn known_worktree_route(
188 &self,
189 wt: &crate::proxy::worktree::WorktreeEntry,
190 daemon_name: &str,
191 ) -> Option<ProxyTlsRoute> {
192 if self.daemon_name != daemon_name {
193 return None;
194 }
195 let branch = wt.sanitized_branch.to_ascii_lowercase();
196 self.worktrees
197 .iter()
198 .any(|known| {
199 known.sanitized_branch.eq_ignore_ascii_case(&branch)
200 && known.path == wt.path
201 && known.namespace == wt.namespace
202 })
203 .then(|| self.worktree_tls.get(&branch).copied())
204 .flatten()
205 }
206}
207
208struct SlugCache {
210 entries: Arc<std::collections::HashMap<String, CachedSlugEntry>>,
211 expires_at: std::time::Instant,
212}
213
214static SLUG_CACHE: once_cell::sync::Lazy<std::sync::RwLock<SlugCache>> =
224 once_cell::sync::Lazy::new(|| {
225 std::sync::RwLock::new(SlugCache {
226 entries: Arc::new(std::collections::HashMap::new()),
227 expires_at: std::time::Instant::now(), })
229 });
230
231fn slug_snapshot() -> Arc<std::collections::HashMap<String, CachedSlugEntry>> {
237 let guard = SLUG_CACHE
238 .read()
239 .unwrap_or_else(std::sync::PoisonError::into_inner);
240 Arc::clone(&guard.entries)
241}
242
243static SLUG_REFRESH: once_cell::sync::Lazy<tokio::sync::Mutex<()>> =
252 once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(()));
253
254fn fresh_slugs() -> Option<Arc<std::collections::HashMap<String, CachedSlugEntry>>> {
256 let cache = SLUG_CACHE
257 .read()
258 .unwrap_or_else(std::sync::PoisonError::into_inner);
259 (std::time::Instant::now() < cache.expires_at).then(|| Arc::clone(&cache.entries))
260}
261
262fn reject_case_colliding_worktrees(
269 wts: Vec<crate::proxy::worktree::WorktreeEntry>,
270) -> (
271 Vec<crate::proxy::worktree::WorktreeEntry>,
272 std::collections::HashSet<String>,
273) {
274 let collisions =
275 crate::proxy::ascii_case_collisions(wts.iter().map(|w| w.sanitized_branch.as_str()));
276 if collisions.is_empty() {
277 return (wts, collisions);
278 }
279
280 let (dropped, kept): (Vec<_>, Vec<_>) = wts
281 .into_iter()
282 .partition(|w| collisions.contains(&w.sanitized_branch.to_ascii_lowercase()));
283
284 let mut folded: Vec<&String> = collisions.iter().collect();
285 folded.sort();
286 for key in folded {
287 let mut branches: Vec<&str> = dropped
288 .iter()
289 .filter(|w| w.sanitized_branch.eq_ignore_ascii_case(key))
290 .map(|w| w.branch.as_str())
291 .collect();
292 branches.sort();
293 log::warn!(
294 "Worktree slug collision: branches [{}] all route to '{key}' under \
295 case-insensitive host matching. None of them will be routed; \
296 rename a branch to disambiguate.",
297 branches.join(", "),
298 );
299 }
300
301 (kept, collisions)
302}
303
304pub(crate) fn read_proxy_tls_route(
327 dir: &std::path::Path,
328 namespace: Option<&str>,
329 daemon_name: &str,
330) -> miette::Result<Option<ProxyTlsRoute>> {
331 let Some(id) = namespace.and_then(|ns| DaemonId::try_new(ns, daemon_name).ok()) else {
332 return Ok(None);
333 };
334 let pt = crate::pitchfork_toml::PitchforkToml::all_merged_from(dir)?;
335 Ok(pt.daemons.get(&id).map(|cfg| ProxyTlsRoute {
336 mode: cfg.proxy_tls.unwrap_or_default(),
337 port: cfg.effective_proxy_tls_port(),
338 }))
339}
340
341fn route_or_last_known(
347 read: miette::Result<Option<ProxyTlsRoute>>,
348 known: Option<ProxyTlsRoute>,
349 dir: &std::path::Path,
350 daemon_name: &str,
351) -> Option<ProxyTlsRoute> {
352 read.unwrap_or_else(|e| {
353 crate::proxy::hostname::warn_once(&format!(
354 "Proxy TLS route for daemon '{daemon_name}': could not read config in {}; \
355 keeping its last known TLS mode until the config is fixed: {e}",
356 dir.display()
357 ));
358 known
359 })
360}
361
362fn build_slug_entries(
374 previous: &std::collections::HashMap<String, CachedSlugEntry>,
375) -> std::collections::HashMap<String, CachedSlugEntry> {
376 let global_slugs = crate::pitchfork_toml::PitchforkToml::read_global_slugs();
377 let collisions = crate::proxy::ascii_case_collisions(global_slugs.keys().map(String::as_str));
378 let mut folded: Vec<&String> = collisions.iter().collect();
379 folded.sort();
380 for key in folded {
381 let mut spellings: Vec<&str> = global_slugs
382 .keys()
383 .filter(|s| s.eq_ignore_ascii_case(key))
384 .map(String::as_str)
385 .collect();
386 spellings.sort();
387 log::warn!(
388 "Slug collision: [{}] differ only by case and host names are case-insensitive. \
389 None of them will be routed; remove or rename all but one.",
390 spellings.join(", "),
391 );
392 }
393
394 let mut entries: std::collections::HashMap<String, CachedSlugEntry> =
395 std::collections::HashMap::with_capacity(global_slugs.len());
396 let worktree_enabled = crate::settings::settings().general.worktree;
397 for (slug, entry) in &global_slugs {
398 let key = slug.to_ascii_lowercase();
399 if collisions.contains(&key) {
400 continue;
401 }
402 let ns = entry.resolve_namespace();
403 let daemon_name = entry.daemon.as_deref().unwrap_or(slug).to_string();
404 let (worktrees, rejected_worktree_prefixes) = if worktree_enabled {
405 let wts = match entry.resolve_dir() {
406 Some(dir) => crate::proxy::worktree::discover_worktrees(&dir),
407 None => vec![],
408 };
409 let wts = wts
410 .into_iter()
411 .map(|mut wt| {
412 wt.namespace =
413 crate::pitchfork_toml::PitchforkToml::namespace_for_dir(&wt.path).ok();
414 wt
415 })
416 .collect();
417 reject_case_colliding_worktrees(wts)
418 } else {
419 (vec![], std::collections::HashSet::new())
420 };
421 let dir = entry.resolve_dir().unwrap_or_default();
422 let prev = previous.get(&key);
423 let tls = route_or_last_known(
427 read_proxy_tls_route(&dir, ns.as_deref(), &daemon_name),
428 prev.and_then(|p| p.known_route(&dir, ns.as_deref(), &daemon_name)),
429 &dir,
430 &daemon_name,
431 )
432 .unwrap_or_default();
433 let worktree_tls = worktrees
437 .iter()
438 .filter_map(|wt| {
439 let branch = wt.sanitized_branch.to_ascii_lowercase();
440 route_or_last_known(
441 read_proxy_tls_route(&wt.path, wt.namespace.as_deref(), &daemon_name),
442 prev.and_then(|p| p.known_worktree_route(wt, &daemon_name)),
443 &wt.path,
444 &daemon_name,
445 )
446 .map(|route| (branch, route))
447 })
448 .collect();
449 entries.insert(
450 key,
451 CachedSlugEntry {
452 slug: slug.clone(),
453 namespace: ns,
454 daemon_name,
455 dir,
456 worktrees,
457 rejected_worktree_prefixes,
458 tls,
459 worktree_tls,
460 },
461 );
462 }
463 entries
464}
465
466pub async fn get_cached_slugs() -> Arc<std::collections::HashMap<String, CachedSlugEntry>> {
476 if let Some(entries) = fresh_slugs() {
478 return entries;
479 }
480
481 let _refreshing = SLUG_REFRESH.lock().await;
483
484 if let Some(entries) = fresh_slugs() {
486 return entries;
487 }
488
489 let previous = slug_snapshot();
491 let new_entries = Arc::new(
492 tokio::task::spawn_blocking(move || build_slug_entries(&previous))
493 .await
494 .unwrap_or_else(|e| {
495 log::warn!("Failed to refresh slug cache: {e}");
496 std::collections::HashMap::new()
497 }),
498 );
499
500 let mut cache = SLUG_CACHE
501 .write()
502 .unwrap_or_else(std::sync::PoisonError::into_inner);
503 cache.entries = Arc::clone(&new_entries);
504 cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
505 new_entries
506}
507
508struct RegistryCache {
516 registry: Arc<crate::proxy::hostname::HostRegistry>,
517 expires_at: std::time::Instant,
518}
519
520static HOST_REGISTRY: once_cell::sync::Lazy<std::sync::RwLock<RegistryCache>> =
527 once_cell::sync::Lazy::new(|| {
528 std::sync::RwLock::new(RegistryCache {
529 registry: Arc::new(crate::proxy::hostname::HostRegistry::default()),
530 expires_at: std::time::Instant::now(), })
532 });
533
534static REGISTRY_REFRESH: once_cell::sync::Lazy<tokio::sync::Mutex<()>> =
537 once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(()));
538
539fn fresh_registry() -> Option<Arc<crate::proxy::hostname::HostRegistry>> {
541 let cache = HOST_REGISTRY
542 .read()
543 .unwrap_or_else(std::sync::PoisonError::into_inner);
544 (std::time::Instant::now() < cache.expires_at).then(|| Arc::clone(&cache.registry))
545}
546
547fn registry_snapshot() -> Arc<crate::proxy::hostname::HostRegistry> {
553 let cache = HOST_REGISTRY
554 .read()
555 .unwrap_or_else(std::sync::PoisonError::into_inner);
556 Arc::clone(&cache.registry)
557}
558
559pub async fn get_cached_host_registry() -> Arc<crate::proxy::hostname::HostRegistry> {
561 if let Some(registry) = fresh_registry() {
562 return registry;
563 }
564
565 let _refreshing = REGISTRY_REFRESH.lock().await;
567 if let Some(registry) = fresh_registry() {
568 return registry;
569 }
570
571 let registry = Arc::new(
572 tokio::task::spawn_blocking(crate::proxy::hostname::HostRegistry::build)
573 .await
574 .unwrap_or_else(|e| {
575 log::warn!("Failed to refresh hostname registry: {e}");
576 crate::proxy::hostname::HostRegistry::default()
577 }),
578 );
579 for err in ®istry.errors {
580 crate::proxy::hostname::warn_once(err);
581 }
582
583 let mut cache = HOST_REGISTRY
584 .write()
585 .unwrap_or_else(std::sync::PoisonError::into_inner);
586 cache.registry = Arc::clone(®istry);
587 cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
588 registry
589}
590
591fn wildcard_slug_lookup<'a>(
601 subdomain: &str,
602 entries: &'a std::collections::HashMap<String, CachedSlugEntry>,
603 wildcard: bool,
604) -> Option<&'a CachedSlugEntry> {
605 let subdomain = subdomain.to_ascii_lowercase();
606
607 entries.get(&subdomain).or_else(|| {
608 if !wildcard {
609 return None;
610 }
611 subdomain
613 .match_indices('.')
614 .map(|(i, _)| &subdomain[i + 1..])
615 .find_map(|candidate| entries.get(candidate))
616 })
617}
618
619#[derive(Debug)]
621enum PrefixMatch<'a> {
622 Worktree(&'a crate::proxy::worktree::WorktreeEntry),
624 Unknown,
627 Ambiguous,
631}
632
633fn match_worktree_prefix<'a>(cached: &'a CachedSlugEntry, prefix: &str) -> PrefixMatch<'a> {
635 if let Some(wt) = cached
636 .worktrees
637 .iter()
638 .find(|w| w.sanitized_branch.eq_ignore_ascii_case(prefix))
639 {
640 return PrefixMatch::Worktree(wt);
641 }
642 if cached
643 .rejected_worktree_prefixes
644 .contains(&prefix.to_ascii_lowercase())
645 {
646 return PrefixMatch::Ambiguous;
647 }
648 PrefixMatch::Unknown
649}
650
651fn worktree_route(cached: &CachedSlugEntry, sanitized_branch: &str) -> ProxyTlsRoute {
665 cached
666 .worktree_tls
667 .get(&sanitized_branch.to_ascii_lowercase())
668 .copied()
669 .unwrap_or(ProxyTlsRoute {
670 mode: cached.tls.mode,
671 port: None,
672 })
673}
674
675fn strip_dot_suffix_ignore_case(s: &str, suffix: &str) -> Option<String> {
679 let needle_len = suffix.len() + 1;
680 if s.len() <= needle_len {
681 return None;
682 }
683 let split = s.len() - needle_len;
684 if !s.is_char_boundary(split) {
685 return None;
686 }
687 let (head, tail) = s.split_at(split);
688 if tail.starts_with('.') && tail[1..].eq_ignore_ascii_case(suffix) {
689 Some(head.to_string())
690 } else {
691 None
692 }
693}
694
695async fn cached_slug_lookup(subdomain: &str) -> Option<CachedSlugEntry> {
702 let entries = get_cached_slugs().await;
703 wildcard_slug_lookup(subdomain, &entries, settings().proxy.wildcard).cloned()
704}
705
706static AUTO_START_IN_PROGRESS: once_cell::sync::Lazy<
713 tokio::sync::Mutex<std::collections::HashSet<DaemonId>>,
714> = once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(std::collections::HashSet::new()));
715
716enum ResolveResult {
718 Ready(u16, Option<ActivityGuard>),
724 Starting { slug: String },
726 NotFound,
728 Page {
731 project: String,
732 worktree: Option<String>,
733 daemons: Vec<String>,
734 dir: Option<std::path::PathBuf>,
738 },
739 Unknown { heading: String, known: Vec<String> },
741 Error(String),
743}
744
745type OnErrorFn = Arc<dyn Fn(&str) + Send + Sync>;
748
749#[derive(Clone)]
750struct ProxyState {
751 client: Arc<Client<HttpConnector, Body>>,
753 tld: String,
755 is_tls: bool,
757 connect_target: Option<SocketAddr>,
759 contact_ip: std::net::IpAddr,
762 cancel: tokio_util::sync::CancellationToken,
764 tunnel_slots: Arc<tokio::sync::Semaphore>,
766 tunnels: tokio_util::task::TaskTracker,
773 on_error: Option<OnErrorFn>,
775}
776
777pub async fn serve(
785 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
786 cancel: tokio_util::sync::CancellationToken,
787) -> crate::Result<()> {
788 let s = settings();
789 let lan_enabled = s.proxy.lan || !s.proxy.lan_ip.is_empty();
790
791 let effective_tld = crate::proxy::effective_tld(&s).to_string();
792
793 let Some(effective_port) = u16::try_from(s.proxy.port).ok().filter(|&p| p > 0) else {
794 let msg = format!(
795 "proxy.port {} is out of valid port range (1-65535), proxy server cannot start",
796 s.proxy.port
797 );
798 let _ = bind_tx.send(Err(msg.clone()));
799 miette::bail!("{msg}");
800 };
801
802 let mut connector = HttpConnector::new();
803 connector.set_connect_timeout(Some(std::time::Duration::from_secs(10)));
807
808 let client = Client::builder(TokioExecutor::new())
809 .pool_idle_timeout(std::time::Duration::from_secs(30))
812 .build(connector);
813
814 let bind_ip: std::net::IpAddr = if lan_enabled && s.proxy.host == "127.0.0.1" {
818 std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)
819 } else {
820 match s.proxy.host.parse() {
821 Ok(ip) => ip,
822 Err(_) => {
823 log::warn!(
824 "proxy.host {:?} is not a valid IP address — falling back to 127.0.0.1. \
825 The proxy will only be reachable on the loopback interface.",
826 s.proxy.host
827 );
828 std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
829 }
830 }
831 };
832 let addr = SocketAddr::from((bind_ip, effective_port));
833 let contact_ip = local_contact_ip(bind_ip);
838 let tunnels = tokio_util::task::TaskTracker::new();
839
840 let state = ProxyState {
841 client: Arc::new(client),
842 tld: effective_tld.clone(),
843 is_tls: s.proxy.https,
844 connect_target: Some(SocketAddr::from((contact_ip, effective_port))),
847 contact_ip,
848 cancel: cancel.clone(),
849 tunnel_slots: Arc::new(tokio::sync::Semaphore::new(MAX_TUNNELS)),
850 tunnels: tunnels.clone(),
851 on_error: None,
852 };
853
854 let plain_state = state.clone();
859 let app = Router::new()
860 .route(crate::proxy::pac::PAC_PATH, axum::routing::any(pac_handler))
866 .fallback(proxy_handler)
867 .with_state(state);
868
869 if s.proxy.https {
870 serve_https_with_http_fallback(
871 app,
872 addr,
873 &s,
874 effective_port,
875 effective_tld,
876 plain_state,
877 bind_tx,
878 cancel,
879 )
880 .await
881 } else {
882 serve_http(app, addr, effective_port, plain_state, bind_tx, cancel).await
886 }
887}
888
889macro_rules! serve_conn_until_cancelled {
896 ($io:expr, $svc:expr, $cancel:expr) => {{
897 let builder = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new());
898 let conn = builder.serve_connection_with_upgrades($io, $svc);
899 tokio::pin!(conn);
900 tokio::select! {
901 r = conn.as_mut() => r,
902 _ = $cancel.cancelled() => {
903 conn.as_mut().graceful_shutdown();
904 conn.await
905 }
906 }
907 }};
908}
909
910async fn serve_http(
912 app: Router,
913 addr: SocketAddr,
914 effective_port: u16,
915 proxy_state: ProxyState,
916 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
917 cancel: tokio_util::sync::CancellationToken,
918) -> crate::Result<()> {
919 let listener = match TcpListener::bind(addr).await {
920 Ok(l) => {
921 if settings().proxy.sync_hosts {
922 crate::proxy::hosts::sync_hosts_from_settings();
923 }
924 let _ = bind_tx.send(Ok(()));
925 l
926 }
927 Err(e) => {
928 let msg = bind_error_message(effective_port, &e);
929 let _ = bind_tx.send(Err(msg.clone()));
930 return Err(miette::miette!("{msg}"));
931 }
932 };
933
934 log::info!("Proxy server listening on http://{addr}");
935 if effective_port < 1024 {
936 log::info!(
937 "Note: port {effective_port} is a privileged port. \
938 The supervisor must be started with sudo to bind to this port."
939 );
940 }
941 let mut conn_tasks: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
948 loop {
949 while conn_tasks.try_join_next().is_some() {}
950 tokio::select! {
951 accept_result = listener.accept() => {
952 let (stream, peer_addr) = match accept_result {
953 Ok(conn) => conn,
954 Err(e) => {
955 log::warn!("Accept error (will retry): {e}");
956 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
957 continue;
958 }
959 };
960 let app = app
963 .clone()
964 .layer(axum::Extension(axum::extract::ConnectInfo(peer_addr)));
965 let cancel = cancel.clone();
966 conn_tasks.spawn(async move {
967 let io = hyper_util::rt::TokioIo::new(stream);
968 let svc = hyper_util::service::TowerToHyperService::new(app);
969 if let Err(e) = serve_conn_until_cancelled!(io, svc, cancel) {
970 log::debug!("Connection error: {e}");
971 }
972 });
973 }
974 _ = cancel.cancelled() => break,
975 }
976 }
977
978 let deadline = tokio::time::Instant::now() + SHUTDOWN_DRAIN_BUDGET;
981 if tokio::time::timeout_at(deadline, async {
982 while conn_tasks.join_next().await.is_some() {}
983 })
984 .await
985 .is_err()
986 {
987 log::debug!("Proxy connections still open after {SHUTDOWN_DRAIN_BUDGET:?}; aborting them");
988 }
989 drop(conn_tasks);
991
992 proxy_state.tunnels.close();
996 let _ = tokio::time::timeout_at(deadline, proxy_state.tunnels.wait()).await;
997 Ok(())
998}
999
1000#[cfg(feature = "proxy-tls")]
1006#[allow(clippy::too_many_arguments)]
1007async fn serve_https_with_http_fallback(
1008 app: Router,
1009 addr: SocketAddr,
1010 s: &crate::settings::Settings,
1011 effective_port: u16,
1012 effective_tld: String,
1013 plain_proxy_state: ProxyState,
1014 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
1015 cancel: tokio_util::sync::CancellationToken,
1016) -> crate::Result<()> {
1017 use rustls::ServerConfig;
1018 use tokio_rustls::TlsAcceptor;
1019
1020 let (cert_path, key_path) = resolve_tls_paths(s)?;
1021
1022 let _ = rustls::crypto::ring::default_provider().install_default();
1024
1025 let resolver: Arc<dyn rustls::server::ResolvesServerCert> = if s.proxy.tls_cert.is_empty() {
1030 if ensure_ca(&cert_path, &key_path, || {
1031 cert_path.exists() && key_path.exists()
1032 })? {
1033 log::info!("Generated local CA certificate at {}", cert_path.display());
1034 log::info!("To trust the CA in your browser, run: pitchfork proxy trust");
1035 }
1036 Arc::new(SniCertResolver::new(
1037 &cert_path,
1038 &key_path,
1039 effective_tld.clone(),
1040 )?)
1041 } else {
1042 log::info!(
1043 "Serving the configured certificate {} (no certificates are minted)",
1044 cert_path.display()
1045 );
1046 Arc::new(StaticCertResolver::new(
1047 &cert_path,
1048 &key_path,
1049 effective_tld.clone(),
1050 )?)
1051 };
1052
1053 let mut tls_config = ServerConfig::builder()
1054 .with_no_client_auth()
1055 .with_cert_resolver(resolver);
1056 tls_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
1059
1060 let acceptor = TlsAcceptor::from(Arc::new(tls_config));
1061
1062 let listener = match TcpListener::bind(addr).await {
1063 Ok(l) => {
1064 if settings().proxy.sync_hosts {
1065 crate::proxy::hosts::sync_hosts_from_settings();
1066 }
1067 let _ = bind_tx.send(Ok(()));
1068 l
1069 }
1070 Err(e) => {
1071 let msg = bind_error_message(effective_port, &e);
1072 let _ = bind_tx.send(Err(msg.clone()));
1073 return Err(miette::miette!("{msg}"));
1074 }
1075 };
1076
1077 log::info!("Proxy server listening on https://{addr} (HTTP also accepted)");
1078 if effective_port < 1024 {
1079 log::info!(
1080 "Note: port {effective_port} is a privileged port. \
1081 The supervisor must be started with sudo to bind to this port."
1082 );
1083 }
1084
1085 let redirect_app = Router::new()
1093 .route(
1094 crate::proxy::pac::PAC_PATH,
1097 axum::routing::any(plain_pac_handler),
1098 )
1099 .fallback(plain_fallback_handler)
1100 .with_state(PlainState {
1101 tld: effective_tld.clone(),
1102 port: effective_port,
1103 proxy: plain_proxy_state.clone(),
1104 });
1105
1106 let mut conn_tasks: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
1108 let handshake_slots = Arc::new(tokio::sync::Semaphore::new(MAX_PENDING_HANDSHAKES));
1109 loop {
1110 while conn_tasks.try_join_next().is_some() {}
1113
1114 tokio::select! {
1115 accept_result = listener.accept() => {
1116 let (stream, peer_addr) = match accept_result {
1117 Ok(conn) => conn,
1118 Err(e) => {
1119 log::warn!("Accept error (will retry): {e}");
1120 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1121 continue;
1122 }
1123 };
1124
1125 let Ok(handshake_permit) = Arc::clone(&handshake_slots).try_acquire_owned() else {
1133 if let Some(suppressed) = REFUSED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL) {
1136 log::warn!(
1137 "Proxy refused a connection: {MAX_PENDING_HANDSHAKES} \
1138 still negotiating \
1139 ({suppressed} similar refusals since the last message)"
1140 );
1141 }
1142 drop(stream);
1143 continue;
1144 };
1145
1146 let acceptor = acceptor.clone();
1147 let app = app
1152 .clone()
1153 .layer(axum::Extension(axum::extract::ConnectInfo(peer_addr)));
1154 let redirect_app = redirect_app.clone();
1155 let tld = effective_tld.clone();
1156 let cancel = cancel.clone();
1157
1158 conn_tasks.spawn(async move {
1159 let handshake_deadline = tokio::time::Instant::now() + HANDSHAKE_TIMEOUT;
1166 let mut peek_buf = [0u8; 1];
1167 match tokio::time::timeout_at(handshake_deadline, stream.peek(&mut peek_buf)).await {
1168 Ok(Ok(0)) | Ok(Err(_)) => return,
1169 Err(_) => {
1170 if let Some(suppressed) =
1171 ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
1172 {
1173 log::debug!(
1174 "Connection sent nothing within the handshake timeout \
1175 ({suppressed} similar since the last message)"
1176 );
1177 }
1178 return;
1179 }
1180 Ok(Ok(_)) => {}
1181 }
1182
1183 if peek_buf[0] == 0x16 {
1184 let sni_budget = SNI_PEEK_TIMEOUT.min(
1192 handshake_deadline.saturating_duration_since(tokio::time::Instant::now()),
1193 );
1194 match peek_sni_host(&stream, sni_budget).await {
1195 SniProbe::Host(host) => {
1196 let Ok(mode) = tokio::time::timeout_at(
1201 handshake_deadline,
1202 resolve_tls_mode(&host, &tld),
1203 )
1204 .await
1205 else {
1206 if let Some(suppressed) =
1207 ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
1208 {
1209 log::debug!(
1210 "Routing lookup for '{host}' did not finish within the \
1211 handshake timeout ({suppressed} similar since the \
1212 last message)"
1213 );
1214 }
1215 return;
1216 };
1217 if mode.is_passthrough() {
1218 drop(handshake_permit);
1222 tokio::select! {
1226 _ = serve_passthrough(stream, &host, &tld) => {}
1227 _ = cancel.cancelled() => {}
1228 }
1229 return;
1230 }
1231 }
1232 SniProbe::NoHost => {}
1233 SniProbe::Undetermined => {
1242 log::debug!(
1243 "Could not read the ClientHello of a TLS connection; \
1244 handing it to the TLS acceptor, which refuses to terminate \
1245 a passthrough hostname."
1246 );
1247 let _ = tokio::time::timeout_at(handshake_deadline, async {
1252 let _ = get_cached_slugs().await;
1253 let _ = get_cached_host_registry().await;
1254 })
1255 .await;
1256 }
1257 }
1258
1259 let accepted = match tokio::time::timeout_at(
1261 handshake_deadline,
1262 acceptor.accept(stream),
1263 )
1264 .await
1265 {
1266 Ok(r) => r,
1267 Err(_) => {
1268 if let Some(suppressed) =
1269 ABANDONED_HANDSHAKE.allow(REFUSAL_LOG_INTERVAL)
1270 {
1271 log::debug!(
1272 "TLS handshake did not complete in time \
1273 ({suppressed} similar since the last message)"
1274 );
1275 }
1276 return;
1277 }
1278 };
1279 drop(handshake_permit);
1282 match accepted {
1283 Ok(tls_stream) => {
1284 let io = hyper_util::rt::TokioIo::new(tls_stream);
1285 let svc = hyper_util::service::TowerToHyperService::new(app);
1286 if let Err(e) = serve_conn_until_cancelled!(io, svc, cancel) {
1287 log::debug!("Connection error: {e}");
1290 }
1291 }
1292 Err(e) => {
1293 log::debug!("TLS handshake error: {e}");
1294 }
1295 }
1296 } else {
1297 drop(handshake_permit);
1300 let io = hyper_util::rt::TokioIo::new(stream);
1301 let svc = hyper_util::service::TowerToHyperService::new(redirect_app);
1302 let _ = serve_conn_until_cancelled!(io, svc, cancel);
1303 }
1304 });
1305
1306 while conn_tasks.try_join_next().is_some() {}
1307 }
1308 _ = cancel.cancelled() => {
1309 log::info!("Proxy server shutting down (cancel signal received)");
1310 break;
1311 }
1312 }
1313 }
1314
1315 let deadline = tokio::time::Instant::now() + SHUTDOWN_DRAIN_BUDGET;
1319
1320 let _ = tokio::time::timeout_at(deadline, async {
1322 while conn_tasks.join_next().await.is_some() {}
1323 })
1324 .await;
1325
1326 plain_proxy_state.tunnels.close();
1330 let _ = tokio::time::timeout_at(deadline, plain_proxy_state.tunnels.wait()).await;
1331
1332 Ok(())
1333}
1334
1335#[cfg(feature = "proxy-tls")]
1340const SNI_PEEK_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
1341
1342#[cfg(feature = "proxy-tls")]
1347const SNI_PEEK_MAX_BYTES: usize = 16 * 1024;
1348
1349#[cfg(feature = "proxy-tls")]
1355#[derive(Debug, PartialEq, Eq)]
1356enum SniProbe {
1357 Host(String),
1359 NoHost,
1362 Undetermined,
1365}
1366
1367#[cfg(feature = "proxy-tls")]
1379async fn peek_sni_host(stream: &TcpStream, timeout: std::time::Duration) -> SniProbe {
1380 use crate::proxy::sni::{SniPeek, parse_sni};
1381
1382 const MIN_PAUSE: std::time::Duration = std::time::Duration::from_millis(10);
1383 const MAX_PAUSE: std::time::Duration = std::time::Duration::from_millis(200);
1384
1385 let deadline = tokio::time::Instant::now() + timeout;
1386 let mut buf = vec![0u8; 2048];
1387 let mut last_n = 0;
1388 let mut pause = MIN_PAUSE;
1389
1390 loop {
1391 let n = match tokio::time::timeout_at(deadline, stream.peek(&mut buf)).await {
1395 Ok(Ok(0)) => return SniProbe::NoHost,
1399 Ok(Ok(n)) => n,
1400 Ok(Err(e)) => {
1401 log::debug!("Failed to peek at a TLS connection: {e}");
1402 return SniProbe::Undetermined;
1403 }
1404 Err(_elapsed) => {
1405 log::debug!("Timed out waiting for a client that sent no ClientHello");
1406 return SniProbe::Undetermined;
1407 }
1408 };
1409 match parse_sni(&buf[..n]) {
1410 SniPeek::Found(host) => return SniProbe::Host(host),
1411 SniPeek::Absent | SniPeek::NotTls => return SniProbe::NoHost,
1412 SniPeek::Incomplete => {}
1413 }
1414
1415 if n == buf.len() && buf.len() < SNI_PEEK_MAX_BYTES {
1417 buf.resize((buf.len() * 2).min(SNI_PEEK_MAX_BYTES), 0);
1418 continue;
1419 }
1420 if n >= SNI_PEEK_MAX_BYTES {
1421 log::debug!("Giving up on SNI after {n} bytes without a complete ClientHello");
1422 return SniProbe::Undetermined;
1423 }
1424 if tokio::time::Instant::now() >= deadline {
1425 log::debug!("Timed out waiting for a complete ClientHello ({n} bytes read)");
1426 return SniProbe::Undetermined;
1427 }
1428 match stream.ready(tokio::io::Interest::READABLE).await {
1434 Ok(ready) if ready.is_read_closed() => {
1435 log::debug!("Client closed after {n} bytes of an incomplete ClientHello");
1436 return SniProbe::Undetermined;
1437 }
1438 Ok(_) => {}
1439 Err(e) => {
1440 log::debug!("Failed to poll a TLS connection: {e}");
1441 return SniProbe::Undetermined;
1442 }
1443 }
1444 pause = if n > last_n {
1447 MIN_PAUSE
1448 } else {
1449 (pause * 2).min(MAX_PAUSE)
1450 };
1451 last_n = n;
1452 tokio::time::sleep_until((tokio::time::Instant::now() + pause).min(deadline)).await;
1453 }
1454}
1455
1456#[cfg(feature = "proxy-tls")]
1469async fn serve_passthrough(mut stream: TcpStream, host: &str, tld: &str) {
1470 let (port, _activity) = match resolve_passthrough_port(host, tld).await {
1474 Ok(ready) => ready,
1475 Err(msg) => {
1476 log::warn!("TLS passthrough for '{host}' failed: {msg}");
1477 return;
1478 }
1479 };
1480
1481 let addr = SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, port));
1482 let mut backend = match connect_backend(addr).await {
1483 Ok(b) => b,
1484 Err(e) => {
1485 log::warn!("TLS passthrough for '{host}': failed to connect to {addr}: {e}");
1486 return;
1487 }
1488 };
1489
1490 log::debug!("TLS passthrough: splicing '{host}' to {addr}");
1491 if let Err(e) = tokio::io::copy_bidirectional(&mut stream, &mut backend).await {
1492 log::debug!("TLS passthrough for '{host}' ended: {e}");
1494 }
1495}
1496
1497#[cfg(feature = "proxy-tls")]
1504const PASSTHROUGH_CONNECT_GRACE: std::time::Duration = std::time::Duration::from_secs(2);
1505
1506#[cfg(feature = "proxy-tls")]
1508async fn connect_backend(addr: SocketAddr) -> std::io::Result<TcpStream> {
1509 let deadline = tokio::time::Instant::now() + PASSTHROUGH_CONNECT_GRACE;
1510 loop {
1511 match TcpStream::connect(addr).await {
1512 Ok(stream) => return Ok(stream),
1513 Err(e) => {
1514 let retryable = e.kind() == std::io::ErrorKind::ConnectionRefused;
1515 if !retryable || tokio::time::Instant::now() >= deadline {
1516 return Err(e);
1517 }
1518 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1519 }
1520 }
1521 }
1522}
1523
1524#[cfg(feature = "proxy-tls")]
1532async fn resolve_passthrough_port(
1533 host: &str,
1534 tld: &str,
1535) -> std::result::Result<(u16, Option<ActivityGuard>), String> {
1536 let budget = settings().proxy_auto_start_timeout();
1537 match tokio::time::timeout(budget, resolve_passthrough_port_inner(host, tld)).await {
1538 Ok(result) => result,
1539 Err(_elapsed) => Err(format!(
1540 "no daemon was ready for '{host}' within proxy.auto_start_timeout ({budget:?})"
1541 )),
1542 }
1543}
1544
1545#[cfg(feature = "proxy-tls")]
1548async fn resolve_passthrough_port_inner(
1549 host: &str,
1550 tld: &str,
1551) -> std::result::Result<(u16, Option<ActivityGuard>), String> {
1552 loop {
1553 match resolve_target(host, tld).await {
1554 ResolveResult::Ready(port, activity) => return Ok((port, activity)),
1555 ResolveResult::Starting { slug: _ } => {
1560 tokio::time::sleep(std::time::Duration::from_millis(250)).await;
1561 }
1562 ResolveResult::NotFound => {
1563 return Err(
1564 "no running daemon with a port matched this hostname, and it could not \
1565 be auto-started"
1566 .to_string(),
1567 );
1568 }
1569 ResolveResult::Page {
1574 project, worktree, ..
1575 } => {
1576 return Err(match worktree {
1577 Some(worktree) => {
1578 format!("'{worktree}' of project '{project}' is a stack page, not a daemon")
1579 }
1580 None => format!("'{project}' is a project page, not a daemon"),
1581 });
1582 }
1583 ResolveResult::Unknown { heading, .. } => return Err(heading),
1584 ResolveResult::Error(msg) => return Err(msg),
1585 }
1586 }
1587}
1588
1589#[cfg(not(feature = "proxy-tls"))]
1591#[allow(clippy::too_many_arguments)]
1592async fn serve_https_with_http_fallback(
1593 _app: Router,
1594 _addr: SocketAddr,
1595 _s: &crate::settings::Settings,
1596 _effective_port: u16,
1597 _effective_tld: String,
1598 _plain_proxy_state: ProxyState,
1599 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
1600 _cancel: tokio_util::sync::CancellationToken,
1601) -> crate::Result<()> {
1602 let msg = "HTTPS proxy support requires the `proxy-tls` feature.\n\
1603 Rebuild pitchfork with: cargo build --features proxy-tls"
1604 .to_string();
1605 let _ = bind_tx.send(Err(msg.clone()));
1606 miette::bail!("{msg}")
1607}
1608
1609#[cfg(feature = "proxy-tls")]
1619fn resolve_tls_paths(
1620 s: &crate::settings::Settings,
1621) -> crate::Result<(std::path::PathBuf, std::path::PathBuf)> {
1622 if let Some(problem) = tls_pair_problem(&s.proxy.tls_cert, &s.proxy.tls_key) {
1623 miette::bail!("{problem}");
1624 }
1625 let proxy_dir = crate::env::PITCHFORK_STATE_DIR.join("proxy");
1626 let resolve = |configured: &str, default: &str| {
1627 if configured.is_empty() {
1628 proxy_dir.join(default)
1629 } else {
1630 std::path::PathBuf::from(configured)
1631 }
1632 };
1633 Ok((
1634 resolve(&s.proxy.tls_cert, "ca.pem"),
1635 resolve(&s.proxy.tls_key, "ca-key.pem"),
1636 ))
1637}
1638
1639pub(crate) fn tls_pair_problem(cert: &str, key: &str) -> Option<String> {
1642 match (cert.is_empty(), key.is_empty()) {
1643 (false, true) => Some(
1644 "proxy.tls_cert is set but proxy.tls_key is empty; set both, or neither to use \
1645 the generated CA"
1646 .to_string(),
1647 ),
1648 (true, false) => Some(
1649 "proxy.tls_key is set but proxy.tls_cert is empty; set both, or neither to use \
1650 the generated CA"
1651 .to_string(),
1652 ),
1653 _ => None,
1654 }
1655}
1656
1657#[cfg(feature = "proxy-tls")]
1667pub fn ensure_ca(
1668 cert_path: &std::path::Path,
1669 key_path: &std::path::Path,
1670 usable: impl FnOnce() -> bool,
1671) -> crate::Result<bool> {
1672 let _lock = xx::fslock::get(cert_path, false)
1673 .map_err(|e| miette::miette!("Failed to lock {}: {e}", cert_path.display()))?;
1674 if usable() {
1675 return Ok(false);
1676 }
1677 generate_ca(cert_path, key_path)?;
1678 clear_host_certs(&host_certs_dir_for(cert_path), None);
1682 Ok(true)
1683}
1684
1685#[cfg(feature = "proxy-tls")]
1687fn host_certs_dir_for(ca_cert_path: &std::path::Path) -> std::path::PathBuf {
1688 ca_cert_path
1689 .parent()
1690 .unwrap_or(std::path::Path::new("."))
1691 .join("host-certs")
1692}
1693
1694#[cfg(feature = "proxy-tls")]
1702fn ca_cache_id(ca_cert_pem: &str) -> String {
1703 let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
1704 for b in ca_cert_pem.bytes() {
1705 hash ^= u64::from(b);
1706 hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
1707 }
1708 format!("ca-{hash:016x}")
1709}
1710
1711#[cfg(feature = "proxy-tls")]
1715fn clear_host_certs(root: &std::path::Path, keep: Option<&str>) {
1716 let Ok(entries) = std::fs::read_dir(root) else {
1717 return;
1718 };
1719 for entry in entries.filter_map(|e| e.ok()) {
1720 let path = entry.path();
1721 let name = entry.file_name();
1722 let result = if path.is_dir() {
1723 let is_ca_dir = name.to_str().is_some_and(|n| n.starts_with("ca-"));
1724 if !is_ca_dir || keep.is_some_and(|k| name == k) {
1725 continue;
1726 }
1727 std::fs::remove_dir_all(&path)
1728 } else if path.extension().is_some_and(|x| x == "pem") {
1729 std::fs::remove_file(&path)
1730 } else {
1731 continue;
1732 };
1733 if let Err(e) = result {
1734 log::debug!(
1735 "Could not remove stale cached certs {}: {e}",
1736 path.display()
1737 );
1738 }
1739 }
1740}
1741
1742#[cfg(feature = "proxy-tls")]
1749fn generate_ca(cert_path: &std::path::Path, key_path: &std::path::Path) -> crate::Result<()> {
1750 use rcgen::{
1751 BasicConstraints, CertificateParams, DistinguishedName, DnType, IsCa, KeyUsagePurpose,
1752 };
1753
1754 if let Some(parent) = cert_path.parent() {
1756 std::fs::create_dir_all(parent)
1757 .map_err(|e| miette::miette!("Failed to create proxy cert directory: {e}"))?;
1758 }
1759
1760 let mut params = CertificateParams::default();
1761 let mut dn = DistinguishedName::new();
1762 dn.push(DnType::CommonName, "Pitchfork Local CA");
1763 dn.push(DnType::OrganizationName, "Pitchfork");
1764 params.distinguished_name = dn;
1765 params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
1766 params.key_usages = vec![KeyUsagePurpose::KeyCertSign, KeyUsagePurpose::CrlSign];
1767
1768 let key_pair = rcgen::KeyPair::generate()
1769 .map_err(|e| miette::miette!("Failed to generate CA key pair: {e}"))?;
1770 let ca_cert = params
1771 .self_signed(&key_pair)
1772 .map_err(|e| miette::miette!("Failed to self-sign CA certificate: {e}"))?;
1773
1774 std::fs::write(cert_path, ca_cert.pem()).map_err(|e| {
1776 miette::miette!(
1777 "Failed to write CA certificate to {}: {e}",
1778 cert_path.display()
1779 )
1780 })?;
1781
1782 {
1786 #[cfg(unix)]
1787 {
1788 use std::io::Write;
1789 use std::os::unix::fs::OpenOptionsExt;
1790 std::fs::OpenOptions::new()
1791 .write(true)
1792 .create(true)
1793 .truncate(true)
1794 .mode(0o600)
1795 .open(key_path)
1796 .and_then(|mut f| f.write_all(key_pair.serialize_pem().as_bytes()))
1797 .map_err(|e| {
1798 miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
1799 })?;
1800 }
1801 #[cfg(not(unix))]
1802 {
1803 std::fs::write(key_path, key_pair.serialize_pem()).map_err(|e| {
1804 miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
1805 })?;
1806 log::debug!(
1807 "CA private key written to {} (file permissions are not restricted \
1808 on non-Unix platforms — consider restricting access manually)",
1809 key_path.display()
1810 );
1811 }
1812 }
1813
1814 Ok(())
1815}
1816
1817#[cfg(feature = "proxy-tls")]
1823struct StaticCertResolver {
1824 certified: Arc<rustls::sign::CertifiedKey>,
1825 tld: String,
1828}
1829
1830#[cfg(feature = "proxy-tls")]
1831impl std::fmt::Debug for StaticCertResolver {
1832 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1833 f.debug_struct("StaticCertResolver").finish_non_exhaustive()
1834 }
1835}
1836
1837#[cfg(feature = "proxy-tls")]
1838impl StaticCertResolver {
1839 fn new(
1840 cert_path: &std::path::Path,
1841 key_path: &std::path::Path,
1842 tld: String,
1843 ) -> crate::Result<Self> {
1844 use rustls::pki_types::CertificateDer;
1845 use rustls_pemfile::{certs, private_key};
1846
1847 let cert_pem = std::fs::read(cert_path).map_err(|e| {
1848 miette::miette!("Failed to read proxy.tls_cert {}: {e}", cert_path.display())
1849 })?;
1850 let key_pem = std::fs::read(key_path).map_err(|e| {
1851 miette::miette!("Failed to read proxy.tls_key {}: {e}", key_path.display())
1852 })?;
1853
1854 let cert_ders: Vec<CertificateDer<'static>> = certs(&mut cert_pem.as_slice())
1855 .collect::<Result<Vec<_>, _>>()
1856 .map_err(|e| miette::miette!("Failed to parse {}: {e}", cert_path.display()))?;
1857 if cert_ders.is_empty() {
1858 miette::bail!("No certificates found in {}", cert_path.display());
1859 }
1860
1861 let key_der = private_key(&mut key_pem.as_slice())
1862 .map_err(|e| miette::miette!("Failed to parse {}: {e}", key_path.display()))?
1863 .ok_or_else(|| miette::miette!("No private key found in {}", key_path.display()))?;
1864 let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
1865 .map_err(|e| miette::miette!("Failed to use the configured private key: {e}"))?;
1866
1867 let certified = rustls::sign::CertifiedKey::new(cert_ders, signing_key);
1868 certified.keys_match().map_err(|e| {
1872 miette::miette!(
1873 "proxy.tls_key {} does not match proxy.tls_cert {}: {e}",
1874 key_path.display(),
1875 cert_path.display()
1876 )
1877 })?;
1878
1879 Ok(Self {
1880 certified: Arc::new(certified),
1881 tld,
1882 })
1883 }
1884}
1885
1886#[cfg(feature = "proxy-tls")]
1887impl rustls::server::ResolvesServerCert for StaticCertResolver {
1888 fn resolve(
1889 &self,
1890 client_hello: rustls::server::ClientHello<'_>,
1891 ) -> Option<Arc<rustls::sign::CertifiedKey>> {
1892 if let Some(domain) = client_hello.server_name()
1896 && resolve_tls_mode_in(domain, &self.tld, &slug_snapshot(), ®istry_snapshot())
1897 .is_passthrough()
1898 {
1899 log::warn!(
1900 "Refusing to terminate TLS for '{domain}', which is configured for \
1901 proxy_tls = \"passthrough\": its ClientHello could not be inspected before the \
1902 handshake, so the stream could not be spliced to the daemon."
1903 );
1904 return None;
1905 }
1906 Some(Arc::clone(&self.certified))
1907 }
1908}
1909
1910#[cfg(feature = "proxy-tls")]
1930struct SniCertResolver {
1931 issuer: rcgen::Issuer<'static, rcgen::KeyPair>,
1933 tld: String,
1937 host_certs_dir: std::path::PathBuf,
1939 cache: std::sync::Mutex<CertCache>,
1942 pending: std::sync::Mutex<std::collections::HashSet<String>>,
1946 pending_cv: std::sync::Condvar,
1948}
1949
1950#[cfg(feature = "proxy-tls")]
1956fn prune_host_certs(dir: &std::path::Path) {
1957 let Ok(entries) = std::fs::read_dir(dir) else {
1958 return;
1959 };
1960 let mut files: Vec<(std::time::SystemTime, std::path::PathBuf)> = entries
1961 .filter_map(|e| e.ok())
1962 .filter(|e| e.path().extension().is_some_and(|x| x == "pem"))
1963 .filter_map(|e| {
1964 let modified = e.metadata().and_then(|m| m.modified()).ok()?;
1965 Some((modified, e.path()))
1966 })
1967 .collect();
1968 if files.len() <= MAX_HOST_CERTS {
1969 return;
1970 }
1971 files.sort_by_key(|(t, _)| *t);
1972 let excess = files.len() - MAX_HOST_CERTS;
1973 for (_, path) in files.into_iter().take(excess) {
1974 if let Err(e) = std::fs::remove_file(&path) {
1975 log::debug!("Could not prune cached cert {}: {e}", path.display());
1976 }
1977 }
1978}
1979
1980#[cfg(feature = "proxy-tls")]
1982#[derive(Default)]
1983struct CertCache {
1984 by_domain: std::collections::HashMap<String, Arc<rustls::sign::CertifiedKey>>,
1985 order: std::collections::VecDeque<String>,
1987}
1988
1989#[cfg(feature = "proxy-tls")]
1990impl CertCache {
1991 fn get(&self, domain: &str) -> Option<&Arc<rustls::sign::CertifiedKey>> {
1992 self.by_domain.get(domain)
1993 }
1994
1995 fn insert(&mut self, domain: String, key: Arc<rustls::sign::CertifiedKey>) -> Option<String> {
1999 if self.by_domain.insert(domain.clone(), key).is_none() {
2000 self.order.push_back(domain);
2001 }
2002 if self.by_domain.len() <= MAX_HOST_CERTS {
2003 return None;
2004 }
2005 let evicted = self.order.pop_front()?;
2006 self.by_domain.remove(&evicted);
2007 Some(evicted)
2008 }
2009}
2010
2011#[cfg(feature = "proxy-tls")]
2012impl std::fmt::Debug for SniCertResolver {
2013 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2014 f.debug_struct("SniCertResolver").finish_non_exhaustive()
2015 }
2016}
2017
2018#[cfg(feature = "proxy-tls")]
2019impl SniCertResolver {
2020 fn new(
2022 ca_cert_path: &std::path::Path,
2023 ca_key_path: &std::path::Path,
2024 tld: String,
2025 ) -> crate::Result<Self> {
2026 let ca_key_pem = std::fs::read_to_string(ca_key_path)
2027 .map_err(|e| miette::miette!("Failed to read CA key {}: {e}", ca_key_path.display()))?;
2028 let ca_cert_pem = std::fs::read_to_string(ca_cert_path).map_err(|e| {
2029 miette::miette!("Failed to read CA cert {}: {e}", ca_cert_path.display())
2030 })?;
2031
2032 if !ca_cert_pem.contains("BEGIN CERTIFICATE") {
2034 miette::bail!("CA cert file does not contain a valid PEM certificate");
2035 }
2036
2037 let ca_key = rcgen::KeyPair::from_pem(&ca_key_pem)
2038 .map_err(|e| miette::miette!("Failed to parse CA key: {e}"))?;
2039
2040 let issuer = rcgen::Issuer::from_ca_cert_pem(&ca_cert_pem, ca_key)
2042 .map_err(|e| miette::miette!("Failed to parse CA cert: {e}"))?;
2043
2044 let host_certs_root = host_certs_dir_for(ca_cert_path);
2047 let cache_id = ca_cache_id(&ca_cert_pem);
2048 clear_host_certs(&host_certs_root, Some(&cache_id));
2049 let host_certs_dir = host_certs_root.join(&cache_id);
2050 std::fs::create_dir_all(&host_certs_dir)
2051 .map_err(|e| miette::miette!("Failed to create host-certs dir: {e}"))?;
2052 prune_host_certs(&host_certs_dir);
2056
2057 Ok(Self {
2058 issuer,
2059 tld,
2060 host_certs_dir,
2061 cache: std::sync::Mutex::new(CertCache::default()),
2062 pending: std::sync::Mutex::new(std::collections::HashSet::new()),
2063 pending_cv: std::sync::Condvar::new(),
2064 })
2065 }
2066
2067 fn get_or_create_checked(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
2076 if !crate::proxy::owns_name(&self.tld, domain) {
2077 if let Some(suppressed) = REFUSED_SNI.allow(REFUSAL_LOG_INTERVAL) {
2080 log::warn!(
2081 "Refusing to issue a certificate for {domain:?}: \
2082 the pitchfork CA only signs names under .{} \
2083 ({suppressed} similar refusals since the last message)",
2084 self.tld
2085 );
2086 }
2087 return None;
2088 }
2089 self.get_or_create(domain)
2090 }
2091
2092 fn get_or_create(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
2115 {
2117 let cache = self.cache.lock().ok()?;
2118 if let Some(ck) = cache.get(domain) {
2119 return Some(Arc::clone(ck));
2120 }
2121 } loop {
2134 {
2135 let mut pending = self.pending.lock().ok()?;
2136 if pending.contains(domain) {
2137 pending = self.pending_cv.wait(pending).ok()?;
2139 drop(pending);
2141 } else {
2142 pending.insert(domain.to_string());
2144 break;
2145 }
2146 } {
2152 let cache = self.cache.lock().ok()?;
2153 if let Some(ck) = cache.get(domain) {
2154 return Some(Arc::clone(ck));
2155 }
2156 } } let result = self.get_or_create_inner(domain);
2160
2161 {
2167 let mut pending = match self.pending.lock() {
2168 Ok(g) => g,
2169 Err(e) => e.into_inner(),
2170 };
2171 pending.remove(domain);
2172 self.pending_cv.notify_all();
2173 }
2174
2175 result
2176 }
2177
2178 fn get_or_create_inner(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
2180 let disk_path = self.disk_path(domain);
2181
2182 if disk_path.exists() {
2184 if let Ok(ck) = self.load_from_disk(&disk_path) {
2185 let ck = Arc::new(ck);
2186 self.remember(domain, &ck);
2187 return Some(ck);
2188 }
2189 let _ = std::fs::remove_file(&disk_path);
2191 }
2192
2193 let ck = self.sign_for_domain(domain).ok()?;
2195
2196 let ck = Arc::new(ck);
2197 self.remember(domain, &ck);
2198 Some(ck)
2199 }
2200
2201 fn remember(&self, domain: &str, ck: &Arc<rustls::sign::CertifiedKey>) {
2206 let evicted = match self.cache.lock() {
2207 Ok(mut cache) => cache.insert(domain.to_string(), Arc::clone(ck)),
2208 Err(_) => return,
2209 };
2210 if let Some(evicted) = evicted {
2211 let path = self.disk_path(&evicted);
2212 if let Err(e) = std::fs::remove_file(&path)
2213 && e.kind() != std::io::ErrorKind::NotFound
2214 {
2215 log::debug!("Could not evict cached cert {}: {e}", path.display());
2216 }
2217 }
2218 }
2219
2220 fn disk_path(&self, domain: &str) -> std::path::PathBuf {
2222 self.host_certs_dir
2223 .join(format!("{}.pem", cert_cache_file_stem(domain)))
2224 }
2225
2226 fn load_from_disk(&self, path: &std::path::Path) -> crate::Result<rustls::sign::CertifiedKey> {
2231 use rustls::pki_types::CertificateDer;
2232 use rustls_pemfile::{certs, private_key};
2233
2234 let pem = std::fs::read_to_string(path)
2235 .map_err(|e| miette::miette!("Failed to read disk cert {}: {e}", path.display()))?;
2236
2237 let cert_ders: Vec<CertificateDer<'static>> = certs(&mut pem.as_bytes())
2238 .collect::<Result<Vec<_>, _>>()
2239 .map_err(|e| miette::miette!("Failed to parse certs from {}: {e}", path.display()))?;
2240
2241 if cert_ders.is_empty() {
2242 miette::bail!("No certificates found in {}", path.display());
2243 }
2244
2245 {
2247 let (_, cert) = x509_parser::parse_x509_certificate(&cert_ders[0]).map_err(|e| {
2248 miette::miette!("Failed to parse certificate from {}: {e}", path.display())
2249 })?;
2250 use chrono::Utc;
2251 let now_ts = Utc::now().timestamp();
2252 let not_after_ts = cert.validity().not_after.timestamp();
2253 if not_after_ts < now_ts {
2254 miette::bail!(
2255 "Cached certificate at {} has expired — will regenerate",
2256 path.display()
2257 );
2258 }
2259 }
2260
2261 let key_der = private_key(&mut pem.as_bytes())
2262 .map_err(|e| miette::miette!("Failed to parse key from {}: {e}", path.display()))?
2263 .ok_or_else(|| miette::miette!("No private key found in {}", path.display()))?;
2264
2265 let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
2266 .map_err(|e| miette::miette!("Failed to create signing key from disk: {e}"))?;
2267
2268 Ok(rustls::sign::CertifiedKey::new(cert_ders, signing_key))
2269 }
2270
2271 fn sign_for_domain(&self, domain: &str) -> crate::Result<rustls::sign::CertifiedKey> {
2279 use rcgen::date_time_ymd;
2280 use rcgen::{CertificateParams, DistinguishedName, DnType, SanType};
2281 use rustls::pki_types::CertificateDer;
2282 use rustls_pemfile::private_key;
2283
2284 let mut params = CertificateParams::default();
2285 let mut dn = DistinguishedName::new();
2286 dn.push(DnType::CommonName, domain);
2287 params.distinguished_name = dn;
2288
2289 {
2291 use chrono::{Datelike, Duration, Utc};
2292 let yesterday = Utc::now() - Duration::days(1);
2293 let expiry = Utc::now() + Duration::days(397);
2296 params.not_before = date_time_ymd(
2297 yesterday.year(),
2298 yesterday.month() as u8,
2299 yesterday.day() as u8,
2300 );
2301 params.not_after =
2302 date_time_ymd(expiry.year(), expiry.month() as u8, expiry.day() as u8);
2303 }
2304
2305 let mut sans =
2307 vec![SanType::DnsName(domain.to_string().try_into().map_err(
2308 |e| miette::miette!("Invalid domain name '{domain}': {e}"),
2309 )?)];
2310 if let Some(dot_pos) = domain.find('.') {
2317 let parent = &domain[dot_pos + 1..];
2318 if crate::proxy::is_strictly_under_tld(&self.tld, parent) {
2319 let wildcard = format!("*.{parent}");
2320 if let Ok(wc) = wildcard.try_into() {
2321 sans.push(SanType::DnsName(wc));
2322 }
2323 }
2324 }
2325 params.subject_alt_names = sans;
2326
2327 let leaf_key = rcgen::KeyPair::generate()
2328 .map_err(|e| miette::miette!("Failed to generate leaf key: {e}"))?;
2329 let leaf_cert = params
2330 .signed_by(&leaf_key, &self.issuer)
2331 .map_err(|e| miette::miette!("Failed to sign leaf cert for '{domain}': {e}"))?;
2332
2333 let cert_der = CertificateDer::from(leaf_cert.der().to_vec());
2335 let key_pem = leaf_key.serialize_pem();
2336 let key_der = private_key(&mut key_pem.as_bytes())
2337 .map_err(|e| miette::miette!("Failed to parse leaf key PEM: {e}"))?
2338 .ok_or_else(|| miette::miette!("No private key found in generated PEM"))?;
2339
2340 let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
2341 .map_err(|e| miette::miette!("Failed to create signing key: {e}"))?;
2342
2343 let disk_path = self.disk_path(domain);
2346 let combined_pem = format!("{}{}", leaf_cert.pem(), key_pem);
2347 {
2348 #[cfg(unix)]
2349 {
2350 use std::io::Write;
2351 use std::os::unix::fs::OpenOptionsExt;
2352 if let Err(e) = std::fs::OpenOptions::new()
2353 .write(true)
2354 .create(true)
2355 .truncate(true)
2356 .mode(0o600)
2357 .open(&disk_path)
2358 .and_then(|mut f| f.write_all(combined_pem.as_bytes()))
2359 {
2360 log::warn!(
2361 "Failed to persist cert for '{domain}' to {}: {e}",
2362 disk_path.display()
2363 );
2364 }
2365 }
2366 #[cfg(not(unix))]
2367 {
2368 if let Err(e) = std::fs::write(&disk_path, combined_pem) {
2369 log::warn!(
2370 "Failed to persist cert for '{domain}' to {}: {e}",
2371 disk_path.display()
2372 );
2373 } else {
2374 log::debug!(
2375 "Leaf cert for '{domain}' written to {} (file permissions are not \
2376 restricted on non-Unix platforms — consider restricting access manually)",
2377 disk_path.display()
2378 );
2379 }
2380 }
2381 }
2382
2383 Ok(rustls::sign::CertifiedKey::new(vec![cert_der], signing_key))
2384 }
2385}
2386
2387#[cfg(feature = "proxy-tls")]
2388impl rustls::server::ResolvesServerCert for SniCertResolver {
2389 fn resolve(
2390 &self,
2391 client_hello: rustls::server::ClientHello<'_>,
2392 ) -> Option<Arc<rustls::sign::CertifiedKey>> {
2393 let domain = client_hello.server_name()?;
2394 if resolve_tls_mode_in(domain, &self.tld, &slug_snapshot(), ®istry_snapshot())
2401 .is_passthrough()
2402 {
2403 log::warn!(
2404 "Refusing to terminate TLS for '{domain}', which is configured for \
2405 proxy_tls = \"passthrough\": its ClientHello could not be inspected before the \
2406 handshake, so the stream could not be spliced to the daemon."
2407 );
2408 return None;
2409 }
2410
2411 self.get_or_create_checked(domain)
2420 }
2421}
2422
2423#[cfg(feature = "proxy-tls")]
2429pub(crate) fn ca_pair_problem(cert: &std::path::Path, key: &std::path::Path) -> Option<String> {
2430 use rcgen::PublicKeyData;
2431
2432 let cert_pem = match std::fs::read_to_string(cert) {
2433 Ok(p) => p,
2434 Err(e) => return Some(format!("cannot read {}: {e}", cert.display())),
2435 };
2436 let key_pem = match std::fs::read_to_string(key) {
2437 Ok(p) => p,
2438 Err(e) => return Some(format!("cannot read {}: {e}", key.display())),
2439 };
2440 let key_pair = match rcgen::KeyPair::from_pem(&key_pem) {
2441 Ok(k) => k,
2442 Err(e) => return Some(format!("cannot parse {}: {e}", key.display())),
2443 };
2444 let Some(Ok(der)) = rustls_pemfile::certs(&mut cert_pem.as_bytes()).next() else {
2445 return Some(format!("no certificate in {}", cert.display()));
2446 };
2447 let parsed = match x509_parser::parse_x509_certificate(&der) {
2448 Ok((_, c)) => c,
2449 Err(e) => return Some(format!("cannot parse {}: {e}", cert.display())),
2450 };
2451 if parsed.public_key().subject_public_key.data.as_ref() != key_pair.der_bytes() {
2452 return Some(format!(
2453 "{} is not the key for {}",
2454 key.display(),
2455 cert.display()
2456 ));
2457 }
2458 if let Err(e) = rcgen::Issuer::from_ca_cert_pem(&cert_pem, key_pair) {
2459 return Some(format!("{} cannot sign certificates: {e}", cert.display()));
2460 }
2461 None
2462}
2463
2464#[cfg(feature = "proxy-tls")]
2471fn cert_cache_file_stem(domain: &str) -> String {
2472 let mut out = String::with_capacity(domain.len());
2473 for b in domain.bytes() {
2474 if b.is_ascii_alphanumeric() || b == b'-' || b == b'.' {
2475 out.push(char::from(b));
2476 } else {
2477 out.push_str(&format!("%{b:02X}"));
2478 }
2479 }
2480 out
2481}
2482
2483#[derive(Clone)]
2485struct PlainState {
2486 tld: String,
2487 port: u16,
2488 proxy: ProxyState,
2491}
2492
2493async fn plain_fallback_handler(State(state): State<PlainState>, req: Request) -> Response {
2498 if req.method() == axum::http::Method::CONNECT {
2499 let raw_host = get_request_host(&req).unwrap_or_default();
2500 return connect_handler(&state.proxy, req, &raw_host).await;
2501 }
2502 redirect_to_https_handler(req).await
2503}
2504
2505fn reject_non_read(method: &axum::http::Method) -> Option<Response> {
2512 if matches!(*method, axum::http::Method::GET | axum::http::Method::HEAD) {
2513 return None;
2514 }
2515 let mut res = error_response(
2516 StatusCode::METHOD_NOT_ALLOWED,
2517 "the proxy auto-config file is read-only\n",
2518 );
2519 res.headers_mut().insert(
2520 axum::http::header::ALLOW,
2521 HeaderValue::from_static("GET, HEAD"),
2522 );
2523 Some(res)
2524}
2525
2526fn pac_response(tld: &str, host: &str, port: u16) -> Response {
2527 match crate::proxy::pac::generate(tld, host, port) {
2528 Ok(body) => (
2529 StatusCode::OK,
2530 [(
2531 axum::http::header::CONTENT_TYPE,
2532 "application/x-ns-proxy-autoconfig",
2533 )],
2534 body,
2535 )
2536 .into_response(),
2537 Err(e) => error_response(StatusCode::INTERNAL_SERVER_ERROR, &e.to_string()),
2538 }
2539}
2540
2541async fn pac_handler(State(state): State<ProxyState>, req: Request) -> Response {
2548 let host = get_request_host(&req).unwrap_or_default();
2549 let bare = host.split(':').next().unwrap_or("");
2550 if !bare.is_empty() && crate::proxy::is_strictly_under_tld(&state.tld, bare) {
2551 return proxy_handler(State(state), req).await;
2552 }
2553 if let Some(deny) = reject_non_read(req.method()) {
2556 return deny;
2557 }
2558 let port = req
2559 .uri()
2560 .authority()
2561 .and_then(|a| a.port_u16())
2562 .or_else(|| host.rsplit(':').next().and_then(|p| p.parse().ok()))
2563 .unwrap_or(if state.is_tls { 443 } else { 80 });
2564 pac_response(&state.tld, &url_host(state.contact_ip), port)
2565}
2566
2567async fn plain_pac_handler(State(state): State<PlainState>, req: Request) -> Response {
2569 let host = get_request_host(&req).unwrap_or_default();
2573 let bare = host.split(':').next().unwrap_or("");
2574 if !bare.is_empty() && crate::proxy::is_strictly_under_tld(&state.tld, bare) {
2575 return redirect_to_https_handler(req).await;
2576 }
2577 if let Some(deny) = reject_non_read(req.method()) {
2578 return deny;
2579 }
2580 pac_response(&state.tld, &url_host(state.proxy.contact_ip), state.port)
2581}
2582
2583fn local_contact_ip(bind_ip: std::net::IpAddr) -> std::net::IpAddr {
2589 match bind_ip {
2590 std::net::IpAddr::V4(ip) if ip.is_unspecified() => {
2591 std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
2592 }
2593 std::net::IpAddr::V6(ip) if ip.is_unspecified() => {
2594 std::net::IpAddr::V6(std::net::Ipv6Addr::LOCALHOST)
2595 }
2596 ip => ip,
2597 }
2598}
2599
2600fn url_host(ip: std::net::IpAddr) -> String {
2602 match ip {
2603 std::net::IpAddr::V6(ip) => format!("[{ip}]"),
2604 std::net::IpAddr::V4(ip) => ip.to_string(),
2605 }
2606}
2607
2608fn connect_port(authority: &str) -> Option<u16> {
2610 let (host, port) = authority.rsplit_once(':')?;
2611 if host.contains(':') && !host.ends_with(']') {
2613 return None;
2614 }
2615 port.parse().ok()
2616}
2617
2618async fn connect_handler(state: &ProxyState, req: Request, raw_host: &str) -> Response {
2633 let authority = req
2634 .uri()
2635 .authority()
2636 .map(|a| a.as_str().to_string())
2637 .unwrap_or_else(|| raw_host.to_string());
2638 let host = authority
2639 .rsplit_once(':')
2640 .map(|(h, _)| h)
2641 .unwrap_or(&authority)
2642 .trim_start_matches('[')
2643 .trim_end_matches(']')
2644 .to_string();
2645
2646 if !crate::proxy::owns_name(&state.tld, &host) {
2649 return error_response(
2650 StatusCode::FORBIDDEN,
2651 &format!(
2652 "pitchfork only tunnels CONNECT for names under .{} — refusing {host}",
2653 state.tld
2654 ),
2655 );
2656 }
2657
2658 if !state.is_tls && connect_port(&authority) == Some(443) {
2663 return error_response(
2664 StatusCode::BAD_GATEWAY,
2665 &format!(
2666 "proxy.https is false, so pitchfork cannot serve https://{host}; \
2667 use http:// or enable proxy.https"
2668 ),
2669 );
2670 }
2671
2672 let Some(target) = state.connect_target else {
2673 return error_response(
2674 StatusCode::SERVICE_UNAVAILABLE,
2675 "The proxy listener address is unknown, so CONNECT cannot be tunnelled",
2676 );
2677 };
2678
2679 if state.cancel.is_cancelled() {
2682 return error_response(
2683 StatusCode::SERVICE_UNAVAILABLE,
2684 "The proxy is shutting down",
2685 );
2686 }
2687
2688 let Ok(permit) = Arc::clone(&state.tunnel_slots).try_acquire_owned() else {
2691 if let Some(suppressed) = REFUSED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
2692 log::warn!(
2693 "Proxy refused a CONNECT tunnel: {MAX_TUNNELS} already open \
2694 ({suppressed} similar refusals since the last message)"
2695 );
2696 }
2697 return error_response(
2698 StatusCode::SERVICE_UNAVAILABLE,
2699 "Too many CONNECT tunnels are open",
2700 );
2701 };
2702 let cancel = state.cancel.clone();
2703
2704 state.tunnels.spawn(async move {
2710 let _permit = permit;
2712
2713 macro_rules! setup_step {
2726 ($what:literal, $fut:expr) => {
2727 tokio::select! {
2728 r = tokio::time::timeout(HANDSHAKE_TIMEOUT, $fut) => match r {
2729 Ok(Ok(v)) => v,
2730 Ok(Err(e)) => {
2731 if let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
2732 log::debug!(
2733 concat!(
2734 "CONNECT {} for {target} failed: {e}",
2735 " ({n} similar since the last message)"
2736 ),
2737 $what,
2738 target = target,
2739 e = e,
2740 n = n
2741 );
2742 }
2743 return;
2744 }
2745 Err(_) => {
2746 if let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL) {
2747 log::debug!(
2748 concat!(
2749 "CONNECT {} for {target} did not finish within",
2750 " {timeout:?}",
2751 " ({n} similar since the last message)"
2752 ),
2753 $what,
2754 target = target,
2755 timeout = HANDSHAKE_TIMEOUT,
2756 n = n
2757 );
2758 }
2759 return;
2760 }
2761 },
2762 _ = cancel.cancelled() => {
2765 log::debug!("CONNECT {} for {target} abandoned by shutdown", $what);
2766 return;
2767 }
2768 }
2769 };
2770 }
2771
2772 let upgraded = setup_step!("upgrade", hyper::upgrade::on(req));
2773 let mut client = hyper_util::rt::TokioIo::new(upgraded);
2774 let mut server = setup_step!("connect", tokio::net::TcpStream::connect(target));
2775
2776 tokio::select! {
2781 r = tokio::io::copy_bidirectional(&mut client, &mut server) => {
2782 if let Err(e) = r
2785 && let Some(n) = ABANDONED_TUNNEL.allow(REFUSAL_LOG_INTERVAL)
2786 {
2787 log::debug!(
2788 "CONNECT tunnel to {target} ended: {e} ({n} similar since the last message)"
2789 );
2790 }
2791 }
2792 _ = cancel.cancelled() => {
2793 log::debug!("CONNECT tunnel to {target} closed by shutdown");
2794 }
2795 }
2796 });
2797
2798 Response::builder()
2799 .status(StatusCode::OK)
2800 .header(PITCHFORK_HEADER, "1")
2801 .body(Body::empty())
2802 .unwrap_or_else(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())
2803}
2804
2805fn get_request_host(req: &Request) -> Option<String> {
2811 let authority = req
2813 .uri()
2814 .authority()
2815 .map(|a| a.as_str().to_string())
2816 .filter(|s| !s.is_empty());
2817
2818 authority.or_else(|| {
2819 req.headers()
2820 .get(HOST)
2821 .and_then(|h| h.to_str().ok())
2822 .map(str::to_string)
2823 })
2824}
2825
2826fn join_cookie_fields(headers: &mut HeaderMap) {
2832 let fields: Vec<&[u8]> = headers
2833 .get_all(COOKIE)
2834 .iter()
2835 .map(HeaderValue::as_bytes)
2836 .collect();
2837 if fields.len() < 2 {
2838 return;
2839 }
2840
2841 let joined = HeaderValue::from_bytes(&fields.join(b"; ".as_slice()))
2842 .expect("valid header values joined with \"; \" form a valid header value");
2843 headers.insert(COOKIE, joined);
2844}
2845
2846fn inject_forwarded_headers(req: &mut Request, is_tls: bool, host_header: &str) {
2857 let remote_addr = req
2858 .extensions()
2859 .get::<axum::extract::ConnectInfo<SocketAddr>>()
2860 .map(|ci| ci.0.ip().to_string())
2861 .unwrap_or_else(|| "127.0.0.1".to_string());
2862
2863 let proto = if is_tls { "https" } else { "http" };
2864 let default_port = if is_tls { "443" } else { "80" };
2865
2866 let forwarded_for = remote_addr.clone();
2869 let forwarded_proto = proto.to_string();
2870 let forwarded_host = host_header.to_string();
2871 let forwarded_port = host_header
2872 .rsplit_once(':')
2873 .map(|(_, port)| port.to_string())
2874 .unwrap_or_else(|| default_port.to_string());
2875
2876 for name in [
2882 "x-forwarded-for",
2883 "x-forwarded-proto",
2884 "x-forwarded-host",
2885 "x-forwarded-port",
2886 "forwarded",
2887 ] {
2888 if let Ok(header_name) = axum::http::HeaderName::from_bytes(name.as_bytes()) {
2889 req.headers_mut().remove(&header_name);
2890 }
2891 }
2892
2893 let headers = [
2894 ("x-forwarded-for", forwarded_for),
2895 ("x-forwarded-proto", forwarded_proto),
2896 ("x-forwarded-host", forwarded_host),
2897 ("x-forwarded-port", forwarded_port),
2898 ];
2899
2900 for (name, value) in headers {
2901 if let Ok(v) = HeaderValue::from_str(&value) {
2902 let header_name = axum::http::HeaderName::from_static(name);
2903 req.headers_mut().insert(header_name, v);
2904 }
2905 }
2906}
2907
2908async fn proxy_handler(State(state): State<ProxyState>, mut req: Request) -> Response {
2913 let Some(raw_host) = get_request_host(&req) else {
2915 return error_response(StatusCode::BAD_REQUEST, "Missing Host header");
2916 };
2917
2918 if req.method() == axum::http::Method::CONNECT {
2922 return connect_handler(&state, req, &raw_host).await;
2923 }
2924 let host = if raw_host.starts_with('[') {
2928 raw_host
2930 .split("]:")
2931 .next()
2932 .unwrap_or(&raw_host)
2933 .trim_start_matches('[')
2934 .trim_end_matches(']')
2935 .to_string()
2936 } else {
2937 raw_host.split(':').next().unwrap_or(&raw_host).to_string()
2939 };
2940 let host = host.trim_end_matches('.').to_string();
2945
2946 let is_from_pitchfork = req.headers().contains_key(PROXY_HOPS_HEADER);
2957 let hops: u64 = if is_from_pitchfork {
2958 req.headers()
2959 .get(PROXY_HOPS_HEADER)
2960 .and_then(|v| v.to_str().ok())
2961 .and_then(|s| s.parse().ok())
2962 .unwrap_or(0)
2963 } else {
2964 0
2966 };
2967 if hops >= MAX_PROXY_HOPS {
2968 return error_response(
2969 StatusCode::LOOP_DETECTED,
2970 &format!(
2971 "Loop detected for '{host}': request has passed through the proxy {hops} times.\n\
2972 This usually means a backend is proxying back through pitchfork without rewriting \n\
2973 the Host header. If you use Vite/webpack proxy, set changeOrigin: true."
2974 ),
2975 );
2976 }
2977
2978 let local_client = is_local_client(&req);
2979
2980 let target_port = if let Some(subdomain) = strip_tld(&host, &state.tld) {
2982 if subdomain.eq_ignore_ascii_case("pitchfork") {
2983 crate::web::port()
2984 } else {
2985 None
2986 }
2987 } else {
2988 None
2989 };
2990
2991 let (target_port, activity) = if let Some(port) = target_port {
2992 (port, None)
2993 } else {
2994 if resolve_tls_mode(&host, &state.tld).await.is_passthrough() {
2998 return error_response(
2999 StatusCode::BAD_GATEWAY,
3000 &passthrough_unroutable_message(&host, state.is_tls),
3001 );
3002 }
3003 match resolve_target(&host, &state.tld).await {
3004 ResolveResult::Ready(port, activity) => (port, activity),
3005 ResolveResult::Starting { slug } => {
3006 return starting_html_response(&slug, &raw_host);
3007 }
3008 ResolveResult::Page {
3009 project,
3010 worktree,
3011 daemons,
3012 dir,
3013 } => {
3014 if !local_client {
3018 return unknown_host_response(&host, "Not found", &[]);
3019 }
3020 if let Some(base) = crate::web::url()
3027 && let Some(resolved) = dir.clone()
3028 && let Some(path) = tokio::task::spawn_blocking(move || {
3029 crate::web::routes::api::projects::page_path_for_dir(&resolved)
3030 })
3031 .await
3032 .ok()
3033 .flatten()
3034 {
3035 return page_redirect_response(&base, &path);
3036 }
3037 return page_placeholder_response(
3038 &project,
3039 worktree.as_deref(),
3040 &daemons,
3041 &state.tld,
3042 &host_port_suffix(&raw_host),
3043 crate::web::url().as_deref(),
3044 );
3045 }
3046 ResolveResult::Unknown { heading, known } => {
3047 return unknown_host_response(
3050 &host,
3051 if local_client { &heading } else { "Not found" },
3052 if local_client { &known } else { &[] },
3053 );
3054 }
3055 ResolveResult::NotFound => {
3056 return error_response(
3057 StatusCode::BAD_GATEWAY,
3058 &format!(
3059 "No daemon found for host '{host}'.\n\
3060 A daemon is reachable once it configures a `port` and its project is \
3061 known to pitchfork; run `pitchfork proxy status` to see the hostnames \
3062 it serves.\n\
3063 Expected format: <daemon>.<project>.{tld}",
3064 tld = state.tld
3065 ),
3066 );
3067 }
3068 ResolveResult::Error(msg) => {
3069 if local_client {
3070 return error_response(StatusCode::BAD_GATEWAY, &msg);
3071 }
3072 log::warn!("Refused '{host}' for a non-local client: {msg}");
3074 return error_response(
3075 StatusCode::BAD_GATEWAY,
3076 &format!("'{host}' is not available."),
3077 );
3078 }
3079 }
3080 };
3081 let path_and_query = req
3083 .uri()
3084 .path_and_query()
3085 .map(|pq| pq.as_str())
3086 .unwrap_or("/");
3087
3088 let forward_uri = match Uri::builder()
3089 .scheme("http")
3090 .authority(format!("localhost:{target_port}"))
3091 .path_and_query(path_and_query)
3092 .build()
3093 {
3094 Ok(uri) => uri,
3095 Err(e) => {
3096 return error_response(
3097 StatusCode::INTERNAL_SERVER_ERROR,
3098 &format!("Failed to build forward URI: {e}"),
3099 );
3100 }
3101 };
3102
3103 *req.uri_mut() = forward_uri;
3105 req.headers_mut().insert(
3106 HOST,
3107 HeaderValue::from_str(&format!("localhost:{target_port}"))
3108 .unwrap_or_else(|_| HeaderValue::from_static("localhost")),
3109 );
3110
3111 inject_forwarded_headers(&mut req, state.is_tls, &raw_host);
3113
3114 if let Ok(v) = HeaderValue::from_str(&(hops + 1).to_string()) {
3116 req.headers_mut()
3117 .insert(axum::http::HeaderName::from_static(PROXY_HOPS_HEADER), v);
3118 }
3119
3120 let pseudo_headers: Vec<_> = req
3125 .headers()
3126 .keys()
3127 .filter(|k| k.as_str().starts_with(':'))
3128 .cloned()
3129 .collect();
3130 for key in pseudo_headers {
3131 req.headers_mut().remove(&key);
3132 }
3133
3134 join_cookie_fields(req.headers_mut());
3135
3136 *req.version_mut() = axum::http::Version::HTTP_11;
3141
3142 let client_upgrade = hyper::upgrade::on(&mut req);
3144
3145 let result = match tokio::time::timeout(
3153 std::time::Duration::from_secs(120),
3154 state.client.request(req),
3155 )
3156 .await
3157 {
3158 Ok(r) => r,
3159 Err(_elapsed) => {
3160 let msg = format!(
3161 "Request to daemon on port {target_port} timed out after 120 s.\n\
3162 The daemon accepted the connection but did not respond in time."
3163 );
3164 log::warn!("{msg}");
3165 if let Some(ref on_error) = state.on_error {
3166 on_error(&msg);
3167 }
3168 return error_response(StatusCode::GATEWAY_TIMEOUT, &msg);
3169 }
3170 };
3171 match result {
3172 Ok(mut resp) => {
3173 let backend_upgrade = hyper::upgrade::on(&mut resp);
3175 let (mut parts, body) = resp.into_parts();
3176
3177 parts.headers.insert(
3179 axum::http::HeaderName::from_static(PITCHFORK_HEADER),
3180 HeaderValue::from_static("1"),
3181 );
3182
3183 parts.headers.remove(PROXY_HOPS_HEADER);
3185
3186 if state.is_tls && parts.status != StatusCode::SWITCHING_PROTOCOLS {
3191 for h in HOP_BY_HOP_HEADERS {
3192 if let Ok(name) = axum::http::HeaderName::from_bytes(h.as_bytes()) {
3193 parts.headers.remove(&name);
3194 }
3195 }
3196 }
3197
3198 if parts.status == StatusCode::SWITCHING_PROTOCOLS {
3200 let cancel = state.cancel.clone();
3208 state.tunnels.spawn(async move {
3209 let _activity = activity;
3212 let splice = async move {
3213 if let (Ok(client_upgraded), Ok(backend_upgraded)) =
3214 (client_upgrade.await, backend_upgrade.await)
3215 {
3216 let mut client_io = hyper_util::rt::TokioIo::new(client_upgraded);
3217 let mut backend_io = hyper_util::rt::TokioIo::new(backend_upgraded);
3218 let _ = tokio::io::copy_bidirectional(&mut client_io, &mut backend_io)
3226 .await;
3227 }
3228 };
3229 tokio::select! {
3233 _ = splice => {}
3234 _ = cancel.cancelled() => {}
3235 }
3236 });
3237 return Response::from_parts(parts, Body::empty());
3238 }
3239
3240 Response::from_parts(parts, Body::new(GuardedBody::new(body, activity)))
3246 }
3247 Err(e) => {
3248 let msg = format!(
3249 "Failed to connect to daemon on port {target_port}: {e}\n\
3250 The daemon may have stopped or is not yet ready."
3251 );
3252 if let Some(ref on_error) = state.on_error {
3253 on_error(&msg);
3254 } else {
3255 log::warn!("{msg}");
3256 }
3257 error_response(StatusCode::BAD_GATEWAY, &msg)
3258 }
3259 }
3260}
3261
3262fn passthrough_unroutable_message(host: &str, is_tls: bool) -> String {
3272 if is_tls {
3273 format!(
3274 "'{host}' uses proxy_tls = \"passthrough\", which routes on the host name \
3275 in the TLS ClientHello.\n\
3276 This connection's TLS handshake named a different host, so it was \
3277 terminated here and the request inside it cannot be spliced to the \
3278 daemon.\n\
3279 Make the connection itself name '{host}' rather than overriding the Host \
3280 header of a connection opened to something else."
3281 )
3282 } else {
3283 format!(
3284 "'{host}' uses proxy_tls = \"passthrough\", but the proxy is serving \
3285 plain HTTP.\n\
3286 Passthrough splices a TLS stream to the daemon, so it requires \
3287 settings.proxy.https = true.\n\
3288 Enable HTTPS on the proxy, or set proxy_tls = \"terminate\" on the daemon."
3289 )
3290 }
3291}
3292
3293async fn resolve_target(host: &str, tld: &str) -> ResolveResult {
3312 let deadline = auto_start_deadline();
3315 let ctx = match resolve_route_context(host, tld).await {
3316 Ok(ctx) => ctx,
3317 Err(result) => return result,
3318 };
3319
3320 let daemons = {
3321 let state_file = SUPERVISOR.state_file.lock().await;
3322 state_file.daemons.clone()
3323 };
3324
3325 let daemon_name = &ctx.cached.daemon_name;
3326 let running_matches: Vec<(&DaemonId, &crate::daemon::Daemon)> = daemons
3327 .iter()
3328 .filter(|(id, d)| {
3329 id.name() == daemon_name
3330 && d.status.is_running()
3331 && match &ctx.expected_namespace {
3332 Some(ns) => id.namespace() == ns,
3333 None => true,
3334 }
3335 })
3336 .collect();
3337
3338 match running_matches.as_slice() {
3339 [] => {
3340 try_auto_start(
3341 &ctx.cached.slug,
3342 &ctx.cached,
3343 ctx.worktree_dir.as_deref(),
3344 ctx.expected_namespace.as_deref(),
3345 &ctx.route,
3346 deadline,
3347 )
3348 .await
3349 }
3350 [(id, d), ..] => match select_daemon_port(&ctx.route, d) {
3353 Some(port) => match begin_running(id, deadline).await {
3354 Running::Yes(activity) => ResolveResult::Ready(port, Some(activity)),
3355 Running::Stopping => ResolveResult::Starting {
3356 slug: ctx.cached.slug.clone(),
3357 },
3358 Running::No => {
3359 try_auto_start(
3360 &ctx.cached.slug,
3361 &ctx.cached,
3362 ctx.worktree_dir.as_deref(),
3363 ctx.expected_namespace.as_deref(),
3364 &ctx.route,
3365 deadline,
3366 )
3367 .await
3368 }
3369 },
3370 None => ResolveResult::NotFound,
3371 },
3372 }
3373}
3374
3375enum Running {
3378 Yes(ActivityGuard),
3380 Stopping,
3382 No,
3385}
3386
3387fn auto_start_deadline() -> tokio::time::Instant {
3390 tokio::time::Instant::now() + settings().proxy_auto_start_timeout()
3391}
3392
3393async fn wait_out_idle_stops(ids: &[DaemonId], deadline: tokio::time::Instant) -> bool {
3400 while ids.iter().any(|id| ACTIVITY.is_idle_stopping(id)) {
3401 if tokio::time::Instant::now() >= deadline {
3402 return false;
3403 }
3404 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
3405 }
3406 true
3407}
3408
3409async fn begin_running(id: &DaemonId, deadline: tokio::time::Instant) -> Running {
3420 let activity = match ACTIVITY.begin(id) {
3421 Some(activity) => activity,
3422 None => {
3423 if !wait_out_idle_stops(std::slice::from_ref(id), deadline).await {
3424 return Running::Stopping;
3425 }
3426 match ACTIVITY.begin(id) {
3427 Some(activity) => activity,
3428 None => return Running::Stopping,
3430 }
3431 }
3432 };
3433 let running = {
3434 let state_file = SUPERVISOR.state_file.lock().await;
3435 state_file
3436 .daemons
3437 .get(id)
3438 .is_some_and(|d| d.status.is_running())
3439 };
3440 if running {
3441 Running::Yes(activity)
3442 } else {
3443 Running::No
3444 }
3445}
3446
3447struct RouteContext {
3451 cached: CachedSlugEntry,
3452 expected_namespace: Option<String>,
3453 worktree_dir: Option<std::path::PathBuf>,
3454 route: ProxyTlsRoute,
3455}
3456
3457async fn resolve_route_context(host: &str, tld: &str) -> Result<RouteContext, ResolveResult> {
3466 let Some(subdomain) = strip_tld(host, tld) else {
3467 return Err(ResolveResult::NotFound);
3468 };
3469
3470 let cached = cached_slug_lookup(&subdomain).await.filter(|cached| {
3471 if crate::proxy::hostname::hostname_fits(&cached.slug) {
3475 return true;
3476 }
3477 crate::proxy::hostname::warn_once(&format!(
3478 "Slug '{}' plus the configured proxy.tld is over the DNS length limit, \
3479 so it is not routed.",
3480 cached.slug
3481 ));
3482 false
3483 });
3484 let Some(cached) = cached else {
3485 return Err(resolve_registry_target(&subdomain).await);
3488 };
3489
3490 let (expected_namespace, worktree_dir, route) = if !subdomain.eq_ignore_ascii_case(&cached.slug)
3494 {
3495 let prefix = strip_dot_suffix_ignore_case(&subdomain, &cached.slug);
3496 match prefix {
3497 Some(ref p) => match match_worktree_prefix(&cached, p) {
3498 PrefixMatch::Worktree(wt) => {
3499 let ns = wt.namespace.clone().or_else(|| {
3500 log::warn!(
3501 "Worktree '{}' has no cached namespace; \
3502 falling back to parent slug namespace.",
3503 wt.path.display()
3504 );
3505 cached.namespace.clone()
3506 });
3507 let route = worktree_route(&cached, &wt.sanitized_branch);
3508 (ns, Some(wt.path.clone()), route)
3509 }
3510 PrefixMatch::Ambiguous => {
3511 return Err(ResolveResult::Error(format!(
3512 "'{host}' is ambiguous: more than one branch or workspace of '{slug}' \
3513 sanitizes to the prefix '{p}', and host names are case-insensitive.\n\
3514 Rename one of them so the prefixes differ by more than case, then \
3515 reload.\n\
3516 The supervisor log lists the colliding branches.",
3517 slug = cached.slug,
3518 )));
3519 }
3520 PrefixMatch::Unknown => (cached.namespace.clone(), None, cached.tls),
3521 },
3522 None => (cached.namespace.clone(), None, cached.tls),
3523 }
3524 } else {
3525 (cached.namespace.clone(), None, cached.tls)
3526 };
3527
3528 Ok(RouteContext {
3529 cached,
3530 expected_namespace,
3531 worktree_dir,
3532 route,
3533 })
3534}
3535
3536pub(crate) async fn resolve_tls_mode(host: &str, tld: &str) -> ProxyTlsMode {
3543 let entries = get_cached_slugs().await;
3544 let registry = get_cached_host_registry().await;
3545 resolve_tls_mode_in(host, tld, &entries, ®istry)
3546}
3547
3548fn resolve_tls_mode_in(
3553 host: &str,
3554 tld: &str,
3555 entries: &std::collections::HashMap<String, CachedSlugEntry>,
3556 registry: &crate::proxy::hostname::HostRegistry,
3557) -> ProxyTlsMode {
3558 let Some(subdomain) = strip_tld(host, tld) else {
3559 return ProxyTlsMode::Terminate;
3560 };
3561 let wildcard = settings().proxy.wildcard;
3562
3563 if let Some(cached) = wildcard_slug_lookup(&subdomain, entries, wildcard)
3565 && crate::proxy::hostname::hostname_fits(&cached.slug)
3566 {
3567 if !subdomain.eq_ignore_ascii_case(&cached.slug)
3569 && let Some(prefix) = strip_dot_suffix_ignore_case(&subdomain, &cached.slug)
3570 && let PrefixMatch::Worktree(wt) = match_worktree_prefix(cached, &prefix)
3571 {
3572 return worktree_route(cached, &wt.sanitized_branch).mode;
3573 }
3574 return cached.tls.mode;
3575 }
3576
3577 match registry.resolve(&subdomain, wildcard) {
3580 crate::proxy::hostname::HostTarget::Daemon { proxy_tls, .. } => {
3581 proxy_tls.unwrap_or_default()
3582 }
3583 _ => ProxyTlsMode::Terminate,
3586 }
3587}
3588
3589fn select_daemon_port(route: &ProxyTlsRoute, daemon: &crate::daemon::Daemon) -> Option<u16> {
3606 let Some(want) = route.port else {
3607 let detected = daemon.active_port.filter(|&p| p != 0);
3616 let first_declared = daemon.resolved_port.iter().copied().find(|&p| p != 0);
3617 return if route.mode.is_passthrough() {
3618 first_declared.or(detected)
3619 } else {
3620 detected.or(first_declared)
3621 };
3622 };
3623
3624 let configured = daemon
3625 .port
3626 .as_ref()
3627 .map(|p| p.expect.as_slice())
3628 .unwrap_or(&[]);
3629 if let Some(idx) = configured.iter().position(|&p| p == want)
3630 && let Some(&resolved) = daemon.resolved_port.get(idx)
3631 {
3632 return (resolved != 0).then_some(resolved);
3638 }
3639 if daemon.resolved_port.contains(&want) {
3640 return Some(want);
3641 }
3642 if daemon.resolved_port.is_empty() {
3643 return Some(want);
3646 }
3647 crate::proxy::hostname::warn_once(&format!(
3650 "Daemon {} has proxy_tls_port {want}, which is not among its resolved ports {:?}; \
3651 refusing to route rather than forwarding to a port it never bound. \
3652 Restart the daemon if its ports changed.",
3653 daemon.id, daemon.resolved_port,
3654 ));
3655 None
3656}
3657
3658struct AutoStartGuard {
3665 daemon_id: DaemonId,
3666}
3667
3668impl Drop for AutoStartGuard {
3669 fn drop(&mut self) {
3670 let daemon_id = self.daemon_id.clone();
3671 tokio::spawn(async move {
3675 AUTO_START_IN_PROGRESS.lock().await.remove(&daemon_id);
3676 });
3677 }
3678}
3679
3680static STARTUP_LOCKS: once_cell::sync::Lazy<
3694 std::sync::Mutex<std::collections::HashMap<DaemonId, std::sync::Weak<tokio::sync::Mutex<()>>>>,
3695> = once_cell::sync::Lazy::new(Default::default);
3696
3697async fn lock_startup_graph(mut ids: Vec<DaemonId>) -> Vec<tokio::sync::OwnedMutexGuard<()>> {
3700 ids.sort();
3701 ids.dedup();
3702 let locks: Vec<Arc<tokio::sync::Mutex<()>>> = {
3703 let mut map = STARTUP_LOCKS.lock().unwrap_or_else(|e| e.into_inner());
3704 map.retain(|_, lock| lock.strong_count() > 0);
3705 ids.into_iter()
3706 .map(|id| {
3707 let entry = map.entry(id).or_default();
3708 entry.upgrade().unwrap_or_else(|| {
3709 let lock = Arc::new(tokio::sync::Mutex::new(()));
3710 *entry = Arc::downgrade(&lock);
3711 lock
3712 })
3713 })
3714 .collect()
3715 };
3716 let mut guards = Vec::with_capacity(locks.len());
3717 for lock in locks {
3718 guards.push(lock.lock_owned().await);
3719 }
3720 guards
3721}
3722
3723async fn try_auto_start(
3737 slug: &str,
3738 cached: &CachedSlugEntry,
3739 worktree_dir: Option<&std::path::Path>,
3740 expected_namespace: Option<&str>,
3741 route: &ProxyTlsRoute,
3742 deadline: tokio::time::Instant,
3743) -> ResolveResult {
3744 let s = settings();
3745 if !s.proxy.auto_start {
3746 return ResolveResult::NotFound;
3747 }
3748
3749 let ns = expected_namespace
3750 .map(|s| s.to_string())
3751 .or_else(|| cached.namespace.clone())
3752 .unwrap_or_else(|| "global".to_string());
3753 let daemon_id = match DaemonId::try_new(&ns, &cached.daemon_name) {
3754 Ok(id) => id,
3755 Err(_) => return ResolveResult::NotFound,
3756 };
3757
3758 {
3759 let mut in_progress = AUTO_START_IN_PROGRESS.lock().await;
3760 if !in_progress.insert(daemon_id.clone()) {
3761 return ResolveResult::Starting {
3762 slug: slug.to_string(),
3763 };
3764 }
3765 }
3766
3767 let guard = Arc::new(AutoStartGuard {
3770 daemon_id: daemon_id.clone(),
3771 });
3772
3773 let timeout = s.proxy_auto_start_timeout();
3776 let timed_out = || {
3777 log::warn!("Auto-start: total timeout ({timeout:?}) exceeded for daemon {daemon_id}");
3778 ResolveResult::Error(format!(
3779 "Auto-start for '{daemon_id}' timed out after {timeout:?}.\n\
3780 The daemon and its dependencies did not all become ready, with the daemon \
3781 bound to a port, within the configured proxy_auto_start_timeout.\n\
3782 Startup continues in the background; reload to check again.\n\
3783 Increase the timeout or check the logs of the daemon and its dependencies \
3784 for slow startup."
3785 ))
3786 };
3787
3788 log::info!("Auto-start: starting daemon {daemon_id} for slug '{slug}'");
3789 let start = tokio::spawn({
3790 let guard = guard.clone();
3791 let daemon_id = daemon_id.clone();
3792 let config_dir = worktree_dir.unwrap_or(&cached.dir).to_path_buf();
3793 async move {
3794 let _guard = guard;
3795 start_with_dependencies(&daemon_id, &config_dir).await
3796 }
3797 });
3798 match tokio::time::timeout_at(deadline, start).await {
3799 Ok(Ok(Ok(()))) => {}
3800 Ok(Ok(Err(result))) => return result,
3801 Ok(Err(e)) => {
3802 log::warn!("Auto-start: start task for {daemon_id} failed: {e}");
3803 return ResolveResult::Error(format!("Failed to start daemon '{daemon_id}': {e}"));
3804 }
3805 Err(_elapsed) => return timed_out(),
3806 }
3807
3808 let result =
3809 tokio::time::timeout_at(deadline, wait_for_active_port(slug, &daemon_id, route)).await;
3810 drop(guard);
3811 result.unwrap_or_else(|_elapsed| timed_out())
3812}
3813
3814async fn start_with_dependencies(
3824 daemon_id: &DaemonId,
3825 config_dir: &std::path::Path,
3826) -> std::result::Result<(), ResolveResult> {
3827 let loaded = {
3828 let dir = config_dir.to_path_buf();
3829 tokio::task::spawn_blocking(move || {
3830 crate::pitchfork_toml::PitchforkToml::all_merged_all_namespaces_from(&dir)
3831 })
3832 .await
3833 };
3834 let pt = match loaded {
3835 Ok(Ok(pt)) => pt,
3836 Ok(Err(e)) => {
3837 log::warn!(
3838 "Auto-start: failed to load config from {}: {e}",
3839 config_dir.display()
3840 );
3841 return Err(ResolveResult::NotFound);
3842 }
3843 Err(e) => {
3844 log::warn!("Auto-start: config loading task failed for {daemon_id}: {e}");
3845 return Err(ResolveResult::Error(format!(
3846 "Failed to load configuration: {e}"
3847 )));
3848 }
3849 };
3850
3851 if !pt.daemons.contains_key(daemon_id) {
3852 log::debug!(
3853 "Auto-start: daemon {daemon_id} not found in config at {}",
3854 config_dir.display()
3855 );
3856 return Err(ResolveResult::NotFound);
3857 }
3858
3859 if SUPERVISOR
3861 .state_file
3862 .lock()
3863 .await
3864 .disabled
3865 .contains(daemon_id)
3866 {
3867 return Err(ResolveResult::Error(format!(
3868 "Daemon '{daemon_id}' is disabled, so it is not started.\n\
3869 Enable it with: pitchfork enable {daemon_id}"
3870 )));
3871 }
3872
3873 let graph: Vec<DaemonId> =
3874 match crate::deps::resolve_dependencies(std::slice::from_ref(daemon_id), &pt.daemons) {
3875 Ok(order) => order.levels.into_iter().flatten().collect(),
3876 Err(e) => {
3877 log::warn!("Auto-start: cannot resolve dependencies of {daemon_id}: {e}");
3878 return Err(ResolveResult::Error(format!(
3879 "Cannot start '{daemon_id}': {e}"
3880 )));
3881 }
3882 };
3883 let _activity = match ACTIVITY.begin_all(&graph) {
3891 Some(activity) => activity,
3892 None => {
3893 wait_out_idle_stops(&graph, auto_start_deadline()).await;
3894 match ACTIVITY.begin_all(&graph) {
3895 Some(activity) => activity,
3896 None => {
3897 return Err(ResolveResult::Error(format!(
3898 "A dependency of '{daemon_id}' is still being stopped for inactivity.\n\
3899 Reload to start it again."
3900 )));
3901 }
3902 }
3903 }
3904 };
3905 let proxy_idle = proxy_idle_timeouts(daemon_id, &graph, &pt).await;
3906 let _locks = lock_startup_graph(graph).await;
3907
3908 let ipc = match crate::ipc::client::IpcClient::connect(false).await {
3911 Ok(ipc) => Arc::new(ipc),
3912 Err(e) => {
3913 log::warn!("Auto-start: failed to connect to the supervisor: {e}");
3914 return Err(ResolveResult::Error(format!(
3915 "Failed to start daemon '{daemon_id}': {e}"
3916 )));
3917 }
3918 };
3919 let opts = crate::ipc::batch::StartOptions {
3920 quiet: true,
3921 proxy_idle: Some(proxy_idle),
3922 ..Default::default()
3923 };
3924 let result = match ipc
3925 .start_daemons_with_config(std::slice::from_ref(daemon_id), opts, pt)
3926 .await
3927 {
3928 Ok(result) => result,
3929 Err(e) => {
3930 log::warn!("Auto-start: failed to start {daemon_id}: {e}");
3931 return Err(ResolveResult::Error(format!(
3932 "Failed to start daemon '{daemon_id}': {e}"
3933 )));
3934 }
3935 };
3936 if !result.any_failed {
3937 return Ok(());
3938 }
3939
3940 let message = match result.failed.first() {
3941 Some((id, reason)) if id == daemon_id => {
3942 format!("Daemon '{daemon_id}' failed to start: {reason}\nCheck its logs for errors.")
3943 }
3944 Some((dep, reason)) => format!(
3945 "Daemon '{daemon_id}' was not started because its dependency '{dep}' failed: \
3946 {reason}\n\
3947 Check the logs of '{dep}' for errors."
3948 ),
3949 None => format!(
3950 "Daemon '{daemon_id}' or one of its dependencies failed to start.\n\
3951 Check the supervisor log for errors."
3952 ),
3953 };
3954 log::warn!("Auto-start: {message}");
3955 Err(ResolveResult::Error(message))
3956}
3957
3958async fn wait_for_active_port(
3960 slug: &str,
3961 daemon_id: &DaemonId,
3962 route: &ProxyTlsRoute,
3963) -> ResolveResult {
3964 let poll_interval = std::time::Duration::from_millis(250);
3965
3966 loop {
3967 let daemons = {
3968 let sf = SUPERVISOR.state_file.lock().await;
3969 sf.daemons.clone()
3970 };
3971
3972 if let Some(d) = daemons.get(daemon_id) {
3973 if d.status.is_running() {
3974 if let Some(port) = select_daemon_port(route, d) {
3979 return match ACTIVITY.begin(daemon_id) {
3987 Some(activity) => {
3988 log::info!("Auto-start: daemon {daemon_id} is ready on port {port}");
3989 ResolveResult::Ready(port, Some(activity))
3990 }
3991 None => ResolveResult::Starting {
3992 slug: slug.to_string(),
3993 },
3994 };
3995 }
3996 } else {
3997 log::warn!(
3998 "Auto-start: daemon {daemon_id} is no longer running (status: {})",
3999 d.status
4000 );
4001 return ResolveResult::Error(format!(
4002 "Daemon '{daemon_id}' started but exited unexpectedly.\n\
4003 Check its logs for errors."
4004 ));
4005 }
4006 } else {
4007 log::warn!("Auto-start: daemon {daemon_id} not found in state file after start");
4008 return ResolveResult::Error(format!(
4009 "Daemon '{daemon_id}' started but disappeared from the state file.\n\
4010 Check its logs for errors."
4011 ));
4012 }
4013
4014 tokio::time::sleep(poll_interval).await;
4015 }
4016}
4017
4018async fn proxy_idle_timeouts(
4030 daemon_id: &DaemonId,
4031 closure: &[DaemonId],
4032 pt: &crate::pitchfork_toml::PitchforkToml,
4033) -> std::collections::HashMap<DaemonId, u64> {
4034 let own = |id: &DaemonId| pt.daemons.get(id).and_then(|d| d.proxy_idle_timeout);
4035 let target = match own(daemon_id) {
4036 Some(configured) => configured.duration(),
4037 None => {
4038 let project_dir = crate::ipc::batch::resolve_config_base_dir(
4040 pt.daemons.get(daemon_id).and_then(|d| d.path.as_deref()),
4041 );
4042 tokio::task::spawn_blocking(move || {
4043 crate::settings::Settings::load_from_dir(&project_dir).proxy_idle_timeout()
4044 })
4045 .await
4046 .ok()
4047 .flatten()
4048 }
4049 };
4050 let millis = |d: std::time::Duration| u64::try_from(d.as_millis()).unwrap_or(u64::MAX);
4051 closure
4052 .iter()
4053 .filter_map(|id| {
4054 let grace = if id == daemon_id {
4055 target
4056 } else {
4057 own(id).map_or(target, |configured| configured.duration())
4058 };
4059 grace.map(|g| (id.clone(), millis(g)))
4060 })
4061 .collect()
4062}
4063
4064async fn resolve_registry_target(subdomain: &str) -> ResolveResult {
4069 let registry = get_cached_host_registry().await;
4070 if !crate::proxy::hostname::hostname_fits(subdomain) {
4071 return ResolveResult::Unknown {
4073 heading: "Host name too long".to_string(),
4074 known: registry.project_labels(),
4075 };
4076 }
4077 match registry.resolve(subdomain, settings().proxy.wildcard) {
4078 crate::proxy::hostname::HostTarget::Daemon {
4079 ref dir,
4080 ref namespace,
4081 ref daemon,
4082 proxy_tls,
4083 proxy_tls_port,
4084 ..
4085 } => {
4086 let route = ProxyTlsRoute {
4091 mode: proxy_tls.unwrap_or_default(),
4092 port: proxy_tls_port,
4093 };
4094 let per_checkout = registry.shares_daemon_id(namespace, daemon);
4099 resolve_registry_daemon(subdomain, dir, namespace, daemon, per_checkout, &route).await
4100 }
4101 crate::proxy::hostname::HostTarget::ProjectPage { project } => {
4102 let entry = registry.projects.get(&project);
4103 ResolveResult::Page {
4104 daemons: entry.map(|p| p.primary.labels()).unwrap_or_default(),
4105 dir: entry.map(|p| p.primary.dir.clone()),
4106 project,
4107 worktree: None,
4108 }
4109 }
4110 crate::proxy::hostname::HostTarget::WorktreePage { project, worktree } => {
4111 let checkout = registry
4112 .projects
4113 .get(&project)
4114 .and_then(|p| p.worktrees.get(&worktree));
4115 ResolveResult::Page {
4116 daemons: checkout.map(|c| c.labels()).unwrap_or_default(),
4117 dir: checkout.map(|c| c.dir.clone()),
4118 project,
4119 worktree: Some(worktree),
4120 }
4121 }
4122 crate::proxy::hostname::HostTarget::UnknownProject { known } => ResolveResult::Unknown {
4123 heading: "Unknown project".to_string(),
4124 known,
4125 },
4126 crate::proxy::hostname::HostTarget::UnknownDaemon {
4127 project,
4128 worktree,
4129 known,
4130 } => ResolveResult::Unknown {
4131 heading: match worktree {
4132 Some(wt) => format!("Unknown daemon in '{wt}' of project '{project}'"),
4133 None => format!("Unknown daemon in project '{project}'"),
4134 },
4135 known,
4136 },
4137 }
4138}
4139
4140async fn resolve_registry_daemon(
4147 host: &str,
4148 dir: &std::path::Path,
4149 namespace: &str,
4150 daemon: &str,
4151 per_checkout: bool,
4152 route: &ProxyTlsRoute,
4153) -> ResolveResult {
4154 let deadline = auto_start_deadline();
4156 let daemons = {
4157 let state_file = SUPERVISOR.state_file.lock().await;
4158 state_file.daemons.clone()
4159 };
4160
4161 let mut matches: Vec<crate::daemon::Daemon> = daemons
4162 .iter()
4163 .filter(|(id, d)| {
4164 id.name() == daemon && id.namespace() == namespace && d.status.is_running()
4165 })
4166 .map(|(_, d)| d.clone())
4167 .collect();
4168 matches = sort_by_checkout(matches, dir).await;
4171
4172 if let Some(d) = matches.first() {
4173 if per_checkout && !runs_in_checkout(d.clone(), dir).await {
4176 return ResolveResult::Error(format!(
4177 "'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
4178 running from {}.\n\
4179 These checkouts share the namespace '{namespace}', so pitchfork cannot run \
4180 both copies at once.\n\
4181 Give each checkout its own top-level `namespace`, or stop the other one first.",
4182 dir.display(),
4183 d.dir
4184 .as_deref()
4185 .map(|p| p.display().to_string())
4186 .unwrap_or_else(|| "an unknown directory".to_string()),
4187 ));
4188 }
4189 let Some(port) = select_daemon_port(route, d) else {
4190 return ResolveResult::NotFound;
4191 };
4192 match begin_running(&d.id, deadline).await {
4193 Running::Yes(activity) => return ResolveResult::Ready(port, Some(activity)),
4194 Running::Stopping => {
4195 return ResolveResult::Starting {
4196 slug: host.to_string(),
4197 };
4198 }
4199 Running::No => {}
4201 }
4202 }
4203
4204 let cached = CachedSlugEntry {
4205 slug: host.to_string(),
4206 namespace: Some(namespace.to_string()),
4207 daemon_name: daemon.to_string(),
4208 dir: dir.to_path_buf(),
4209 worktrees: vec![],
4210 rejected_worktree_prefixes: std::collections::HashSet::new(),
4211 tls: *route,
4212 worktree_tls: std::collections::HashMap::new(),
4213 };
4214 let result = try_auto_start(host, &cached, None, Some(namespace), route, deadline).await;
4215
4216 if per_checkout && let ResolveResult::Ready(..) = result {
4220 let started = {
4221 let state_file = SUPERVISOR.state_file.lock().await;
4222 state_file
4223 .daemons
4224 .iter()
4225 .find(|(id, _)| id.name() == daemon && id.namespace() == namespace)
4226 .map(|(_, d)| d.clone())
4227 };
4228 if let Some(d) = started
4229 && !runs_in_checkout(d.clone(), dir).await
4230 {
4231 return ResolveResult::Error(format!(
4232 "'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
4233 running from {}.\n\
4234 These checkouts share the namespace '{namespace}', so pitchfork cannot run \
4235 both copies at once.\n\
4236 Give each checkout its own top-level `namespace`, or stop the other one first.",
4237 dir.display(),
4238 d.dir
4239 .as_deref()
4240 .map(|p| p.display().to_string())
4241 .unwrap_or_else(|| "an unknown directory".to_string()),
4242 ));
4243 }
4244 }
4245
4246 result
4247}
4248
4249fn is_local_client(req: &Request) -> bool {
4256 req.extensions()
4259 .get::<axum::extract::ConnectInfo<SocketAddr>>()
4260 .is_some_and(|ci| ci.0.ip().is_loopback())
4261}
4262
4263fn daemon_runs_in(daemon: &crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
4272 daemon
4273 .dir
4274 .as_deref()
4275 .is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
4276}
4277
4278async fn runs_in_checkout(daemon: crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
4280 let checkout = checkout.to_path_buf();
4281 tokio::task::spawn_blocking(move || daemon_runs_in(&daemon, &checkout))
4282 .await
4283 .unwrap_or(false)
4284}
4285
4286async fn sort_by_checkout(
4292 daemons: Vec<crate::daemon::Daemon>,
4293 checkout: &std::path::Path,
4294) -> Vec<crate::daemon::Daemon> {
4295 if daemons.len() < 2 {
4296 return daemons;
4297 }
4298 let dirs: Vec<Option<std::path::PathBuf>> = daemons.iter().map(|d| d.dir.clone()).collect();
4299 let checkout = checkout.to_path_buf();
4300 let here = tokio::task::spawn_blocking(move || {
4301 dirs.iter()
4302 .map(|dir| {
4303 dir.as_deref()
4304 .is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
4305 })
4306 .collect::<Vec<bool>>()
4307 })
4308 .await;
4309
4310 match here {
4311 Ok(here) => {
4312 let mut ordered: Vec<(bool, crate::daemon::Daemon)> =
4313 here.into_iter().zip(daemons).collect();
4314 ordered.sort_by_key(|(here, _)| !here);
4315 ordered.into_iter().map(|(_, d)| d).collect()
4316 }
4317 Err(e) => {
4318 log::warn!("Checkout attribution task failed: {e}");
4319 daemons
4320 }
4321 }
4322}
4323
4324fn escape_html(s: &str) -> String {
4326 s.replace('&', "&")
4327 .replace('<', "<")
4328 .replace('>', ">")
4329 .replace('"', """)
4330 .replace('\'', "'")
4331}
4332
4333fn html_page(status: StatusCode, title: &str, body: String) -> Response {
4335 let html = format!(
4336 r##"<!DOCTYPE html>
4337<html lang="en">
4338<head>
4339 <meta charset="UTF-8">
4340 <meta name="viewport" content="width=device-width, initial-scale=1">
4341 <title>{title} — pitchfork</title>
4342 <style>
4343 * {{ margin: 0; padding: 0; box-sizing: border-box; }}
4344 body {{
4345 font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
4346 background: #0f1117;
4347 color: #e1e4e8;
4348 display: flex;
4349 align-items: center;
4350 justify-content: center;
4351 min-height: 100vh;
4352 }}
4353 .container {{ max-width: 640px; padding: 2rem; }}
4354 h1 {{ font-size: 1.5rem; font-weight: 600; margin-bottom: 0.75rem; }}
4355 p {{ color: #8b949e; font-size: 0.9rem; margin-bottom: 0.75rem; }}
4356 ul {{ list-style: none; margin: 0.5rem 0 1rem; }}
4357 li {{ margin: 0.25rem 0; }}
4358 code, a {{
4359 color: #58a6ff;
4360 font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
4361 text-decoration: none;
4362 }}
4363 </style>
4364</head>
4365<body>
4366 <div class="container">{body}</div>
4367</body>
4368</html>"##
4369 );
4370 Response::builder()
4371 .status(status)
4372 .header("content-type", "text/html; charset=utf-8")
4373 .body(Body::from(html))
4374 .unwrap_or_else(|_| (status, title.to_string()).into_response())
4375}
4376
4377fn host_port_suffix(raw_host: &str) -> String {
4382 let port = if raw_host.starts_with('[') {
4383 raw_host.split_once("]:").map(|(_, port)| port)
4384 } else {
4385 raw_host.rsplit_once(':').map(|(_, port)| port)
4386 };
4387 port.filter(|p| p.chars().all(|c| c.is_ascii_digit()) && !p.is_empty())
4388 .map(|p| format!(":{p}"))
4389 .unwrap_or_default()
4390}
4391
4392fn page_redirect_response(base: &str, path: &str) -> Response {
4398 let target = format!("{base}{path}");
4399 Response::builder()
4400 .status(StatusCode::FOUND)
4401 .header(axum::http::header::LOCATION, &target)
4402 .header(axum::http::header::CACHE_CONTROL, "no-store")
4403 .body(axum::body::Body::from(format!(
4404 "This page is at {target}\n"
4405 )))
4406 .unwrap_or_else(|_| {
4407 html_page(
4408 StatusCode::INTERNAL_SERVER_ERROR,
4409 "pitchfork",
4410 String::new(),
4411 )
4412 })
4413}
4414
4415fn page_placeholder_response(
4422 project: &str,
4423 worktree: Option<&str>,
4424 daemons: &[String],
4425 tld: &str,
4426 port_suffix: &str,
4427 web_url: Option<&str>,
4428) -> Response {
4429 let heading = match worktree {
4430 Some(wt) => format!("{} · {}", escape_html(project), escape_html(wt)),
4431 None => escape_html(project),
4432 };
4433 let suffix = match worktree {
4434 Some(wt) => format!(
4435 "{}.{}.{}",
4436 escape_html(wt),
4437 escape_html(project),
4438 escape_html(tld)
4439 ),
4440 None => format!("{}.{}", escape_html(project), escape_html(tld)),
4441 };
4442 let list = if daemons.is_empty() {
4443 "<p>No daemon in this checkout has a port configured.</p>".to_string()
4444 } else {
4445 let items: String = daemons
4446 .iter()
4447 .map(|d| {
4448 let d = escape_html(d);
4449 format!("<li><a href=\"//{d}.{suffix}{port_suffix}\">{d}.{suffix}</a></li>")
4450 })
4451 .collect();
4452 format!("<p>Daemons here:</p><ul>{items}</ul>")
4453 };
4454 let page = if worktree.is_some() {
4455 "stack"
4456 } else {
4457 "project"
4458 };
4459 let explanation = match web_url {
4464 Some(url) => format!(
4465 "<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>",
4466 url = escape_html(url),
4467 ),
4468 None => format!(
4469 "<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>"
4470 ),
4471 };
4472 let body = format!("<h1>{heading}</h1>{explanation}{list}");
4473 html_page(StatusCode::OK, "pitchfork", body)
4474}
4475
4476fn unknown_host_response(host: &str, heading: &str, known: &[String]) -> Response {
4478 let list = if known.is_empty() {
4479 "<p>Nothing is registered under this name yet.</p>".to_string()
4480 } else {
4481 let items: String = known
4482 .iter()
4483 .map(|k| format!("<li><code>{}</code></li>", escape_html(k)))
4484 .collect();
4485 format!("<p>Known names:</p><ul>{items}</ul>")
4486 };
4487 let body = format!(
4488 "<h1>{heading}</h1><p>No route for <code>{host}</code>.</p>{list}",
4489 heading = escape_html(heading),
4490 host = escape_html(host),
4491 );
4492 html_page(StatusCode::NOT_FOUND, "Not found", body)
4493}
4494
4495fn strip_tld(host: &str, tld: &str) -> Option<String> {
4510 strip_dot_suffix_ignore_case(host.trim_end_matches('.'), tld)
4511}
4512
4513fn bind_error_message(port: u16, err: &std::io::Error) -> String {
4515 if port < 1024 {
4516 format!(
4517 "Failed to bind proxy server to port {port}: {err}\n\
4518 Hint: ports below 1024 require elevated privileges. Run \
4519 `pitchfork proxy setup`, which grants the bind capability on Linux, \
4520 or set an unprivileged proxy.port and let setup redirect {port} to it."
4521 )
4522 } else {
4523 format!(
4524 "Failed to bind proxy server to port {port}: {err}\n\
4525 Hint: another process may already be using this port."
4526 )
4527 }
4528}
4529
4530fn starting_html_response(slug: &str, raw_host: &str) -> Response {
4535 let escaped_slug = slug
4536 .replace('&', "&")
4537 .replace('<', "<")
4538 .replace('>', ">")
4539 .replace('"', """)
4540 .replace('\'', "'");
4541 let escaped_host = raw_host
4542 .replace('&', "&")
4543 .replace('<', "<")
4544 .replace('>', ">")
4545 .replace('"', """)
4546 .replace('\'', "'");
4547
4548 let html = format!(
4549 r##"<!DOCTYPE html>
4550<html lang="en">
4551<head>
4552 <meta charset="UTF-8">
4553 <meta name="viewport" content="width=device-width, initial-scale=1">
4554 <meta http-equiv="refresh" content="2">
4555 <title>Starting {escaped_slug}… — pitchfork</title>
4556 <style>
4557 * {{ margin: 0; padding: 0; box-sizing: border-box; }}
4558 body {{
4559 font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
4560 background: #0f1117;
4561 color: #e1e4e8;
4562 display: flex;
4563 align-items: center;
4564 justify-content: center;
4565 min-height: 100vh;
4566 }}
4567 .container {{
4568 text-align: center;
4569 max-width: 480px;
4570 padding: 2rem;
4571 }}
4572 .spinner {{
4573 width: 48px;
4574 height: 48px;
4575 border: 4px solid rgba(255, 255, 255, 0.1);
4576 border-top-color: #58a6ff;
4577 border-radius: 50%;
4578 animation: spin 0.8s linear infinite;
4579 margin: 0 auto 1.5rem;
4580 }}
4581 @keyframes spin {{
4582 to {{ transform: rotate(360deg); }}
4583 }}
4584 h1 {{
4585 font-size: 1.5rem;
4586 font-weight: 600;
4587 margin-bottom: 0.5rem;
4588 }}
4589 .slug {{
4590 color: #58a6ff;
4591 font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
4592 }}
4593 .host {{
4594 color: #8b949e;
4595 font-size: 0.875rem;
4596 margin-top: 0.25rem;
4597 }}
4598 .hint {{
4599 color: #8b949e;
4600 font-size: 0.8rem;
4601 margin-top: 1.5rem;
4602 }}
4603 </style>
4604</head>
4605<body>
4606 <div class="container">
4607 <div class="spinner"></div>
4608 <h1>Starting <span class="slug">{escaped_slug}</span>…</h1>
4609 <p class="host">{escaped_host}</p>
4610 <p class="hint">This page will refresh automatically when the daemon is ready.</p>
4611 </div>
4612</body>
4613</html>"##
4614 );
4615
4616 Response::builder()
4617 .status(StatusCode::SERVICE_UNAVAILABLE)
4618 .header("content-type", "text/html; charset=utf-8")
4619 .header("retry-after", "2")
4620 .body(Body::from(html))
4621 .unwrap_or_else(|_| (StatusCode::SERVICE_UNAVAILABLE, "Starting…").into_response())
4622}
4623
4624async fn redirect_to_https_handler(req: Request) -> Response {
4633 if req.headers().contains_key("upgrade") {
4635 log::warn!("Dropping plain-HTTP WebSocket upgrade attempt — use wss:// instead of ws://");
4636 return (
4637 StatusCode::BAD_REQUEST,
4638 "WebSocket over plain HTTP is not supported on the HTTPS port. Use wss:// instead.",
4639 )
4640 .into_response();
4641 }
4642
4643 let raw_host = get_request_host(&req);
4644 let Some(raw_host) = raw_host else {
4645 return (StatusCode::BAD_REQUEST, "Missing Host header").into_response();
4646 };
4647
4648 let hostname = if raw_host.starts_with('[') {
4650 raw_host
4652 .split_once("]:")
4653 .map(|(host, _)| host)
4654 .unwrap_or(&raw_host)
4655 .trim_start_matches('[')
4656 .trim_end_matches(']')
4657 } else {
4658 let mut parts = raw_host.rsplitn(2, ':');
4660 let last = parts.next().unwrap_or(&raw_host);
4661 parts.next().unwrap_or(last)
4662 };
4663
4664 let path = req
4665 .uri()
4666 .path_and_query()
4667 .map(|pq| pq.as_str())
4668 .unwrap_or("/");
4669
4670 let https_port = match u16::try_from(settings().proxy.port).ok().filter(|&p| p > 0) {
4671 Some(443) | None => String::new(),
4672 Some(port) => format!(":{port}"),
4673 };
4674
4675 let host_for_url = if raw_host.starts_with('[') {
4676 format!("[{hostname}]")
4677 } else {
4678 hostname.to_string()
4679 };
4680
4681 let location = format!("https://{host_for_url}{https_port}{path}");
4682 (
4683 StatusCode::FOUND,
4684 [
4685 (axum::http::header::LOCATION, location),
4686 (
4689 axum::http::HeaderName::from_static(PITCHFORK_HEADER),
4690 "1".to_string(),
4691 ),
4692 ],
4693 )
4694 .into_response()
4695}
4696
4697fn error_response(status: StatusCode, message: &str) -> Response {
4699 (
4703 status,
4704 [(
4705 axum::http::HeaderName::from_static(PITCHFORK_HEADER),
4706 HeaderValue::from_static("1"),
4707 )],
4708 message.to_string(),
4709 )
4710 .into_response()
4711}
4712
4713#[cfg(test)]
4714mod tests {
4715 use super::*;
4716
4717 #[tokio::test]
4721 async fn a_request_forwards_after_an_idle_stop_is_called_off() {
4722 let id = DaemonId::new("calledoff", "api");
4723 SUPERVISOR.state_file.lock().await.daemons.insert(
4724 id.clone(),
4725 crate::daemon::Daemon {
4726 id: id.clone(),
4727 status: crate::daemon_status::DaemonStatus::Running,
4728 ..Default::default()
4729 },
4730 );
4731 assert!(ACTIVITY.claim_idle_stop(&id, std::time::Duration::ZERO));
4732 let releaser = {
4733 let id = id.clone();
4734 tokio::spawn(async move {
4735 tokio::time::sleep(std::time::Duration::from_millis(300)).await;
4736 ACTIVITY.release_idle_stop(&id);
4737 })
4738 };
4739 assert!(matches!(
4740 begin_running(&id, auto_start_deadline()).await,
4741 Running::Yes(_)
4742 ));
4743 releaser.await.unwrap();
4744 SUPERVISOR.state_file.lock().await.daemons.remove(&id);
4745 }
4746
4747 #[tokio::test]
4750 async fn a_request_waits_out_an_idle_stop() {
4751 let id = DaemonId::new("waitproj", "api");
4752 assert!(ACTIVITY.claim_idle_stop(&id, std::time::Duration::ZERO));
4753 let releaser = {
4754 let id = id.clone();
4755 tokio::spawn(async move {
4756 tokio::time::sleep(std::time::Duration::from_millis(300)).await;
4757 ACTIVITY.release_idle_stop(&id);
4758 })
4759 };
4760 assert!(matches!(
4761 begin_running(&id, auto_start_deadline()).await,
4762 Running::No
4763 ));
4764 releaser.await.unwrap();
4765 assert!(ACTIVITY.begin(&id).is_some());
4766 }
4767
4768 #[tokio::test]
4774 async fn a_get_only_route_never_reaches_the_fallback() {
4775 async fn routed() -> &'static str {
4776 "routed"
4777 }
4778 async fn fell_through() -> &'static str {
4779 "fell-through"
4780 }
4781
4782 async fn post_to(app: Router) -> String {
4785 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4786
4787 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
4788 let addr = listener.local_addr().unwrap();
4789 let server = tokio::spawn(async move {
4790 let _ = axum::serve(listener, app).await;
4791 });
4792
4793 let mut sock = tokio::net::TcpStream::connect(addr).await.unwrap();
4794 sock.write_all(
4795 b"POST /proxy.pac HTTP/1.1\r\nHost: example\r\nConnection: close\r\n\r\n",
4796 )
4797 .await
4798 .unwrap();
4799 let mut raw = Vec::new();
4800 sock.read_to_end(&mut raw).await.unwrap();
4801 server.abort();
4802
4803 let text = String::from_utf8_lossy(&raw).into_owned();
4804 let status = text.split_whitespace().nth(1).unwrap_or("").to_string();
4805 if status == "405" {
4806 return status;
4807 }
4808 text.rsplit("\r\n").next().unwrap_or("").to_string()
4809 }
4810
4811 let get_only = Router::new()
4812 .route("/proxy.pac", axum::routing::get(routed))
4813 .fallback(fell_through);
4814 assert_eq!(
4815 post_to(get_only).await,
4816 "405",
4817 "axum reached the fallback on a method mismatch, so `any` is unnecessary"
4818 );
4819
4820 let any_method = Router::new()
4821 .route("/proxy.pac", axum::routing::any(routed))
4822 .fallback(fell_through);
4823 assert_eq!(
4824 post_to(any_method).await,
4825 "routed",
4826 "`any` did not deliver the POST to the handler"
4827 );
4828 }
4829
4830 #[cfg(feature = "proxy-tls")]
4831 #[test]
4832 fn a_ca_pair_is_checked_not_just_found() {
4833 let dir = tempfile::tempdir().unwrap();
4834 let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
4835 generate_ca(&cert, &key).unwrap();
4836 assert_eq!(ca_pair_problem(&cert, &key), None);
4837
4838 let (other_cert, other_key) = (dir.path().join("b.pem"), dir.path().join("b-key.pem"));
4840 generate_ca(&other_cert, &other_key).unwrap();
4841 assert!(ca_pair_problem(&cert, &other_key).is_some());
4842
4843 std::fs::write(&cert, "-----BEGIN CERTIFICATE-----\nAAAA\n").unwrap();
4845 assert!(ca_pair_problem(&cert, &key).is_some());
4846 }
4847
4848 #[cfg(feature = "proxy-tls")]
4849 #[test]
4850 fn a_new_ca_clears_leaves_signed_by_the_old_one() {
4851 let dir = tempfile::tempdir().unwrap();
4854 let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
4855 let host_certs = host_certs_dir_for(&cert);
4856 std::fs::create_dir_all(&host_certs).unwrap();
4857 std::fs::write(host_certs.join("api.localhost.pem"), "old leaf").unwrap();
4858 std::fs::write(host_certs.join("notes.txt"), "keep me").unwrap();
4859
4860 assert!(!ensure_ca(&cert, &key, || true).unwrap());
4862 assert!(host_certs.join("api.localhost.pem").exists());
4863
4864 assert!(ensure_ca(&cert, &key, || false).unwrap());
4866 assert!(!host_certs.join("api.localhost.pem").exists());
4867 assert!(host_certs.join("notes.txt").exists());
4868 }
4869
4870 #[cfg(feature = "proxy-tls")]
4871 #[test]
4872 fn leaves_are_cached_per_ca_so_a_replaced_ca_is_never_served() {
4873 let dir = tempfile::tempdir().unwrap();
4877 let (cert, key) = (dir.path().join("ca.pem"), dir.path().join("ca-key.pem"));
4878 let _ = rustls::crypto::ring::default_provider().install_default();
4879 ensure_ca(&cert, &key, || false).unwrap();
4880 let first = SniCertResolver::new(&cert, &key, "localhost".into()).unwrap();
4881 assert!(first.get_or_create_checked("api.localhost").is_some());
4882 let old_dir = first.host_certs_dir.clone();
4883 assert!(std::fs::read_dir(&old_dir).unwrap().count() > 0);
4884
4885 ensure_ca(&cert, &key, || false).unwrap();
4886 std::fs::create_dir_all(&old_dir).unwrap();
4888 std::fs::write(old_dir.join("late.localhost.pem"), "old leaf").unwrap();
4889
4890 let second = SniCertResolver::new(&cert, &key, "localhost".into()).unwrap();
4891 assert_ne!(second.host_certs_dir, old_dir);
4892 assert!(!old_dir.exists());
4894 }
4895
4896 #[cfg(feature = "proxy-tls")]
4897 #[test]
4898 fn cert_cache_file_names_do_not_collide() {
4899 assert_ne!(
4901 cert_cache_file_stem("a_b.localhost"),
4902 cert_cache_file_stem("a.b.localhost")
4903 );
4904 assert_eq!(cert_cache_file_stem("api.localhost"), "api.localhost");
4905 assert_eq!(cert_cache_file_stem("a_b.localhost"), "a%5Fb.localhost");
4906 assert_eq!(
4907 cert_cache_file_stem("*.proj.localhost"),
4908 "%2A.proj.localhost"
4909 );
4910 assert!(!cert_cache_file_stem("../x").contains('/'));
4912 }
4913
4914 #[test]
4915 fn connect_port_reads_the_authority() {
4916 assert_eq!(connect_port("api.localhost:443"), Some(443));
4917 assert_eq!(connect_port("api.localhost:80"), Some(80));
4918 assert_eq!(connect_port("[::1]:443"), Some(443));
4919 assert_eq!(connect_port("api.localhost"), None);
4920 assert_eq!(connect_port("::1"), None);
4921 }
4922
4923 #[test]
4924 fn a_half_configured_certificate_pair_is_refused() {
4925 assert!(tls_pair_problem("", "/k.pem").is_some());
4929 assert!(tls_pair_problem("/c.pem", "").is_some());
4930 assert!(tls_pair_problem("", "").is_none());
4931 assert!(tls_pair_problem("/c.pem", "/k.pem").is_none());
4932 }
4933
4934 #[cfg(feature = "proxy-tls")]
4936 fn test_resolver(tld: &str) -> (SniCertResolver, tempfile::TempDir) {
4937 let dir = tempfile::tempdir().unwrap();
4938 let cert = dir.path().join("ca.pem");
4939 let key = dir.path().join("ca-key.pem");
4940 generate_ca(&cert, &key).unwrap();
4941 let _ = rustls::crypto::ring::default_provider().install_default();
4942 (
4943 SniCertResolver::new(&cert, &key, tld.to_string()).unwrap(),
4944 dir,
4945 )
4946 }
4947
4948 #[cfg(feature = "proxy-tls")]
4950 fn sans_for(resolver: &SniCertResolver, domain: &str) -> Vec<String> {
4951 let ck = resolver.get_or_create(domain).expect("a certificate");
4952 let (_, cert) = x509_parser::parse_x509_certificate(&ck.cert[0]).unwrap();
4953 cert.subject_alternative_name()
4954 .unwrap()
4955 .map(|ext| {
4956 ext.value
4957 .general_names
4958 .iter()
4959 .filter_map(|n| match n {
4960 x509_parser::extensions::GeneralName::DNSName(d) => Some(d.to_string()),
4961 _ => None,
4962 })
4963 .collect()
4964 })
4965 .unwrap_or_default()
4966 }
4967
4968 #[cfg(feature = "proxy-tls")]
4969 #[test]
4970 fn the_ca_refuses_to_sign_for_a_name_outside_the_tld() {
4971 use rustls::server::ResolvesServerCert;
4972
4973 let (resolver, _dir) = test_resolver("localhost");
4974 let cache = resolver.host_certs_dir.clone();
4975
4976 for foreign in [
4980 "login.microsoftonline.com",
4981 "example.com",
4982 "notlocalhost",
4983 "localhost.evil.com",
4984 ] {
4985 assert!(
4986 resolver.get_or_create_checked(foreign).is_none(),
4987 "expected {foreign:?} to be refused"
4988 );
4989 }
4990
4991 let cached: Vec<String> = std::fs::read_dir(&cache)
4993 .map(|rd| {
4994 rd.filter_map(|e| e.ok())
4995 .map(|e| e.file_name().to_string_lossy().into_owned())
4996 .collect()
4997 })
4998 .unwrap_or_default();
4999 assert!(
5000 cached.is_empty(),
5001 "unexpected cached certificates: {cached:?}"
5002 );
5003
5004 for ours in ["localhost", "api.localhost", "core.wt.proj.localhost"] {
5006 assert!(
5007 resolver.get_or_create_checked(ours).is_some(),
5008 "expected {ours:?} to be issued"
5009 );
5010 }
5011 let _ = &resolver as &dyn ResolvesServerCert;
5013 }
5014
5015 #[cfg(feature = "proxy-tls")]
5016 #[test]
5017 fn the_certificate_cache_is_bounded_and_evicts_the_oldest() {
5018 let (resolver, _dir) = test_resolver("localhost");
5021 let host_certs = resolver.host_certs_dir.clone();
5022
5023 for i in 0..MAX_HOST_CERTS + 8 {
5024 assert!(
5025 resolver
5026 .get_or_create_checked(&format!("h{i}.localhost"))
5027 .is_some()
5028 );
5029 }
5030
5031 let cached = resolver.cache.lock().unwrap();
5032 assert_eq!(cached.by_domain.len(), MAX_HOST_CERTS);
5033 assert_eq!(cached.order.len(), MAX_HOST_CERTS);
5034 assert!(cached.get("h0.localhost").is_none());
5036 assert!(
5037 cached
5038 .get(&format!("h{}.localhost", MAX_HOST_CERTS + 7))
5039 .is_some()
5040 );
5041 drop(cached);
5042
5043 let on_disk = std::fs::read_dir(&host_certs).unwrap().count();
5045 assert!(
5046 on_disk <= MAX_HOST_CERTS,
5047 "{on_disk} files cached, expected at most {MAX_HOST_CERTS}"
5048 );
5049 }
5050
5051 #[cfg(feature = "proxy-tls")]
5052 #[test]
5053 fn the_disk_cache_is_pruned_at_startup() {
5054 let dir = tempfile::tempdir().unwrap();
5058 let host_certs = dir.path().join("host-certs");
5059 std::fs::create_dir_all(&host_certs).unwrap();
5060 for i in 0..MAX_HOST_CERTS + 20 {
5061 std::fs::write(host_certs.join(format!("old{i}.pem")), "stale").unwrap();
5062 }
5063 std::fs::write(host_certs.join("notes.txt"), "keep me").unwrap();
5065
5066 prune_host_certs(&host_certs);
5067
5068 let pems = std::fs::read_dir(&host_certs)
5069 .unwrap()
5070 .filter_map(|e| e.ok())
5071 .filter(|e| e.path().extension().is_some_and(|x| x == "pem"))
5072 .count();
5073 assert_eq!(pems, MAX_HOST_CERTS);
5074 assert!(host_certs.join("notes.txt").exists());
5075
5076 let small = dir.path().join("small");
5078 std::fs::create_dir_all(&small).unwrap();
5079 std::fs::write(small.join("a.pem"), "x").unwrap();
5080 prune_host_certs(&small);
5081 assert!(small.join("a.pem").exists());
5082 }
5083
5084 #[cfg(feature = "proxy-tls")]
5085 #[test]
5086 fn a_minted_certificate_never_wildcards_the_whole_tld() {
5087 let (resolver, _dir) = test_resolver("localhost");
5088
5089 let sans = sans_for(&resolver, "api.localhost");
5091 assert!(sans.contains(&"api.localhost".to_string()));
5092 assert!(
5093 !sans.iter().any(|s| s.starts_with('*')),
5094 "unexpected wildcard in {sans:?}"
5095 );
5096
5097 let sans = sans_for(&resolver, "core.wt.proj.localhost");
5099 assert!(sans.contains(&"*.wt.proj.localhost".to_string()));
5100
5101 let (resolver, _dir) = test_resolver("dev.internal");
5103 let sans = sans_for(&resolver, "api.dev.internal");
5104 assert!(
5105 !sans.iter().any(|s| s.starts_with('*')),
5106 "unexpected wildcard in {sans:?}"
5107 );
5108 }
5109
5110 #[tokio::test]
5114 async fn test_page_placeholder_explains_which_step_is_missing() {
5115 async fn body_of(response: Response) -> String {
5116 let (_, body) = response.into_parts();
5117 let bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap();
5118 String::from_utf8(bytes.to_vec()).unwrap()
5119 }
5120
5121 let disabled = page_placeholder_response("shop", None, &[], "localhost", "", None);
5122 assert_eq!(disabled.status(), StatusCode::OK);
5123 let disabled = body_of(disabled).await;
5124 assert!(disabled.contains("auto_start"), "{disabled}");
5125 assert!(!disabled.contains("[namespaces]"));
5126
5127 let running = page_placeholder_response(
5128 "shop",
5129 Some("feature-a"),
5130 &[],
5131 "localhost",
5132 "",
5133 Some("http://127.0.0.1:3120"),
5134 );
5135 let running = body_of(running).await;
5136 assert!(running.contains("[namespaces]"), "{running}");
5137 assert!(running.contains("http://127.0.0.1:3120/projects"));
5138 assert!(!running.contains("auto_start"));
5139 }
5140
5141 #[test]
5144 fn test_page_redirect_targets_the_web_ui() {
5145 let response = page_redirect_response("http://127.0.0.1:3120", "/projects/shop/feature-a");
5146 assert_eq!(response.status(), StatusCode::FOUND);
5147 assert_eq!(
5148 response
5149 .headers()
5150 .get(axum::http::header::LOCATION)
5151 .unwrap(),
5152 "http://127.0.0.1:3120/projects/shop/feature-a"
5153 );
5154
5155 let project = page_redirect_response("http://127.0.0.1:3120/ps", "/projects/shop");
5156 assert_eq!(
5157 project.headers().get(axum::http::header::LOCATION).unwrap(),
5158 "http://127.0.0.1:3120/ps/projects/shop"
5160 );
5161 }
5162
5163 #[tokio::test]
5164 async fn overlapping_startup_graphs_wait_for_each_other() {
5165 use std::time::Duration;
5166 let id = |name: &str| DaemonId::try_new("startup-lock-test", name).unwrap();
5167
5168 let first = lock_startup_graph(vec![id("shared"), id("app-a")]).await;
5169
5170 let second = tokio::spawn(lock_startup_graph(vec![id("shared"), id("app-b")]));
5173 tokio::time::sleep(Duration::from_millis(50)).await;
5174 assert!(
5175 !second.is_finished(),
5176 "a graph sharing a daemon must wait for the graph starting it"
5177 );
5178
5179 let disjoint = tokio::time::timeout(
5181 Duration::from_secs(1),
5182 lock_startup_graph(vec![id("other")]),
5183 )
5184 .await;
5185 assert!(disjoint.is_ok());
5186
5187 drop(first);
5188 let second = tokio::time::timeout(Duration::from_secs(1), second)
5189 .await
5190 .expect("the waiting graph proceeds once the first releases")
5191 .unwrap();
5192 assert_eq!(second.len(), 2);
5193 drop(second);
5194 drop(disjoint);
5195
5196 drop(lock_startup_graph(vec![id("unrelated")]).await);
5198 let map = STARTUP_LOCKS.lock().unwrap();
5199 assert!(
5200 !map.keys()
5201 .any(|k| k.namespace() == "startup-lock-test" && k.name() != "unrelated"),
5202 "released startup locks must be pruned"
5203 );
5204 }
5205
5206 #[test]
5207 fn test_strip_tld() {
5208 assert_eq!(
5209 strip_tld("api.myproject.localhost", "localhost"),
5210 Some("api.myproject".to_string())
5211 );
5212 assert_eq!(
5214 strip_tld("API.MyProject.LOCALHOST", "localhost"),
5215 Some("API.MyProject".to_string())
5216 );
5217 assert_eq!(
5218 strip_tld("api.localhost", "LOCALHOST"),
5219 Some("api".to_string())
5220 );
5221 assert_eq!(
5222 strip_tld("api.localhost", "localhost"),
5223 Some("api".to_string())
5224 );
5225 assert_eq!(strip_tld("localhost", "localhost"), None);
5226 assert_eq!(
5231 strip_tld("API.LocalHost", "localhost"),
5232 Some("API".to_string())
5233 );
5234 assert_eq!(
5235 strip_tld("api.localhost.", "localhost"),
5236 Some("api".to_string())
5237 );
5238 assert_eq!(
5239 strip_tld("API.MyProject.LOCALHOST.", "localhost"),
5240 Some("API.MyProject".to_string())
5241 );
5242 assert_eq!(strip_tld("localhost.", "localhost"), None);
5243 assert_eq!(
5244 strip_tld("api.myproject.test", "test"),
5245 Some("api.myproject".to_string())
5246 );
5247 assert_eq!(strip_tld("other.com", "localhost"), None);
5248 }
5249
5250 fn make_entry(name: &str) -> CachedSlugEntry {
5251 CachedSlugEntry {
5252 slug: name.to_string(),
5253 namespace: None,
5254 daemon_name: name.to_string(),
5255 dir: std::path::PathBuf::from(format!("/tmp/{name}")),
5256 worktrees: vec![],
5257 rejected_worktree_prefixes: std::collections::HashSet::new(),
5258 tls: ProxyTlsRoute::default(),
5259 worktree_tls: std::collections::HashMap::new(),
5260 }
5261 }
5262
5263 fn make_daemon(
5265 configured: &[u16],
5266 resolved: &[u16],
5267 active: Option<u16>,
5268 ) -> crate::daemon::Daemon {
5269 crate::daemon::Daemon {
5270 id: DaemonId::try_new("proj", "api").unwrap(),
5271 port: crate::config_types::PortConfig::from_parts(
5272 configured.to_vec(),
5273 crate::config_types::PortBump(0),
5274 ),
5275 resolved_port: resolved.to_vec(),
5276 active_port: active,
5277 ..crate::daemon::Daemon::default()
5278 }
5279 }
5280
5281 #[test]
5284 fn test_select_daemon_port_prefers_active_port() {
5285 let route = ProxyTlsRoute::default();
5286 let d = make_daemon(&[8443, 9443], &[8443, 9443], Some(8443));
5287 assert_eq!(select_daemon_port(&route, &d), Some(8443));
5288 }
5289
5290 #[test]
5294 fn test_select_daemon_port_skips_port_zero() {
5295 for mode in [ProxyTlsMode::Passthrough, ProxyTlsMode::Terminate] {
5296 let route = ProxyTlsRoute { mode, port: None };
5297
5298 let mixed = make_daemon(&[0, 8443], &[0, 8443], None);
5300 assert_eq!(select_daemon_port(&route, &mixed), Some(8443));
5301
5302 let detected_placeholder = make_daemon(&[0, 8443], &[0, 8443], Some(0));
5305 assert_eq!(
5306 select_daemon_port(&route, &detected_placeholder),
5307 Some(8443)
5308 );
5309
5310 let unresolved = make_daemon(&[0], &[0], Some(0));
5312 assert_eq!(select_daemon_port(&route, &unresolved), None);
5313 }
5314 }
5315
5316 #[test]
5321 fn test_select_daemon_port_passthrough_prefers_declared_first_port() {
5322 let route = ProxyTlsRoute {
5323 mode: ProxyTlsMode::Passthrough,
5324 port: None,
5325 };
5326 let d = make_daemon(&[8443, 9080], &[8443, 9080], Some(9080));
5327 assert_eq!(select_daemon_port(&route, &d), Some(8443));
5328
5329 let detected_only = make_daemon(&[], &[], Some(9080));
5332 assert_eq!(select_daemon_port(&route, &detected_only), Some(9080));
5333 }
5334
5335 #[test]
5337 fn test_select_daemon_port_falls_back_to_first_resolved() {
5338 let route = ProxyTlsRoute::default();
5339 let d = make_daemon(&[8443, 9443], &[8443, 9443], None);
5340 assert_eq!(select_daemon_port(&route, &d), Some(8443));
5341
5342 let none = make_daemon(&[], &[], None);
5344 assert_eq!(select_daemon_port(&route, &none), None);
5345 }
5346
5347 #[test]
5350 fn test_select_daemon_port_honors_configured_port() {
5351 let route = ProxyTlsRoute {
5352 mode: ProxyTlsMode::Passthrough,
5353 port: Some(9443),
5354 };
5355 let d = make_daemon(&[8443, 9443], &[8443, 9443], Some(8443));
5356 assert_eq!(select_daemon_port(&route, &d), Some(9443));
5357 }
5358
5359 #[test]
5363 fn test_select_daemon_port_follows_auto_bump() {
5364 let route = ProxyTlsRoute {
5365 mode: ProxyTlsMode::Passthrough,
5366 port: Some(9443),
5367 };
5368 let d = make_daemon(&[8443, 9443], &[8444, 9444], Some(8444));
5369 assert_eq!(select_daemon_port(&route, &d), Some(9444));
5370 }
5371
5372 #[test]
5376 fn test_select_daemon_port_skips_a_configured_port_resolved_to_zero() {
5377 let route = ProxyTlsRoute {
5378 mode: ProxyTlsMode::Passthrough,
5379 port: Some(9443),
5380 };
5381 let pending = make_daemon(&[8443, 9443], &[8443, 0], Some(8443));
5384 assert_eq!(select_daemon_port(&route, &pending), None);
5385
5386 let ready = make_daemon(&[8443, 9443], &[8443, 9444], Some(8443));
5388 assert_eq!(select_daemon_port(&route, &ready), Some(9444));
5389 }
5390
5391 #[test]
5396 fn test_select_daemon_port_refuses_a_port_the_daemon_never_bound() {
5397 let route = ProxyTlsRoute {
5398 mode: ProxyTlsMode::Passthrough,
5399 port: Some(9443),
5400 };
5401 let stale = make_daemon(&[8443], &[8443], Some(8443));
5402 assert_eq!(select_daemon_port(&route, &stale), None);
5403
5404 let bare = make_daemon(&[], &[], None);
5407 assert_eq!(select_daemon_port(&route, &bare), Some(9443));
5408 }
5409
5410 #[test]
5413 fn test_select_daemon_port_honors_configured_port_when_terminating() {
5414 let route = ProxyTlsRoute {
5415 mode: ProxyTlsMode::Terminate,
5416 port: Some(9080),
5417 };
5418 let d = make_daemon(&[8080, 9080], &[8080, 9080], Some(8080));
5419 assert_eq!(select_daemon_port(&route, &d), Some(9080));
5420
5421 let bumped = make_daemon(&[8080, 9080], &[8081, 9081], Some(8081));
5423 assert_eq!(select_daemon_port(&route, &bumped), Some(9081));
5424 }
5425
5426 #[test]
5430 fn test_read_proxy_tls_route_absent_without_config() {
5431 let dir = tempfile::tempdir().unwrap();
5432
5433 assert_eq!(
5435 read_proxy_tls_route(dir.path(), Some("proj"), "api").unwrap(),
5436 None
5437 );
5438
5439 std::fs::write(
5441 dir.path().join("pitchfork.toml"),
5442 "[daemons.other]\nrun = \"serve\"\n",
5443 )
5444 .unwrap();
5445 assert_eq!(
5446 read_proxy_tls_route(dir.path(), Some("proj"), "api").unwrap(),
5447 None,
5448 "a config without this daemon says nothing about it"
5449 );
5450
5451 assert_eq!(read_proxy_tls_route(dir.path(), None, "api").unwrap(), None);
5453 }
5454
5455 #[test]
5459 fn test_unreadable_config_keeps_the_last_known_route() {
5460 let dir = tempfile::tempdir().unwrap();
5461 std::fs::write(
5462 dir.path().join("pitchfork.toml"),
5463 "[daemons.api]\nrun = \"serve\"\nport = 8443\nproxy_tls = \"passthrough\"\n\n\
5464 [daemons.broken]\nrun = \"serve\"\nproxy_tls = \"passthrough\"\n",
5465 )
5466 .unwrap();
5467 let read = read_proxy_tls_route(dir.path(), Some("proj"), "api");
5468 assert!(
5469 read.is_err(),
5470 "an invalid sibling makes the config unreadable"
5471 );
5472
5473 let known = ProxyTlsRoute {
5474 mode: ProxyTlsMode::Passthrough,
5475 port: Some(8443),
5476 };
5477 assert_eq!(
5478 route_or_last_known(read, Some(known), dir.path(), "api"),
5479 Some(known)
5480 );
5481
5482 assert_eq!(
5484 route_or_last_known(Ok(None), Some(known), dir.path(), "api"),
5485 None
5486 );
5487 }
5488
5489 #[test]
5492 fn test_known_route_requires_the_same_target() {
5493 let known = ProxyTlsRoute {
5494 mode: ProxyTlsMode::Passthrough,
5495 port: Some(8443),
5496 };
5497 let mut entry = make_entry("api");
5498 entry.namespace = Some("proj".to_string());
5499 entry.tls = known;
5500 let dir = entry.dir.clone();
5501
5502 assert_eq!(entry.known_route(&dir, Some("proj"), "api"), Some(known));
5503 assert_eq!(entry.known_route(&dir, Some("proj"), "web"), None);
5504 assert_eq!(entry.known_route(&dir, Some("other"), "api"), None);
5505 assert_eq!(
5506 entry.known_route(std::path::Path::new("/elsewhere"), Some("proj"), "api"),
5507 None
5508 );
5509
5510 let wt = make_worktree("feature/x", "feature-x");
5511 entry.worktrees = vec![wt.clone()];
5512 entry.worktree_tls.insert("feature-x".to_string(), known);
5513 assert_eq!(entry.known_worktree_route(&wt, "api"), Some(known));
5514 assert_eq!(entry.known_worktree_route(&wt, "web"), None);
5515 let moved = crate::proxy::worktree::WorktreeEntry {
5516 path: std::path::PathBuf::from("/elsewhere/feature-x"),
5517 ..wt
5518 };
5519 assert_eq!(entry.known_worktree_route(&moved, "api"), None);
5520 }
5521
5522 #[test]
5526 fn test_worktree_route_inherits_when_unknown() {
5527 let mut entry = make_entry("spliced");
5528 entry.tls = ProxyTlsRoute {
5529 mode: ProxyTlsMode::Passthrough,
5530 port: Some(8443),
5531 };
5532 entry.worktrees = vec![
5533 make_worktree("feature/known", "feature-known"),
5534 make_worktree("feature/unknown", "feature-unknown"),
5535 ];
5536 entry.worktree_tls.insert(
5538 "feature-known".to_string(),
5539 ProxyTlsRoute {
5540 mode: ProxyTlsMode::Terminate,
5541 port: None,
5542 },
5543 );
5544 let mut entries = std::collections::HashMap::new();
5545 entries.insert("spliced".to_string(), entry);
5546
5547 let mode =
5548 |host: &str| resolve_tls_mode_in(host, "localhost", &entries, &Default::default());
5549 assert_eq!(
5550 mode("feature-known.spliced.localhost"),
5551 ProxyTlsMode::Terminate,
5552 "an explicit worktree setting wins"
5553 );
5554 assert_eq!(
5555 mode("feature-unknown.spliced.localhost"),
5556 ProxyTlsMode::Passthrough,
5557 "a worktree with nothing recorded inherits the slug"
5558 );
5559
5560 let cached = entries.get("spliced").unwrap();
5564 assert_eq!(
5565 worktree_route(cached, "feature-unknown"),
5566 ProxyTlsRoute {
5567 mode: ProxyTlsMode::Passthrough,
5568 port: None,
5569 }
5570 );
5571 assert_eq!(
5573 worktree_route(cached, "feature-known"),
5574 ProxyTlsRoute {
5575 mode: ProxyTlsMode::Terminate,
5576 port: None,
5577 }
5578 );
5579 }
5580
5581 #[test]
5584 fn test_worktree_route_lookup() {
5585 let mut entry = make_entry("myapp");
5586 entry.tls = ProxyTlsRoute {
5587 mode: ProxyTlsMode::Terminate,
5588 port: None,
5589 };
5590 entry.worktrees = vec![
5591 make_worktree("feature/b", "feature-b"),
5592 make_worktree("feature/c", "feature-c"),
5593 ];
5594 entry.worktree_tls.insert(
5595 "feature-b".to_string(),
5596 ProxyTlsRoute {
5597 mode: ProxyTlsMode::Passthrough,
5598 port: Some(9443),
5599 },
5600 );
5601
5602 let lookup = |prefix: &str| match match_worktree_prefix(&entry, prefix) {
5603 PrefixMatch::Worktree(wt) => entry
5604 .worktree_tls
5605 .get(&wt.sanitized_branch.to_ascii_lowercase())
5606 .copied()
5607 .unwrap_or(entry.tls),
5608 _ => entry.tls,
5609 };
5610
5611 assert_eq!(lookup("feature-b").mode, ProxyTlsMode::Passthrough);
5612 assert_eq!(lookup("feature-b").port, Some(9443));
5613 assert_eq!(lookup("feature-c").mode, ProxyTlsMode::Terminate);
5615 assert_eq!(lookup("tenant").mode, ProxyTlsMode::Terminate);
5617 }
5618
5619 #[test]
5620 fn test_wildcard_slug_lookup_exact_match() {
5621 let mut entries = std::collections::HashMap::new();
5622 entries.insert("myapp".to_string(), make_entry("myapp"));
5623 let result = wildcard_slug_lookup("myapp", &entries, true);
5625 assert!(result.is_some());
5626 assert_eq!(result.unwrap().daemon_name, "myapp");
5627 }
5628
5629 #[test]
5630 fn test_wildcard_slug_lookup_subdomain_fallback() {
5631 let mut entries = std::collections::HashMap::new();
5632 entries.insert("myapp".to_string(), make_entry("myapp"));
5633 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
5635 assert!(result.is_some());
5636 assert_eq!(result.unwrap().daemon_name, "myapp");
5637 }
5638
5639 #[test]
5640 fn test_wildcard_slug_lookup_nested_fallback() {
5641 let mut entries = std::collections::HashMap::new();
5642 entries.insert("myapp".to_string(), make_entry("myapp"));
5643 let result = wildcard_slug_lookup("a.b.myapp", &entries, true);
5645 assert!(result.is_some());
5646 assert_eq!(result.unwrap().daemon_name, "myapp");
5647 }
5648
5649 #[test]
5650 fn test_wildcard_slug_lookup_no_match() {
5651 let entries = std::collections::HashMap::new();
5652 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
5654 assert!(result.is_none());
5655 }
5656
5657 #[test]
5658 fn test_wildcard_slug_lookup_disabled() {
5659 let mut entries = std::collections::HashMap::new();
5660 entries.insert("myapp".to_string(), make_entry("myapp"));
5661 let result = wildcard_slug_lookup("tenant.myapp", &entries, false);
5663 assert!(result.is_none());
5664 let result = wildcard_slug_lookup("myapp", &entries, false);
5666 assert!(result.is_some());
5667 }
5668
5669 #[test]
5670 fn test_wildcard_slug_lookup_exact_beats_wildcard() {
5671 let mut entries = std::collections::HashMap::new();
5672 entries.insert("myapp".to_string(), make_entry("myapp"));
5673 let mut tenant_entry = make_entry("tenant-daemon");
5674 tenant_entry.slug = "tenant.myapp".to_string();
5675 entries.insert("tenant.myapp".to_string(), tenant_entry);
5676 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
5678 assert!(result.is_some());
5679 assert_eq!(result.unwrap().daemon_name, "tenant-daemon");
5680 }
5681
5682 #[test]
5683 fn test_wildcard_slug_lookup_ignores_case() {
5684 let mut entries = std::collections::HashMap::new();
5685 entries.insert("myapp".to_string(), make_entry("myapp"));
5686 for host in ["MyApp", "MYAPP", "myapp"] {
5688 let result = wildcard_slug_lookup(host, &entries, true);
5689 assert!(result.is_some(), "exact lookup failed for {host}");
5690 assert_eq!(result.unwrap().daemon_name, "myapp");
5691 }
5692 for host in ["Tenant.MyApp", "tenant.MYAPP", "A.B.MyApp"] {
5694 let result = wildcard_slug_lookup(host, &entries, true);
5695 assert!(result.is_some(), "wildcard lookup failed for {host}");
5696 assert_eq!(result.unwrap().daemon_name, "myapp");
5697 }
5698 }
5699
5700 #[test]
5701 fn test_wildcard_slug_lookup_case_insensitive_registration() {
5702 let mut entries = std::collections::HashMap::new();
5703 let mut entry = make_entry("upper");
5704 entry.slug = "MyApp".to_string();
5705 entries.insert("myapp".to_string(), entry);
5708 for host in ["myapp", "MyApp", "tenant.MYAPP"] {
5709 let result = wildcard_slug_lookup(host, &entries, true);
5710 assert!(result.is_some(), "lookup failed for {host}");
5711 assert_eq!(result.unwrap().daemon_name, "upper");
5712 }
5713 }
5714
5715 fn make_worktree(branch: &str, sanitized: &str) -> crate::proxy::worktree::WorktreeEntry {
5716 crate::proxy::worktree::WorktreeEntry {
5717 path: std::path::PathBuf::from(format!("/tmp/{sanitized}")),
5718 branch: branch.to_string(),
5719 sanitized_branch: sanitized.to_string(),
5720 namespace: Some(sanitized.to_string()),
5721 }
5722 }
5723
5724 #[test]
5725 fn test_reject_case_colliding_worktrees_drops_both_sides() {
5726 let wts = vec![
5727 make_worktree("Feature-A", "Feature-A"),
5728 make_worktree("feature-a", "feature-a"),
5729 make_worktree("main", "main"),
5730 ];
5731 let (kept, rejected) = reject_case_colliding_worktrees(wts);
5732 assert_eq!(kept.len(), 1);
5735 assert_eq!(kept[0].sanitized_branch, "main");
5736 assert!(rejected.contains("feature-a"));
5739 }
5740
5741 #[test]
5742 fn test_reject_case_colliding_worktrees_keeps_unambiguous() {
5743 let wts = vec![
5744 make_worktree("main", "main"),
5745 make_worktree("feature/a", "feature-a"),
5746 ];
5747 let (kept, rejected) = reject_case_colliding_worktrees(wts);
5748 assert_eq!(kept.len(), 2);
5749 assert!(rejected.is_empty());
5750 }
5751
5752 #[test]
5753 fn test_reject_case_colliding_worktrees_drops_sanitize_duplicates() {
5754 let wts = vec![
5757 make_worktree("feature/a", "feature-a"),
5758 make_worktree("feature.a", "feature-a"),
5759 ];
5760 let (kept, rejected) = reject_case_colliding_worktrees(wts);
5761 assert!(kept.is_empty());
5762 assert!(rejected.contains("feature-a"));
5763 }
5764
5765 #[test]
5766 fn test_match_worktree_prefix() {
5767 let mut entry = make_entry("myapp");
5768 entry.worktrees = vec![make_worktree("feature/b", "feature-b")];
5769 entry
5770 .rejected_worktree_prefixes
5771 .insert("feature-a".to_string());
5772
5773 assert!(matches!(
5774 match_worktree_prefix(&entry, "feature-b"),
5775 PrefixMatch::Worktree(_)
5776 ));
5777 assert!(matches!(
5779 match_worktree_prefix(&entry, "Feature-B"),
5780 PrefixMatch::Worktree(_)
5781 ));
5782 assert!(matches!(
5784 match_worktree_prefix(&entry, "feature-a"),
5785 PrefixMatch::Ambiguous
5786 ));
5787 assert!(matches!(
5788 match_worktree_prefix(&entry, "FEATURE-A"),
5789 PrefixMatch::Ambiguous
5790 ));
5791 assert!(matches!(
5793 match_worktree_prefix(&entry, "tenant"),
5794 PrefixMatch::Unknown
5795 ));
5796 }
5797
5798 #[test]
5801 fn test_passthrough_unroutable_message() {
5802 let no_tls = passthrough_unroutable_message("api.localhost", false);
5803 assert!(no_tls.contains("api.localhost"), "{no_tls}");
5804 assert!(no_tls.contains("settings.proxy.https = true"), "{no_tls}");
5805
5806 let mismatch = passthrough_unroutable_message("api.localhost", true);
5810 assert!(mismatch.contains("named a different host"), "{mismatch}");
5811 assert!(
5812 !mismatch.contains("settings.proxy.https"),
5813 "a host mismatch is not an HTTPS configuration problem: {mismatch}"
5814 );
5815 }
5816
5817 #[tokio::test]
5821 async fn test_slug_snapshot_is_the_cached_table() {
5822 let from_async = get_cached_slugs().await;
5823 let from_sync = slug_snapshot();
5824 assert!(
5825 Arc::ptr_eq(&from_async, &from_sync),
5826 "the synchronous read must see the same table routing does"
5827 );
5828 }
5829
5830 #[tokio::test]
5835 async fn test_concurrent_refresh_returns_the_published_table() {
5836 let (first, second) = tokio::join!(get_cached_slugs(), get_cached_slugs());
5839 assert!(
5840 Arc::ptr_eq(&first, &second),
5841 "overlapping refreshes must agree on one table"
5842 );
5843 assert!(
5844 Arc::ptr_eq(&first, &slug_snapshot()),
5845 "and it must be the published one"
5846 );
5847 }
5848
5849 #[test]
5853 fn test_resolve_tls_mode_in() {
5854 let mut entries = std::collections::HashMap::new();
5855
5856 let mut spliced = make_entry("spliced");
5857 spliced.tls = ProxyTlsRoute {
5858 mode: ProxyTlsMode::Passthrough,
5859 port: Some(8443),
5860 };
5861 spliced.worktrees = vec![
5862 make_worktree("feature/b", "feature-b"),
5863 make_worktree("feature/c", "feature-c"),
5864 ];
5865 spliced.worktree_tls.insert(
5866 "feature-b".to_string(),
5867 ProxyTlsRoute {
5868 mode: ProxyTlsMode::Terminate,
5869 port: None,
5870 },
5871 );
5872 entries.insert("spliced".to_string(), spliced);
5873 entries.insert("plain".to_string(), make_entry("plain"));
5874
5875 let mode =
5876 |host: &str| resolve_tls_mode_in(host, "localhost", &entries, &Default::default());
5877
5878 assert_eq!(mode("spliced.localhost"), ProxyTlsMode::Passthrough);
5879 assert_eq!(mode("SPLICED.localhost"), ProxyTlsMode::Passthrough);
5883 assert_eq!(mode("Spliced.LocalHost"), ProxyTlsMode::Passthrough);
5884 assert_eq!(mode("spliced.localhost."), ProxyTlsMode::Passthrough);
5885 assert_eq!(
5886 mode("FEATURE-C.Spliced.LOCALHOST."),
5887 ProxyTlsMode::Passthrough
5888 );
5889 assert_eq!(mode("tenant.spliced.localhost"), ProxyTlsMode::Passthrough);
5891 assert_eq!(mode("feature-b.spliced.localhost"), ProxyTlsMode::Terminate);
5893 assert_eq!(
5895 mode("feature-c.spliced.localhost"),
5896 ProxyTlsMode::Passthrough
5897 );
5898 assert_eq!(mode("plain.localhost"), ProxyTlsMode::Terminate);
5901 assert_eq!(mode("unknown.localhost"), ProxyTlsMode::Terminate);
5902 assert_eq!(mode("localhost"), ProxyTlsMode::Terminate);
5903 assert_eq!(mode("spliced.example.com"), ProxyTlsMode::Terminate);
5904 }
5905
5906 #[test]
5911 fn test_resolve_tls_mode_in_uses_the_hostname_registry() {
5912 let dir = tempfile::tempdir().unwrap();
5913 let project = dir.path().join("autoproj");
5914 std::fs::create_dir_all(&project).unwrap();
5915 std::fs::write(
5916 project.join("pitchfork.toml"),
5917 "[daemons.secure]\nrun = \"serve\"\nport = 8443\nproxy_tls = \"passthrough\"\n\
5918 [daemons.plain]\nrun = \"serve\"\nport = 8080\n",
5919 )
5920 .unwrap();
5921
5922 let registry =
5923 crate::proxy::hostname::HostRegistry::from_dirs(std::slice::from_ref(&project));
5924 let slugs = std::collections::HashMap::new();
5925 let mode = |host: &str| resolve_tls_mode_in(host, "localhost", &slugs, ®istry);
5926
5927 assert_eq!(mode("secure.autoproj.localhost"), ProxyTlsMode::Passthrough);
5928 assert_eq!(mode("plain.autoproj.localhost"), ProxyTlsMode::Terminate);
5929
5930 std::fs::write(project.join("pitchfork.toml"), "this is not toml = [\n").unwrap();
5934 assert_eq!(mode("secure.autoproj.localhost"), ProxyTlsMode::Passthrough);
5935 assert_eq!(mode("autoproj.localhost"), ProxyTlsMode::Terminate);
5938 assert_eq!(mode("nothing.autoproj.localhost"), ProxyTlsMode::Terminate);
5939 assert_eq!(mode("unknown.localhost"), ProxyTlsMode::Terminate);
5940 }
5941
5942 #[test]
5945 fn test_resolve_tls_mode_in_empty_table() {
5946 let entries = std::collections::HashMap::new();
5947 assert_eq!(
5948 resolve_tls_mode_in(
5949 "spliced.localhost",
5950 "localhost",
5951 &entries,
5952 &Default::default()
5953 ),
5954 ProxyTlsMode::Terminate
5955 );
5956 }
5957
5958 #[test]
5959 fn test_strip_dot_suffix_ignore_case() {
5960 assert_eq!(
5961 strip_dot_suffix_ignore_case("feature-a.myapp", "myapp"),
5962 Some("feature-a".to_string())
5963 );
5964 assert_eq!(
5965 strip_dot_suffix_ignore_case("Feature-A.MyApp", "myapp"),
5966 Some("Feature-A".to_string())
5967 );
5968 assert_eq!(
5969 strip_dot_suffix_ignore_case("feature-a.myapp", "MYAPP"),
5970 Some("feature-a".to_string())
5971 );
5972 assert_eq!(strip_dot_suffix_ignore_case("xmyapp", "myapp"), None);
5974 assert_eq!(strip_dot_suffix_ignore_case(".myapp", "myapp"), None);
5975 assert_eq!(strip_dot_suffix_ignore_case("myapp", "myapp"), None);
5976 assert_eq!(
5977 strip_dot_suffix_ignore_case("feature-a.other", "myapp"),
5978 None
5979 );
5980 assert_eq!(
5982 strip_dot_suffix_ignore_case("café.myapp", "myapp"),
5983 Some("café".to_string())
5984 );
5985 assert_eq!(strip_dot_suffix_ignore_case("café", "afé"), None);
5986 }
5987
5988 #[cfg(feature = "proxy-tls")]
5990 fn client_hello_wire(host: &str) -> Vec<u8> {
5991 let mut entry = vec![0u8];
5992 entry.extend_from_slice(&(host.len() as u16).to_be_bytes());
5993 entry.extend_from_slice(host.as_bytes());
5994 let mut sni = (entry.len() as u16).to_be_bytes().to_vec();
5995 sni.extend_from_slice(&entry);
5996
5997 let mut ext = vec![0x00, 0x00];
5998 ext.extend_from_slice(&(sni.len() as u16).to_be_bytes());
5999 ext.extend_from_slice(&sni);
6000
6001 let mut body = vec![0x03, 0x03];
6002 body.extend_from_slice(&[0x22; 32]);
6003 body.push(0);
6004 body.extend_from_slice(&[0x00, 0x02, 0x13, 0x01]);
6005 body.extend_from_slice(&[0x01, 0x00]);
6006 body.extend_from_slice(&(ext.len() as u16).to_be_bytes());
6007 body.extend_from_slice(&ext);
6008
6009 let mut msg = vec![0x01];
6010 let len = body.len() as u32;
6011 msg.extend_from_slice(&[(len >> 16) as u8, (len >> 8) as u8, len as u8]);
6012 msg.extend_from_slice(&body);
6013
6014 let mut record = vec![0x16, 0x03, 0x01];
6015 record.extend_from_slice(&(msg.len() as u16).to_be_bytes());
6016 record.extend_from_slice(&msg);
6017 record
6018 }
6019
6020 #[cfg(feature = "proxy-tls")]
6023 async fn probe_over_socket(
6024 writes: Vec<Vec<u8>>,
6025 gap: std::time::Duration,
6026 timeout: std::time::Duration,
6027 ) -> (SniProbe, Vec<u8>) {
6028 probe_over_socket_then(writes, gap, timeout, false).await
6029 }
6030
6031 #[cfg(feature = "proxy-tls")]
6034 async fn probe_over_socket_then(
6035 writes: Vec<Vec<u8>>,
6036 gap: std::time::Duration,
6037 timeout: std::time::Duration,
6038 close_after: bool,
6039 ) -> (SniProbe, Vec<u8>) {
6040 use tokio::io::{AsyncReadExt, AsyncWriteExt};
6041
6042 let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
6043 .await
6044 .unwrap();
6045 let addr = listener.local_addr().unwrap();
6046 let total: usize = writes.iter().map(Vec::len).sum();
6047
6048 let client = tokio::spawn(async move {
6049 let mut sock = TcpStream::connect(addr).await.unwrap();
6050 for chunk in writes {
6051 sock.write_all(&chunk).await.unwrap();
6052 sock.flush().await.unwrap();
6053 tokio::time::sleep(gap).await;
6054 }
6055 if close_after {
6056 sock.shutdown().await.unwrap();
6057 }
6058 tokio::time::sleep(std::time::Duration::from_secs(2)).await;
6061 });
6062
6063 let (stream, _) = listener.accept().await.unwrap();
6064 let probe = peek_sni_host(&stream, timeout).await;
6065
6066 let mut replayed = vec![0u8; total];
6068 let mut stream = stream;
6069 let read = tokio::time::timeout(
6070 std::time::Duration::from_secs(2),
6071 stream.read_exact(&mut replayed),
6072 )
6073 .await;
6074 let replayed = match read {
6075 Ok(Ok(_)) => replayed,
6076 _ => vec![],
6077 };
6078 client.abort();
6079 (probe, replayed)
6080 }
6081
6082 #[cfg(feature = "proxy-tls")]
6086 #[tokio::test]
6087 async fn test_peek_sni_host_reads_a_hello_split_across_writes() {
6088 let wire = client_hello_wire("api.localhost");
6089 let writes: Vec<Vec<u8>> = wire.chunks(3).map(<[u8]>::to_vec).collect();
6090 let (probe, replayed) = probe_over_socket(
6091 writes,
6092 std::time::Duration::from_millis(5),
6093 std::time::Duration::from_secs(5),
6094 )
6095 .await;
6096
6097 assert_eq!(probe, SniProbe::Host("api.localhost".to_string()));
6098 assert_eq!(replayed, wire, "peeked bytes must still be readable");
6099 }
6100
6101 #[cfg(feature = "proxy-tls")]
6105 #[tokio::test]
6106 async fn test_peek_sni_host_undetermined_when_a_hello_stalls() {
6107 let wire = client_hello_wire("api.localhost");
6108 let truncated = wire[..wire.len() / 2].to_vec();
6109 let (probe, _) = probe_over_socket(
6110 vec![truncated],
6111 std::time::Duration::ZERO,
6112 std::time::Duration::from_millis(150),
6113 )
6114 .await;
6115
6116 assert_eq!(probe, SniProbe::Undetermined);
6117 }
6118
6119 #[cfg(feature = "proxy-tls")]
6123 #[tokio::test]
6124 async fn test_peek_sni_host_notices_a_client_that_closes_mid_hello() {
6125 let wire = client_hello_wire("api.localhost");
6126 let truncated = wire[..wire.len() / 2].to_vec();
6127 let started = std::time::Instant::now();
6128 let (probe, _) = probe_over_socket_then(
6129 vec![truncated],
6130 std::time::Duration::ZERO,
6131 std::time::Duration::from_secs(5),
6132 true,
6133 )
6134 .await;
6135
6136 assert_eq!(probe, SniProbe::Undetermined);
6137 assert!(
6138 started.elapsed() < std::time::Duration::from_secs(1),
6139 "the probe must end when the client closes, not at its deadline"
6140 );
6141 }
6142
6143 #[cfg(feature = "proxy-tls")]
6147 #[tokio::test]
6148 async fn test_peek_sni_host_gives_up_on_a_silent_peer() {
6149 let started = std::time::Instant::now();
6150 let (probe, _) = probe_over_socket(
6151 vec![],
6152 std::time::Duration::ZERO,
6153 std::time::Duration::from_millis(150),
6154 )
6155 .await;
6156
6157 assert_eq!(probe, SniProbe::Undetermined);
6158 assert!(
6159 started.elapsed() < std::time::Duration::from_secs(1),
6160 "the probe must end at its deadline, not wait on the peer"
6161 );
6162 }
6163
6164 #[cfg(feature = "proxy-tls")]
6167 #[tokio::test]
6168 async fn test_peek_sni_host_reports_no_host_for_non_tls() {
6169 let (probe, _) = probe_over_socket(
6170 vec![b"GET / HTTP/1.1\r\n\r\n".to_vec()],
6171 std::time::Duration::ZERO,
6172 std::time::Duration::from_secs(5),
6173 )
6174 .await;
6175
6176 assert_eq!(probe, SniProbe::NoHost);
6177 }
6178
6179 #[cfg(feature = "proxy-tls")]
6180 #[test]
6181 fn test_generate_ca() {
6182 let dir = tempfile::tempdir().unwrap();
6183 let cert_path = dir.path().join("ca.pem");
6184 let key_path = dir.path().join("ca-key.pem");
6185
6186 generate_ca(&cert_path, &key_path).unwrap();
6187
6188 assert!(cert_path.exists(), "ca.pem should be created");
6189 assert!(key_path.exists(), "ca-key.pem should be created");
6190
6191 let cert_pem = std::fs::read_to_string(&cert_path).unwrap();
6192 let key_pem = std::fs::read_to_string(&key_path).unwrap();
6193
6194 assert!(cert_pem.contains("BEGIN CERTIFICATE"), "should be PEM cert");
6195 assert!(
6196 key_pem.contains("BEGIN") && key_pem.contains("PRIVATE KEY"),
6197 "should be PEM key"
6198 );
6199 }
6200
6201 fn cookie_fields(headers: &HeaderMap) -> Vec<&[u8]> {
6203 headers
6204 .get_all(COOKIE)
6205 .iter()
6206 .map(HeaderValue::as_bytes)
6207 .collect()
6208 }
6209
6210 #[test]
6213 fn test_join_cookie_fields_joins_with_semicolon_space() {
6214 let mut headers = HeaderMap::new();
6215 headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
6216 headers.append(COOKIE, HeaderValue::from_static("consent=ads,stats"));
6217 headers.append(COOKIE, HeaderValue::from_static("theme=dark"));
6218
6219 join_cookie_fields(&mut headers);
6220
6221 assert_eq!(
6222 cookie_fields(&headers),
6223 vec![&b"_session=abc123; consent=ads,stats; theme=dark"[..]]
6224 );
6225 }
6226
6227 #[test]
6230 fn test_join_cookie_fields_joins_bytes_outside_ascii() {
6231 let mut headers = HeaderMap::new();
6232 headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
6233 headers.append(
6234 COOKIE,
6235 HeaderValue::from_bytes(b"name=Jos\xc3\xa9").unwrap(),
6236 );
6237
6238 join_cookie_fields(&mut headers);
6239
6240 assert_eq!(
6241 cookie_fields(&headers),
6242 vec![&b"_session=abc123; name=Jos\xc3\xa9"[..]]
6243 );
6244 }
6245
6246 #[test]
6248 fn test_join_cookie_fields_without_cookies() {
6249 let mut headers = HeaderMap::new();
6250 headers.insert(HOST, HeaderValue::from_static("app.localhost"));
6251
6252 join_cookie_fields(&mut headers);
6253
6254 assert!(headers.get(COOKIE).is_none());
6255 }
6256
6257 #[test]
6258 fn test_host_port_suffix() {
6259 assert_eq!(host_port_suffix("api.myproj.localhost:8088"), ":8088");
6260 assert_eq!(host_port_suffix("api.myproj.localhost"), "");
6261 assert_eq!(host_port_suffix("[::1]:8088"), ":8088");
6262 assert_eq!(host_port_suffix("[::1]"), "");
6263 assert_eq!(host_port_suffix("host:notaport"), "");
6265 }
6266
6267 #[test]
6271 fn test_daemon_runs_in() {
6272 let temp = tempfile::tempdir().unwrap();
6273 let repo = temp.path().join("my-repo");
6274 std::fs::create_dir_all(repo.join(".git/worktrees/feature")).unwrap();
6275 std::fs::create_dir_all(repo.join("sub")).unwrap();
6276 let nested = repo.join(".worktrees/feature");
6278 std::fs::create_dir_all(&nested).unwrap();
6279 std::fs::write(
6280 nested.join(".git"),
6281 format!(
6282 "gitdir: {}\n",
6283 repo.join(".git/worktrees/feature").display()
6284 ),
6285 )
6286 .unwrap();
6287
6288 let root = |p: &std::path::Path| crate::proxy::hostname::checkout_root_of(p);
6289 let repo_root = root(&repo);
6290 let nested_root = root(&nested);
6291
6292 let mut daemon = crate::daemon::Daemon {
6293 dir: Some(repo.join("sub")),
6294 ..Default::default()
6295 };
6296 assert!(daemon_runs_in(&daemon, &repo_root));
6297 assert!(!daemon_runs_in(&daemon, &nested_root));
6298
6299 daemon.dir = Some(nested.clone());
6302 assert!(daemon_runs_in(&daemon, &nested_root));
6303 assert!(!daemon_runs_in(&daemon, &repo_root));
6304
6305 daemon.dir = Some(temp.path().join("elsewhere"));
6307 assert!(!daemon_runs_in(&daemon, &repo_root));
6308
6309 daemon.dir = None;
6310 assert!(!daemon_runs_in(&daemon, &repo_root));
6311 }
6312
6313 #[test]
6316 fn test_is_local_client() {
6317 let build = |info: Option<SocketAddr>| {
6318 let mut req = Request::new(Body::empty());
6319 if let Some(addr) = info {
6320 req.extensions_mut()
6321 .insert(axum::extract::ConnectInfo(addr));
6322 }
6323 req
6324 };
6325
6326 assert!(is_local_client(&build(Some(
6327 "127.0.0.1:5000".parse().unwrap()
6328 ))));
6329 assert!(is_local_client(&build(Some("[::1]:5000".parse().unwrap()))));
6330 assert!(!is_local_client(&build(Some(
6331 "192.168.1.42:5000".parse().unwrap()
6332 ))));
6333 assert!(!is_local_client(&build(None)));
6334 }
6335}