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 HOP_BY_HOP_HEADERS: &[&str] = &[
33 "connection",
34 "keep-alive",
35 "proxy-connection",
36 "transfer-encoding",
37 "upgrade",
38];
39
40use hyper_util::client::legacy::Client;
41use hyper_util::client::legacy::connect::HttpConnector;
42use hyper_util::rt::TokioExecutor;
43use tokio::net::TcpListener;
44
45use crate::daemon_id::DaemonId;
46use crate::settings::settings;
47use crate::supervisor::SUPERVISOR;
48
49const SLUG_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(2);
62
63#[derive(Clone, Debug)]
65pub struct CachedSlugEntry {
66 pub slug: String,
68 pub namespace: Option<String>,
70 pub daemon_name: String,
72 pub dir: std::path::PathBuf,
74 pub worktrees: Vec<crate::proxy::worktree::WorktreeEntry>,
76 pub rejected_worktree_prefixes: std::collections::HashSet<String>,
80}
81
82struct SlugCache {
84 entries: Arc<std::collections::HashMap<String, CachedSlugEntry>>,
85 expires_at: std::time::Instant,
86}
87
88static SLUG_CACHE: once_cell::sync::Lazy<tokio::sync::Mutex<SlugCache>> =
89 once_cell::sync::Lazy::new(|| {
90 tokio::sync::Mutex::new(SlugCache {
91 entries: Arc::new(std::collections::HashMap::new()),
92 expires_at: std::time::Instant::now(), })
94 });
95
96fn reject_case_colliding_worktrees(
103 wts: Vec<crate::proxy::worktree::WorktreeEntry>,
104) -> (
105 Vec<crate::proxy::worktree::WorktreeEntry>,
106 std::collections::HashSet<String>,
107) {
108 let collisions =
109 crate::proxy::ascii_case_collisions(wts.iter().map(|w| w.sanitized_branch.as_str()));
110 if collisions.is_empty() {
111 return (wts, collisions);
112 }
113
114 let (dropped, kept): (Vec<_>, Vec<_>) = wts
115 .into_iter()
116 .partition(|w| collisions.contains(&w.sanitized_branch.to_ascii_lowercase()));
117
118 let mut folded: Vec<&String> = collisions.iter().collect();
119 folded.sort();
120 for key in folded {
121 let mut branches: Vec<&str> = dropped
122 .iter()
123 .filter(|w| w.sanitized_branch.eq_ignore_ascii_case(key))
124 .map(|w| w.branch.as_str())
125 .collect();
126 branches.sort();
127 log::warn!(
128 "Worktree slug collision: branches [{}] all route to '{key}' under \
129 case-insensitive host matching. None of them will be routed; \
130 rename a branch to disambiguate.",
131 branches.join(", "),
132 );
133 }
134
135 (kept, collisions)
136}
137
138fn build_slug_entries() -> std::collections::HashMap<String, CachedSlugEntry> {
145 let global_slugs = crate::pitchfork_toml::PitchforkToml::read_global_slugs();
146 let collisions = crate::proxy::ascii_case_collisions(global_slugs.keys().map(String::as_str));
147 let mut folded: Vec<&String> = collisions.iter().collect();
148 folded.sort();
149 for key in folded {
150 let mut spellings: Vec<&str> = global_slugs
151 .keys()
152 .filter(|s| s.eq_ignore_ascii_case(key))
153 .map(String::as_str)
154 .collect();
155 spellings.sort();
156 log::warn!(
157 "Slug collision: [{}] differ only by case and host names are case-insensitive. \
158 None of them will be routed; remove or rename all but one.",
159 spellings.join(", "),
160 );
161 }
162
163 let mut entries: std::collections::HashMap<String, CachedSlugEntry> =
164 std::collections::HashMap::with_capacity(global_slugs.len());
165 let worktree_enabled = crate::settings::settings().general.worktree;
166 for (slug, entry) in &global_slugs {
167 let key = slug.to_ascii_lowercase();
168 if collisions.contains(&key) {
169 continue;
170 }
171 let ns = entry.resolve_namespace();
172 let daemon_name = entry.daemon.as_deref().unwrap_or(slug).to_string();
173 let (worktrees, rejected_worktree_prefixes) = if worktree_enabled {
174 let wts = match entry.resolve_dir() {
175 Some(dir) => crate::proxy::worktree::discover_worktrees(&dir),
176 None => vec![],
177 };
178 let wts = wts
179 .into_iter()
180 .map(|mut wt| {
181 wt.namespace =
182 crate::pitchfork_toml::PitchforkToml::namespace_for_dir(&wt.path).ok();
183 wt
184 })
185 .collect();
186 reject_case_colliding_worktrees(wts)
187 } else {
188 (vec![], std::collections::HashSet::new())
189 };
190 entries.insert(
191 key,
192 CachedSlugEntry {
193 slug: slug.clone(),
194 namespace: ns,
195 daemon_name,
196 dir: entry.resolve_dir().unwrap_or_default(),
197 worktrees,
198 rejected_worktree_prefixes,
199 },
200 );
201 }
202 entries
203}
204
205pub async fn get_cached_slugs() -> Arc<std::collections::HashMap<String, CachedSlugEntry>> {
211 {
213 let cache = SLUG_CACHE.lock().await;
214 if std::time::Instant::now() < cache.expires_at {
215 return Arc::clone(&cache.entries);
216 }
217 } let new_entries = Arc::new(
221 tokio::task::spawn_blocking(build_slug_entries)
222 .await
223 .unwrap_or_else(|e| {
224 log::warn!("Failed to refresh slug cache: {e}");
225 std::collections::HashMap::new()
226 }),
227 );
228
229 {
231 let mut cache = SLUG_CACHE.lock().await;
232 cache.entries = Arc::clone(&new_entries);
233 cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
234 }
235
236 new_entries
237}
238
239struct RegistryCache {
247 registry: Arc<crate::proxy::hostname::HostRegistry>,
248 expires_at: std::time::Instant,
249}
250
251static HOST_REGISTRY: once_cell::sync::Lazy<tokio::sync::Mutex<RegistryCache>> =
252 once_cell::sync::Lazy::new(|| {
253 tokio::sync::Mutex::new(RegistryCache {
254 registry: Arc::new(crate::proxy::hostname::HostRegistry::default()),
255 expires_at: std::time::Instant::now(), })
257 });
258
259pub async fn get_cached_host_registry() -> Arc<crate::proxy::hostname::HostRegistry> {
261 {
262 let cache = HOST_REGISTRY.lock().await;
263 if std::time::Instant::now() < cache.expires_at {
264 return Arc::clone(&cache.registry);
265 }
266 } let registry = Arc::new(
269 tokio::task::spawn_blocking(crate::proxy::hostname::HostRegistry::build)
270 .await
271 .unwrap_or_else(|e| {
272 log::warn!("Failed to refresh hostname registry: {e}");
273 crate::proxy::hostname::HostRegistry::default()
274 }),
275 );
276 for err in ®istry.errors {
277 crate::proxy::hostname::warn_once(err);
278 }
279
280 {
281 let mut cache = HOST_REGISTRY.lock().await;
282 cache.registry = Arc::clone(®istry);
283 cache.expires_at = std::time::Instant::now() + SLUG_CACHE_TTL;
284 }
285
286 registry
287}
288
289fn wildcard_slug_lookup<'a>(
299 subdomain: &str,
300 entries: &'a std::collections::HashMap<String, CachedSlugEntry>,
301 wildcard: bool,
302) -> Option<&'a CachedSlugEntry> {
303 let subdomain = subdomain.to_ascii_lowercase();
304
305 entries.get(&subdomain).or_else(|| {
306 if !wildcard {
307 return None;
308 }
309 subdomain
311 .match_indices('.')
312 .map(|(i, _)| &subdomain[i + 1..])
313 .find_map(|candidate| entries.get(candidate))
314 })
315}
316
317#[derive(Debug)]
319enum PrefixMatch<'a> {
320 Worktree(&'a crate::proxy::worktree::WorktreeEntry),
322 Unknown,
325 Ambiguous,
329}
330
331fn match_worktree_prefix<'a>(cached: &'a CachedSlugEntry, prefix: &str) -> PrefixMatch<'a> {
333 if let Some(wt) = cached
334 .worktrees
335 .iter()
336 .find(|w| w.sanitized_branch.eq_ignore_ascii_case(prefix))
337 {
338 return PrefixMatch::Worktree(wt);
339 }
340 if cached
341 .rejected_worktree_prefixes
342 .contains(&prefix.to_ascii_lowercase())
343 {
344 return PrefixMatch::Ambiguous;
345 }
346 PrefixMatch::Unknown
347}
348
349fn strip_dot_suffix_ignore_case(s: &str, suffix: &str) -> Option<String> {
353 let needle_len = suffix.len() + 1;
354 if s.len() <= needle_len {
355 return None;
356 }
357 let split = s.len() - needle_len;
358 if !s.is_char_boundary(split) {
359 return None;
360 }
361 let (head, tail) = s.split_at(split);
362 if tail.starts_with('.') && tail[1..].eq_ignore_ascii_case(suffix) {
363 Some(head.to_string())
364 } else {
365 None
366 }
367}
368
369async fn cached_slug_lookup(subdomain: &str) -> Option<CachedSlugEntry> {
376 let entries = get_cached_slugs().await;
377 wildcard_slug_lookup(subdomain, &entries, settings().proxy.wildcard).cloned()
378}
379
380static AUTO_START_IN_PROGRESS: once_cell::sync::Lazy<
387 tokio::sync::Mutex<std::collections::HashSet<DaemonId>>,
388> = once_cell::sync::Lazy::new(|| tokio::sync::Mutex::new(std::collections::HashSet::new()));
389
390enum ResolveResult {
392 Ready(u16),
395 Starting { slug: String },
397 NotFound,
399 Page {
402 project: String,
403 worktree: Option<String>,
404 daemons: Vec<String>,
405 },
406 Unknown { heading: String, known: Vec<String> },
408 Error(String),
410}
411
412type OnErrorFn = Arc<dyn Fn(&str) + Send + Sync>;
415
416#[derive(Clone)]
417struct ProxyState {
418 client: Arc<Client<HttpConnector, Body>>,
420 tld: String,
422 is_tls: bool,
424 on_error: Option<OnErrorFn>,
426}
427
428pub async fn serve(
436 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
437 cancel: tokio_util::sync::CancellationToken,
438) -> crate::Result<()> {
439 let s = settings();
440 let lan_enabled = s.proxy.lan || !s.proxy.lan_ip.is_empty();
441
442 let effective_tld = if lan_enabled {
443 "local".to_string()
444 } else {
445 s.proxy.tld.clone()
446 };
447
448 let Some(effective_port) = u16::try_from(s.proxy.port).ok().filter(|&p| p > 0) else {
449 let msg = format!(
450 "proxy.port {} is out of valid port range (1-65535), proxy server cannot start",
451 s.proxy.port
452 );
453 let _ = bind_tx.send(Err(msg.clone()));
454 miette::bail!("{msg}");
455 };
456
457 let mut connector = HttpConnector::new();
458 connector.set_connect_timeout(Some(std::time::Duration::from_secs(10)));
462
463 let client = Client::builder(TokioExecutor::new())
464 .pool_idle_timeout(std::time::Duration::from_secs(30))
467 .build(connector);
468
469 let state = ProxyState {
470 client: Arc::new(client),
471 tld: effective_tld.clone(),
472 is_tls: s.proxy.https,
473 on_error: None,
474 };
475
476 let app = Router::new().fallback(proxy_handler).with_state(state);
477
478 let bind_ip: std::net::IpAddr = if lan_enabled && s.proxy.host == "127.0.0.1" {
482 std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)
483 } else {
484 match s.proxy.host.parse() {
485 Ok(ip) => ip,
486 Err(_) => {
487 log::warn!(
488 "proxy.host {:?} is not a valid IP address — falling back to 127.0.0.1. \
489 The proxy will only be reachable on the loopback interface.",
490 s.proxy.host
491 );
492 std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
493 }
494 }
495 };
496 let addr = SocketAddr::from((bind_ip, effective_port));
497
498 if s.proxy.https {
499 serve_https_with_http_fallback(app, addr, &s, effective_port, bind_tx, cancel).await
500 } else {
501 serve_http(app, addr, effective_port, bind_tx, cancel).await
502 }
503}
504
505async fn serve_http(
507 app: Router,
508 addr: SocketAddr,
509 effective_port: u16,
510 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
511 cancel: tokio_util::sync::CancellationToken,
512) -> crate::Result<()> {
513 let listener = match TcpListener::bind(addr).await {
514 Ok(l) => {
515 if settings().proxy.sync_hosts {
516 crate::proxy::hosts::sync_hosts_from_settings();
517 }
518 let _ = bind_tx.send(Ok(()));
519 l
520 }
521 Err(e) => {
522 let msg = bind_error_message(effective_port, &e);
523 let _ = bind_tx.send(Err(msg.clone()));
524 return Err(miette::miette!("{msg}"));
525 }
526 };
527
528 log::info!("Proxy server listening on http://{addr}");
529 if effective_port < 1024 {
530 log::info!(
531 "Note: port {effective_port} is a privileged port. \
532 The supervisor must be started with sudo to bind to this port."
533 );
534 }
535 let shutdown_signal = cancel.clone().cancelled_owned();
536 axum::serve(
537 listener,
538 app.into_make_service_with_connect_info::<SocketAddr>(),
539 )
540 .with_graceful_shutdown(shutdown_signal)
541 .await
542 .map_err(|e| miette::miette!("Proxy server error: {e}"))?;
543 Ok(())
544}
545
546#[cfg(feature = "proxy-tls")]
552async fn serve_https_with_http_fallback(
553 app: Router,
554 addr: SocketAddr,
555 s: &crate::settings::Settings,
556 effective_port: u16,
557 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
558 cancel: tokio_util::sync::CancellationToken,
559) -> crate::Result<()> {
560 use rustls::ServerConfig;
561 use tokio_rustls::TlsAcceptor;
562
563 let (ca_cert_path, ca_key_path) = resolve_tls_paths(s);
564
565 if !ca_cert_path.exists() || !ca_key_path.exists() {
567 generate_ca(&ca_cert_path, &ca_key_path)?;
568 log::info!(
569 "Generated local CA certificate at {}",
570 ca_cert_path.display()
571 );
572 log::info!("To trust the CA in your browser, run: pitchfork proxy trust");
573 }
574
575 let _ = rustls::crypto::ring::default_provider().install_default();
577
578 let resolver = SniCertResolver::new(&ca_cert_path, &ca_key_path)?;
580
581 let mut tls_config = ServerConfig::builder()
582 .with_no_client_auth()
583 .with_cert_resolver(Arc::new(resolver));
584 tls_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
587
588 let acceptor = TlsAcceptor::from(Arc::new(tls_config));
589
590 let listener = match TcpListener::bind(addr).await {
591 Ok(l) => {
592 if settings().proxy.sync_hosts {
593 crate::proxy::hosts::sync_hosts_from_settings();
594 }
595 let _ = bind_tx.send(Ok(()));
596 l
597 }
598 Err(e) => {
599 let msg = bind_error_message(effective_port, &e);
600 let _ = bind_tx.send(Err(msg.clone()));
601 return Err(miette::miette!("{msg}"));
602 }
603 };
604
605 log::info!("Proxy server listening on https://{addr} (HTTP also accepted)");
606 if effective_port < 1024 {
607 log::info!(
608 "Note: port {effective_port} is a privileged port. \
609 The supervisor must be started with sudo to bind to this port."
610 );
611 }
612
613 let redirect_app = Router::new().fallback(redirect_to_https_handler);
615
616 let mut conn_tasks: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
618 loop {
619 while conn_tasks.try_join_next().is_some() {}
622
623 tokio::select! {
624 accept_result = listener.accept() => {
625 let (stream, peer_addr) = match accept_result {
626 Ok(conn) => conn,
627 Err(e) => {
628 log::warn!("Accept error (will retry): {e}");
629 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
630 continue;
631 }
632 };
633
634 let acceptor = acceptor.clone();
635 let app = app
640 .clone()
641 .layer(axum::Extension(axum::extract::ConnectInfo(peer_addr)));
642 let redirect_app = redirect_app.clone();
643
644 conn_tasks.spawn(async move {
645 let mut peek_buf = [0u8; 1];
648 match stream.peek(&mut peek_buf).await {
649 Ok(0) | Err(_) => return,
650 _ => {}
651 }
652
653 if peek_buf[0] == 0x16 {
654 match acceptor.accept(stream).await {
656 Ok(tls_stream) => {
657 let io = hyper_util::rt::TokioIo::new(tls_stream);
658 let svc = hyper_util::service::TowerToHyperService::new(app);
659 if let Err(e) = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new())
660 .serve_connection_with_upgrades(io, svc)
661 .await
662 {
663 log::debug!("Connection error: {e}");
666 }
667 }
668 Err(e) => {
669 log::debug!("TLS handshake error: {e}");
670 }
671 }
672 } else {
673 let io = hyper_util::rt::TokioIo::new(stream);
675 let svc = hyper_util::service::TowerToHyperService::new(redirect_app);
676 let _ = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new())
677 .serve_connection_with_upgrades(io, svc)
678 .await;
679 }
680 });
681
682 while conn_tasks.try_join_next().is_some() {}
683 }
684 _ = cancel.cancelled() => {
685 log::info!("Proxy server shutting down (cancel signal received)");
686 break;
687 }
688 }
689 }
690
691 let drain_timeout = std::time::Duration::from_secs(10);
693 let _ = tokio::time::timeout(drain_timeout, async {
694 while conn_tasks.join_next().await.is_some() {}
695 })
696 .await;
697
698 Ok(())
699}
700
701#[cfg(not(feature = "proxy-tls"))]
703async fn serve_https_with_http_fallback(
704 _app: Router,
705 _addr: SocketAddr,
706 _s: &crate::settings::Settings,
707 _effective_port: u16,
708 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
709 _cancel: tokio_util::sync::CancellationToken,
710) -> crate::Result<()> {
711 let msg = "HTTPS proxy support requires the `proxy-tls` feature.\n\
712 Rebuild pitchfork with: cargo build --features proxy-tls"
713 .to_string();
714 let _ = bind_tx.send(Err(msg.clone()));
715 miette::bail!("{msg}")
716}
717
718#[cfg(feature = "proxy-tls")]
723fn resolve_tls_paths(s: &crate::settings::Settings) -> (std::path::PathBuf, std::path::PathBuf) {
724 let proxy_dir = crate::env::PITCHFORK_STATE_DIR.join("proxy");
725 let resolve = |configured: &str, default: &str| {
726 if configured.is_empty() {
727 proxy_dir.join(default)
728 } else {
729 std::path::PathBuf::from(configured)
730 }
731 };
732 (
733 resolve(&s.proxy.tls_cert, "ca.pem"),
734 resolve(&s.proxy.tls_key, "ca-key.pem"),
735 )
736}
737
738#[cfg(feature = "proxy-tls")]
743pub fn generate_ca(cert_path: &std::path::Path, key_path: &std::path::Path) -> crate::Result<()> {
744 use rcgen::{
745 BasicConstraints, CertificateParams, DistinguishedName, DnType, IsCa, KeyUsagePurpose,
746 };
747
748 if let Some(parent) = cert_path.parent() {
750 std::fs::create_dir_all(parent)
751 .map_err(|e| miette::miette!("Failed to create proxy cert directory: {e}"))?;
752 }
753
754 let mut params = CertificateParams::default();
755 let mut dn = DistinguishedName::new();
756 dn.push(DnType::CommonName, "Pitchfork Local CA");
757 dn.push(DnType::OrganizationName, "Pitchfork");
758 params.distinguished_name = dn;
759 params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
760 params.key_usages = vec![KeyUsagePurpose::KeyCertSign, KeyUsagePurpose::CrlSign];
761
762 let key_pair = rcgen::KeyPair::generate()
763 .map_err(|e| miette::miette!("Failed to generate CA key pair: {e}"))?;
764 let ca_cert = params
765 .self_signed(&key_pair)
766 .map_err(|e| miette::miette!("Failed to self-sign CA certificate: {e}"))?;
767
768 std::fs::write(cert_path, ca_cert.pem()).map_err(|e| {
770 miette::miette!(
771 "Failed to write CA certificate to {}: {e}",
772 cert_path.display()
773 )
774 })?;
775
776 {
780 #[cfg(unix)]
781 {
782 use std::io::Write;
783 use std::os::unix::fs::OpenOptionsExt;
784 std::fs::OpenOptions::new()
785 .write(true)
786 .create(true)
787 .truncate(true)
788 .mode(0o600)
789 .open(key_path)
790 .and_then(|mut f| f.write_all(key_pair.serialize_pem().as_bytes()))
791 .map_err(|e| {
792 miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
793 })?;
794 }
795 #[cfg(not(unix))]
796 {
797 std::fs::write(key_path, key_pair.serialize_pem()).map_err(|e| {
798 miette::miette!("Failed to write CA key to {}: {e}", key_path.display())
799 })?;
800 log::debug!(
801 "CA private key written to {} (file permissions are not restricted \
802 on non-Unix platforms — consider restricting access manually)",
803 key_path.display()
804 );
805 }
806 }
807
808 Ok(())
809}
810
811#[cfg(feature = "proxy-tls")]
831struct SniCertResolver {
832 issuer: rcgen::Issuer<'static, rcgen::KeyPair>,
834 host_certs_dir: std::path::PathBuf,
836 cache: std::sync::Mutex<std::collections::HashMap<String, Arc<rustls::sign::CertifiedKey>>>,
838 pending: std::sync::Mutex<std::collections::HashSet<String>>,
842 pending_cv: std::sync::Condvar,
844}
845
846#[cfg(feature = "proxy-tls")]
847impl std::fmt::Debug for SniCertResolver {
848 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
849 f.debug_struct("SniCertResolver").finish_non_exhaustive()
850 }
851}
852
853#[cfg(feature = "proxy-tls")]
854impl SniCertResolver {
855 fn new(ca_cert_path: &std::path::Path, ca_key_path: &std::path::Path) -> crate::Result<Self> {
857 let ca_key_pem = std::fs::read_to_string(ca_key_path)
858 .map_err(|e| miette::miette!("Failed to read CA key {}: {e}", ca_key_path.display()))?;
859 let ca_cert_pem = std::fs::read_to_string(ca_cert_path).map_err(|e| {
860 miette::miette!("Failed to read CA cert {}: {e}", ca_cert_path.display())
861 })?;
862
863 if !ca_cert_pem.contains("BEGIN CERTIFICATE") {
865 miette::bail!("CA cert file does not contain a valid PEM certificate");
866 }
867
868 let ca_key = rcgen::KeyPair::from_pem(&ca_key_pem)
869 .map_err(|e| miette::miette!("Failed to parse CA key: {e}"))?;
870
871 let issuer = rcgen::Issuer::from_ca_cert_pem(&ca_cert_pem, ca_key)
873 .map_err(|e| miette::miette!("Failed to parse CA cert: {e}"))?;
874
875 let host_certs_dir = ca_cert_path
877 .parent()
878 .unwrap_or(std::path::Path::new("."))
879 .join("host-certs");
880 std::fs::create_dir_all(&host_certs_dir)
881 .map_err(|e| miette::miette!("Failed to create host-certs dir: {e}"))?;
882
883 Ok(Self {
884 issuer,
885 host_certs_dir,
886 cache: std::sync::Mutex::new(std::collections::HashMap::new()),
887 pending: std::sync::Mutex::new(std::collections::HashSet::new()),
888 pending_cv: std::sync::Condvar::new(),
889 })
890 }
891
892 fn get_or_create(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
915 {
917 let cache = self.cache.lock().ok()?;
918 if let Some(ck) = cache.get(domain) {
919 return Some(Arc::clone(ck));
920 }
921 } loop {
934 {
935 let mut pending = self.pending.lock().ok()?;
936 if pending.contains(domain) {
937 pending = self.pending_cv.wait(pending).ok()?;
939 drop(pending);
941 } else {
942 pending.insert(domain.to_string());
944 break;
945 }
946 } {
952 let cache = self.cache.lock().ok()?;
953 if let Some(ck) = cache.get(domain) {
954 return Some(Arc::clone(ck));
955 }
956 } } let result = self.get_or_create_inner(domain);
960
961 {
967 let mut pending = match self.pending.lock() {
968 Ok(g) => g,
969 Err(e) => e.into_inner(),
970 };
971 pending.remove(domain);
972 self.pending_cv.notify_all();
973 }
974
975 result
976 }
977
978 fn get_or_create_inner(&self, domain: &str) -> Option<Arc<rustls::sign::CertifiedKey>> {
980 let safe_name = domain.replace('.', "_").replace('*', "wildcard");
981 let disk_path = self.host_certs_dir.join(format!("{safe_name}.pem"));
982
983 if disk_path.exists() {
985 if let Ok(ck) = self.load_from_disk(&disk_path) {
986 let ck = Arc::new(ck);
987 if let Ok(mut cache) = self.cache.lock() {
988 cache.insert(domain.to_string(), Arc::clone(&ck));
989 }
990 return Some(ck);
991 }
992 let _ = std::fs::remove_file(&disk_path);
994 }
995
996 let ck = self.sign_for_domain(domain).ok()?;
998
999 let ck = Arc::new(ck);
1000 if let Ok(mut cache) = self.cache.lock() {
1001 cache.insert(domain.to_string(), Arc::clone(&ck));
1002 }
1003 Some(ck)
1004 }
1005
1006 fn load_from_disk(&self, path: &std::path::Path) -> crate::Result<rustls::sign::CertifiedKey> {
1011 use rustls::pki_types::CertificateDer;
1012 use rustls_pemfile::{certs, private_key};
1013
1014 let pem = std::fs::read_to_string(path)
1015 .map_err(|e| miette::miette!("Failed to read disk cert {}: {e}", path.display()))?;
1016
1017 let cert_ders: Vec<CertificateDer<'static>> = certs(&mut pem.as_bytes())
1018 .collect::<Result<Vec<_>, _>>()
1019 .map_err(|e| miette::miette!("Failed to parse certs from {}: {e}", path.display()))?;
1020
1021 if cert_ders.is_empty() {
1022 miette::bail!("No certificates found in {}", path.display());
1023 }
1024
1025 {
1027 let (_, cert) = x509_parser::parse_x509_certificate(&cert_ders[0]).map_err(|e| {
1028 miette::miette!("Failed to parse certificate from {}: {e}", path.display())
1029 })?;
1030 use chrono::Utc;
1031 let now_ts = Utc::now().timestamp();
1032 let not_after_ts = cert.validity().not_after.timestamp();
1033 if not_after_ts < now_ts {
1034 miette::bail!(
1035 "Cached certificate at {} has expired — will regenerate",
1036 path.display()
1037 );
1038 }
1039 }
1040
1041 let key_der = private_key(&mut pem.as_bytes())
1042 .map_err(|e| miette::miette!("Failed to parse key from {}: {e}", path.display()))?
1043 .ok_or_else(|| miette::miette!("No private key found in {}", path.display()))?;
1044
1045 let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
1046 .map_err(|e| miette::miette!("Failed to create signing key from disk: {e}"))?;
1047
1048 Ok(rustls::sign::CertifiedKey::new(cert_ders, signing_key))
1049 }
1050
1051 fn sign_for_domain(&self, domain: &str) -> crate::Result<rustls::sign::CertifiedKey> {
1059 use rcgen::date_time_ymd;
1060 use rcgen::{CertificateParams, DistinguishedName, DnType, SanType};
1061 use rustls::pki_types::CertificateDer;
1062 use rustls_pemfile::private_key;
1063
1064 let mut params = CertificateParams::default();
1065 let mut dn = DistinguishedName::new();
1066 dn.push(DnType::CommonName, domain);
1067 params.distinguished_name = dn;
1068
1069 {
1071 use chrono::{Datelike, Duration, Utc};
1072 let yesterday = Utc::now() - Duration::days(1);
1073 let expiry = Utc::now() + Duration::days(397);
1076 params.not_before = date_time_ymd(
1077 yesterday.year(),
1078 yesterday.month() as u8,
1079 yesterday.day() as u8,
1080 );
1081 params.not_after =
1082 date_time_ymd(expiry.year(), expiry.month() as u8, expiry.day() as u8);
1083 }
1084
1085 let mut sans =
1087 vec![SanType::DnsName(domain.to_string().try_into().map_err(
1088 |e| miette::miette!("Invalid domain name '{domain}': {e}"),
1089 )?)];
1090 if let Some(dot_pos) = domain.find('.') {
1092 let parent = &domain[dot_pos + 1..];
1093 if parent.contains('.') {
1095 let wildcard = format!("*.{parent}");
1096 if let Ok(wc) = wildcard.try_into() {
1097 sans.push(SanType::DnsName(wc));
1098 }
1099 }
1100 }
1101 params.subject_alt_names = sans;
1102
1103 let leaf_key = rcgen::KeyPair::generate()
1104 .map_err(|e| miette::miette!("Failed to generate leaf key: {e}"))?;
1105 let leaf_cert = params
1106 .signed_by(&leaf_key, &self.issuer)
1107 .map_err(|e| miette::miette!("Failed to sign leaf cert for '{domain}': {e}"))?;
1108
1109 let cert_der = CertificateDer::from(leaf_cert.der().to_vec());
1111 let key_pem = leaf_key.serialize_pem();
1112 let key_der = private_key(&mut key_pem.as_bytes())
1113 .map_err(|e| miette::miette!("Failed to parse leaf key PEM: {e}"))?
1114 .ok_or_else(|| miette::miette!("No private key found in generated PEM"))?;
1115
1116 let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der)
1117 .map_err(|e| miette::miette!("Failed to create signing key: {e}"))?;
1118
1119 let safe_name = domain.replace('.', "_").replace('*', "wildcard");
1122 let disk_path = self.host_certs_dir.join(format!("{safe_name}.pem"));
1123 let combined_pem = format!("{}{}", leaf_cert.pem(), key_pem);
1124 {
1125 #[cfg(unix)]
1126 {
1127 use std::io::Write;
1128 use std::os::unix::fs::OpenOptionsExt;
1129 if let Err(e) = std::fs::OpenOptions::new()
1130 .write(true)
1131 .create(true)
1132 .truncate(true)
1133 .mode(0o600)
1134 .open(&disk_path)
1135 .and_then(|mut f| f.write_all(combined_pem.as_bytes()))
1136 {
1137 log::warn!(
1138 "Failed to persist cert for '{domain}' to {}: {e}",
1139 disk_path.display()
1140 );
1141 }
1142 }
1143 #[cfg(not(unix))]
1144 {
1145 if let Err(e) = std::fs::write(&disk_path, combined_pem) {
1146 log::warn!(
1147 "Failed to persist cert for '{domain}' to {}: {e}",
1148 disk_path.display()
1149 );
1150 } else {
1151 log::debug!(
1152 "Leaf cert for '{domain}' written to {} (file permissions are not \
1153 restricted on non-Unix platforms — consider restricting access manually)",
1154 disk_path.display()
1155 );
1156 }
1157 }
1158 }
1159
1160 Ok(rustls::sign::CertifiedKey::new(vec![cert_der], signing_key))
1161 }
1162}
1163
1164#[cfg(feature = "proxy-tls")]
1165impl rustls::server::ResolvesServerCert for SniCertResolver {
1166 fn resolve(
1167 &self,
1168 client_hello: rustls::server::ClientHello<'_>,
1169 ) -> Option<Arc<rustls::sign::CertifiedKey>> {
1170 let domain = client_hello.server_name()?;
1171 self.get_or_create(domain)
1172 }
1173}
1174
1175fn get_request_host(req: &Request) -> Option<String> {
1181 let authority = req
1183 .uri()
1184 .authority()
1185 .map(|a| a.as_str().to_string())
1186 .filter(|s| !s.is_empty());
1187
1188 authority.or_else(|| {
1189 req.headers()
1190 .get(HOST)
1191 .and_then(|h| h.to_str().ok())
1192 .map(str::to_string)
1193 })
1194}
1195
1196fn join_cookie_fields(headers: &mut HeaderMap) {
1202 let fields: Vec<&[u8]> = headers
1203 .get_all(COOKIE)
1204 .iter()
1205 .map(HeaderValue::as_bytes)
1206 .collect();
1207 if fields.len() < 2 {
1208 return;
1209 }
1210
1211 let joined = HeaderValue::from_bytes(&fields.join(b"; ".as_slice()))
1212 .expect("valid header values joined with \"; \" form a valid header value");
1213 headers.insert(COOKIE, joined);
1214}
1215
1216fn inject_forwarded_headers(req: &mut Request, is_tls: bool, host_header: &str) {
1227 let remote_addr = req
1228 .extensions()
1229 .get::<axum::extract::ConnectInfo<SocketAddr>>()
1230 .map(|ci| ci.0.ip().to_string())
1231 .unwrap_or_else(|| "127.0.0.1".to_string());
1232
1233 let proto = if is_tls { "https" } else { "http" };
1234 let default_port = if is_tls { "443" } else { "80" };
1235
1236 let forwarded_for = remote_addr.clone();
1239 let forwarded_proto = proto.to_string();
1240 let forwarded_host = host_header.to_string();
1241 let forwarded_port = host_header
1242 .rsplit_once(':')
1243 .map(|(_, port)| port.to_string())
1244 .unwrap_or_else(|| default_port.to_string());
1245
1246 for name in [
1252 "x-forwarded-for",
1253 "x-forwarded-proto",
1254 "x-forwarded-host",
1255 "x-forwarded-port",
1256 "forwarded",
1257 ] {
1258 if let Ok(header_name) = axum::http::HeaderName::from_bytes(name.as_bytes()) {
1259 req.headers_mut().remove(&header_name);
1260 }
1261 }
1262
1263 let headers = [
1264 ("x-forwarded-for", forwarded_for),
1265 ("x-forwarded-proto", forwarded_proto),
1266 ("x-forwarded-host", forwarded_host),
1267 ("x-forwarded-port", forwarded_port),
1268 ];
1269
1270 for (name, value) in headers {
1271 if let Ok(v) = HeaderValue::from_str(&value) {
1272 let header_name = axum::http::HeaderName::from_static(name);
1273 req.headers_mut().insert(header_name, v);
1274 }
1275 }
1276}
1277
1278async fn proxy_handler(State(state): State<ProxyState>, mut req: Request) -> Response {
1283 let Some(raw_host) = get_request_host(&req) else {
1285 return error_response(StatusCode::BAD_REQUEST, "Missing Host header");
1286 };
1287 let host = if raw_host.starts_with('[') {
1291 raw_host
1293 .split("]:")
1294 .next()
1295 .unwrap_or(&raw_host)
1296 .trim_start_matches('[')
1297 .trim_end_matches(']')
1298 .to_string()
1299 } else {
1300 raw_host.split(':').next().unwrap_or(&raw_host).to_string()
1302 };
1303
1304 let is_from_pitchfork = req.headers().contains_key(PROXY_HOPS_HEADER);
1315 let hops: u64 = if is_from_pitchfork {
1316 req.headers()
1317 .get(PROXY_HOPS_HEADER)
1318 .and_then(|v| v.to_str().ok())
1319 .and_then(|s| s.parse().ok())
1320 .unwrap_or(0)
1321 } else {
1322 0
1324 };
1325 if hops >= MAX_PROXY_HOPS {
1326 return error_response(
1327 StatusCode::LOOP_DETECTED,
1328 &format!(
1329 "Loop detected for '{host}': request has passed through the proxy {hops} times.\n\
1330 This usually means a backend is proxying back through pitchfork without rewriting \n\
1331 the Host header. If you use Vite/webpack proxy, set changeOrigin: true."
1332 ),
1333 );
1334 }
1335
1336 let local_client = is_local_client(&req);
1337
1338 let target_port = if let Some(subdomain) = strip_tld(&host, &state.tld) {
1340 if subdomain.eq_ignore_ascii_case("pitchfork") {
1341 crate::web::port()
1342 } else {
1343 None
1344 }
1345 } else {
1346 None
1347 };
1348
1349 let target_port = if let Some(port) = target_port {
1350 port
1351 } else {
1352 match resolve_target(&host, &state.tld).await {
1353 ResolveResult::Ready(port) => port,
1354 ResolveResult::Starting { slug } => {
1355 return starting_html_response(&slug, &raw_host);
1356 }
1357 ResolveResult::Page {
1358 project,
1359 worktree,
1360 daemons,
1361 } => {
1362 if !local_client {
1366 return unknown_host_response(&host, "Not found", &[]);
1367 }
1368 return page_placeholder_response(
1369 &project,
1370 worktree.as_deref(),
1371 &daemons,
1372 &state.tld,
1373 &host_port_suffix(&raw_host),
1374 );
1375 }
1376 ResolveResult::Unknown { heading, known } => {
1377 return unknown_host_response(
1380 &host,
1381 if local_client { &heading } else { "Not found" },
1382 if local_client { &known } else { &[] },
1383 );
1384 }
1385 ResolveResult::NotFound => {
1386 return error_response(
1387 StatusCode::BAD_GATEWAY,
1388 &format!(
1389 "No daemon found for host '{host}'.\n\
1390 A daemon is reachable once it configures a `port` and its project is \
1391 known to pitchfork; run `pitchfork proxy status` to see the hostnames \
1392 it serves.\n\
1393 Expected format: <daemon>.<project>.{tld}",
1394 tld = state.tld
1395 ),
1396 );
1397 }
1398 ResolveResult::Error(msg) => {
1399 if local_client {
1400 return error_response(StatusCode::BAD_GATEWAY, &msg);
1401 }
1402 log::warn!("Refused '{host}' for a non-local client: {msg}");
1404 return error_response(
1405 StatusCode::BAD_GATEWAY,
1406 &format!("'{host}' is not available."),
1407 );
1408 }
1409 }
1410 };
1411 let path_and_query = req
1413 .uri()
1414 .path_and_query()
1415 .map(|pq| pq.as_str())
1416 .unwrap_or("/");
1417
1418 let forward_uri = match Uri::builder()
1419 .scheme("http")
1420 .authority(format!("localhost:{target_port}"))
1421 .path_and_query(path_and_query)
1422 .build()
1423 {
1424 Ok(uri) => uri,
1425 Err(e) => {
1426 return error_response(
1427 StatusCode::INTERNAL_SERVER_ERROR,
1428 &format!("Failed to build forward URI: {e}"),
1429 );
1430 }
1431 };
1432
1433 *req.uri_mut() = forward_uri;
1435 req.headers_mut().insert(
1436 HOST,
1437 HeaderValue::from_str(&format!("localhost:{target_port}"))
1438 .unwrap_or_else(|_| HeaderValue::from_static("localhost")),
1439 );
1440
1441 inject_forwarded_headers(&mut req, state.is_tls, &raw_host);
1443
1444 if let Ok(v) = HeaderValue::from_str(&(hops + 1).to_string()) {
1446 req.headers_mut()
1447 .insert(axum::http::HeaderName::from_static(PROXY_HOPS_HEADER), v);
1448 }
1449
1450 let pseudo_headers: Vec<_> = req
1455 .headers()
1456 .keys()
1457 .filter(|k| k.as_str().starts_with(':'))
1458 .cloned()
1459 .collect();
1460 for key in pseudo_headers {
1461 req.headers_mut().remove(&key);
1462 }
1463
1464 join_cookie_fields(req.headers_mut());
1465
1466 *req.version_mut() = axum::http::Version::HTTP_11;
1471
1472 let client_upgrade = hyper::upgrade::on(&mut req);
1474
1475 let result = match tokio::time::timeout(
1483 std::time::Duration::from_secs(120),
1484 state.client.request(req),
1485 )
1486 .await
1487 {
1488 Ok(r) => r,
1489 Err(_elapsed) => {
1490 let msg = format!(
1491 "Request to daemon on port {target_port} timed out after 120 s.\n\
1492 The daemon accepted the connection but did not respond in time."
1493 );
1494 log::warn!("{msg}");
1495 if let Some(ref on_error) = state.on_error {
1496 on_error(&msg);
1497 }
1498 return error_response(StatusCode::GATEWAY_TIMEOUT, &msg);
1499 }
1500 };
1501 match result {
1502 Ok(mut resp) => {
1503 let backend_upgrade = hyper::upgrade::on(&mut resp);
1505 let (mut parts, body) = resp.into_parts();
1506
1507 parts.headers.insert(
1509 axum::http::HeaderName::from_static(PITCHFORK_HEADER),
1510 HeaderValue::from_static("1"),
1511 );
1512
1513 parts.headers.remove(PROXY_HOPS_HEADER);
1515
1516 if state.is_tls && parts.status != StatusCode::SWITCHING_PROTOCOLS {
1521 for h in HOP_BY_HOP_HEADERS {
1522 if let Ok(name) = axum::http::HeaderName::from_bytes(h.as_bytes()) {
1523 parts.headers.remove(&name);
1524 }
1525 }
1526 }
1527
1528 if parts.status == StatusCode::SWITCHING_PROTOCOLS {
1530 tokio::spawn(async move {
1535 if let (Ok(client_upgraded), Ok(backend_upgraded)) =
1536 (client_upgrade.await, backend_upgrade.await)
1537 {
1538 let mut client_io = hyper_util::rt::TokioIo::new(client_upgraded);
1539 let mut backend_io = hyper_util::rt::TokioIo::new(backend_upgraded);
1540 let _ =
1548 tokio::io::copy_bidirectional(&mut client_io, &mut backend_io).await;
1549 }
1550 });
1551 return Response::from_parts(parts, Body::empty());
1552 }
1553
1554 Response::from_parts(parts, Body::new(body))
1557 }
1558 Err(e) => {
1559 let msg = format!(
1560 "Failed to connect to daemon on port {target_port}: {e}\n\
1561 The daemon may have stopped or is not yet ready."
1562 );
1563 if let Some(ref on_error) = state.on_error {
1564 on_error(&msg);
1565 } else {
1566 log::warn!("{msg}");
1567 }
1568 error_response(StatusCode::BAD_GATEWAY, &msg)
1569 }
1570 }
1571}
1572
1573async fn resolve_target(host: &str, tld: &str) -> ResolveResult {
1592 let Some(subdomain) = strip_tld(host, tld) else {
1593 return ResolveResult::NotFound;
1594 };
1595
1596 let cached = cached_slug_lookup(&subdomain).await.filter(|cached| {
1597 if crate::proxy::hostname::hostname_fits(&cached.slug) {
1601 return true;
1602 }
1603 crate::proxy::hostname::warn_once(&format!(
1604 "Slug '{}' plus the configured proxy.tld is over the DNS length limit, \
1605 so it is not routed.",
1606 cached.slug
1607 ));
1608 false
1609 });
1610 let Some(cached) = cached else {
1611 return resolve_registry_target(&subdomain).await;
1614 };
1615
1616 let (expected_namespace, worktree_dir) = if !subdomain.eq_ignore_ascii_case(&cached.slug) {
1620 let prefix = strip_dot_suffix_ignore_case(&subdomain, &cached.slug);
1621 match prefix {
1622 Some(ref p) => match match_worktree_prefix(&cached, p) {
1623 PrefixMatch::Worktree(wt) => {
1624 let ns = wt.namespace.clone().or_else(|| {
1625 log::warn!(
1626 "Worktree '{}' has no cached namespace; \
1627 falling back to parent slug namespace.",
1628 wt.path.display()
1629 );
1630 cached.namespace.clone()
1631 });
1632 (ns, Some(wt.path.clone()))
1633 }
1634 PrefixMatch::Ambiguous => {
1635 return ResolveResult::Error(format!(
1636 "'{host}' is ambiguous: more than one branch or workspace of '{slug}' \
1637 sanitizes to the prefix '{p}', and host names are case-insensitive.\n\
1638 Rename one of them so the prefixes differ by more than case, then \
1639 reload.\n\
1640 The supervisor log lists the colliding branches.",
1641 slug = cached.slug,
1642 ));
1643 }
1644 PrefixMatch::Unknown => (cached.namespace.clone(), None),
1645 },
1646 None => (cached.namespace.clone(), None),
1647 }
1648 } else {
1649 (cached.namespace.clone(), None)
1650 };
1651
1652 let daemon_name = &cached.daemon_name;
1653
1654 let daemons = {
1655 let state_file = SUPERVISOR.state_file.lock().await;
1656 state_file.daemons.clone()
1657 };
1658
1659 let running_matches: Vec<(&DaemonId, &crate::daemon::Daemon)> = daemons
1660 .iter()
1661 .filter(|(id, d)| {
1662 id.name() == daemon_name
1663 && d.status.is_running()
1664 && match &expected_namespace {
1665 Some(ns) => id.namespace() == ns,
1666 None => true,
1667 }
1668 })
1669 .collect();
1670
1671 match running_matches.as_slice() {
1672 [] => {
1673 try_auto_start(
1674 &cached.slug,
1675 &cached,
1676 worktree_dir.as_deref(),
1677 expected_namespace.as_deref(),
1678 )
1679 .await
1680 }
1681 [(_, d)] => {
1682 if let Some(port) = d.active_port.or_else(|| d.resolved_port.first().copied()) {
1683 ResolveResult::Ready(port)
1684 } else {
1685 ResolveResult::NotFound
1686 }
1687 }
1688 _ => {
1689 let d = running_matches[0].1;
1690 if let Some(port) = d.active_port.or_else(|| d.resolved_port.first().copied()) {
1691 ResolveResult::Ready(port)
1692 } else {
1693 ResolveResult::NotFound
1694 }
1695 }
1696 }
1697}
1698
1699struct AutoStartGuard {
1706 daemon_id: DaemonId,
1707}
1708
1709impl Drop for AutoStartGuard {
1710 fn drop(&mut self) {
1711 let daemon_id = self.daemon_id.clone();
1712 tokio::spawn(async move {
1716 AUTO_START_IN_PROGRESS.lock().await.remove(&daemon_id);
1717 });
1718 }
1719}
1720
1721async fn try_auto_start(
1732 slug: &str,
1733 cached: &CachedSlugEntry,
1734 worktree_dir: Option<&std::path::Path>,
1735 expected_namespace: Option<&str>,
1736) -> ResolveResult {
1737 let s = settings();
1738 if !s.proxy.auto_start {
1739 return ResolveResult::NotFound;
1740 }
1741
1742 let ns = expected_namespace
1743 .map(|s| s.to_string())
1744 .or_else(|| cached.namespace.clone())
1745 .unwrap_or_else(|| "global".to_string());
1746 let daemon_id = match DaemonId::try_new(&ns, &cached.daemon_name) {
1747 Ok(id) => id,
1748 Err(_) => return ResolveResult::NotFound,
1749 };
1750
1751 {
1752 let mut in_progress = AUTO_START_IN_PROGRESS.lock().await;
1753 if !in_progress.insert(daemon_id.clone()) {
1754 return ResolveResult::Starting {
1755 slug: slug.to_string(),
1756 };
1757 }
1758 }
1759
1760 let _guard = AutoStartGuard {
1761 daemon_id: daemon_id.clone(),
1762 };
1763
1764 let timeout = s.proxy_auto_start_timeout();
1765
1766 match tokio::time::timeout(
1767 timeout,
1768 try_auto_start_inner(slug, cached, &daemon_id, worktree_dir),
1769 )
1770 .await
1771 {
1772 Ok(result) => result,
1773 Err(_elapsed) => {
1774 log::warn!("Auto-start: total timeout ({timeout:?}) exceeded for daemon {daemon_id}");
1775 ResolveResult::Error(format!(
1776 "Auto-start for '{daemon_id}' timed out after {timeout:?}.\n\
1777 The daemon did not become ready and bind a port within the configured \
1778 proxy_auto_start_timeout.\n\
1779 Increase the timeout or check the daemon's logs for slow startup."
1780 ))
1781 }
1782 }
1783}
1784
1785async fn try_auto_start_inner(
1789 slug: &str,
1790 cached: &CachedSlugEntry,
1791 daemon_id: &DaemonId,
1792 worktree_dir: Option<&std::path::Path>,
1793) -> ResolveResult {
1794 let config_dir = worktree_dir.unwrap_or(&cached.dir);
1795
1796 let pt = match crate::pitchfork_toml::PitchforkToml::all_merged_from(config_dir) {
1797 Ok(pt) => pt,
1798 Err(e) => {
1799 log::warn!(
1800 "Auto-start: failed to load config from {}: {e}",
1801 config_dir.display()
1802 );
1803 return ResolveResult::NotFound;
1804 }
1805 };
1806
1807 let mut daemon_config = match pt.daemons.get(daemon_id) {
1808 Some(cfg) => cfg.clone(),
1809 None => {
1810 log::debug!(
1811 "Auto-start: daemon {daemon_id} not found in config at {}",
1812 config_dir.display()
1813 );
1814 return ResolveResult::NotFound;
1815 }
1816 };
1817
1818 let rendered = {
1822 let id = daemon_id.clone();
1823 let mut config = daemon_config.clone();
1824 tokio::task::spawn_blocking(move || {
1825 crate::ipc::batch::render_daemon_config(&id, &mut config, &pt).map(|()| config)
1826 })
1827 .await
1828 };
1829 daemon_config = match rendered {
1830 Ok(Ok(config)) => config,
1831 Ok(Err(e)) => {
1832 log::warn!("Auto-start: failed to render templates for {daemon_id}: {e}");
1833 return ResolveResult::Error(format!("Failed to render templates: {e}"));
1834 }
1835 Err(e) => {
1836 log::warn!("Auto-start: template rendering task failed for {daemon_id}: {e}");
1837 return ResolveResult::Error(format!("Failed to render templates: {e}"));
1838 }
1839 };
1840
1841 let opts = crate::ipc::batch::StartOptions {
1842 quiet: true,
1843 ..crate::ipc::batch::StartOptions::default()
1844 };
1845 let mut run_opts =
1846 match crate::ipc::batch::build_run_options(daemon_id, &daemon_config, Some(&opts)).await {
1847 Ok(o) => o,
1848 Err(e) => {
1849 log::warn!("Auto-start: failed to build run options for {daemon_id}: {e}");
1850 return ResolveResult::Error(format!("Failed to build run options: {e}"));
1851 }
1852 };
1853
1854 if run_opts.dir.0.as_os_str().is_empty() {
1857 run_opts.dir = crate::config_types::Dir(config_dir.to_path_buf());
1858 }
1859
1860 log::info!("Auto-start: starting daemon {daemon_id} for slug '{slug}'");
1861
1862 let run_result = SUPERVISOR.run(run_opts).await;
1863
1864 if let Err(e) = run_result {
1865 log::warn!("Auto-start: failed to start daemon {daemon_id}: {e}");
1866 return ResolveResult::Error(format!("Failed to start daemon: {e}"));
1867 }
1868
1869 let poll_interval = std::time::Duration::from_millis(250);
1870
1871 loop {
1872 let daemons = {
1873 let sf = SUPERVISOR.state_file.lock().await;
1874 sf.daemons.clone()
1875 };
1876
1877 if let Some(d) = daemons.get(daemon_id) {
1878 if d.status.is_running() {
1879 if let Some(port) = d.active_port.or_else(|| d.resolved_port.first().copied()) {
1880 log::info!("Auto-start: daemon {daemon_id} is ready on port {port}");
1881 return ResolveResult::Ready(port);
1882 }
1883 } else {
1884 log::warn!(
1885 "Auto-start: daemon {daemon_id} is no longer running (status: {})",
1886 d.status
1887 );
1888 return ResolveResult::Error(format!(
1889 "Daemon '{daemon_id}' started but exited unexpectedly.\n\
1890 Check its logs for errors."
1891 ));
1892 }
1893 } else {
1894 log::warn!("Auto-start: daemon {daemon_id} not found in state file after start");
1895 return ResolveResult::Error(format!(
1896 "Daemon '{daemon_id}' started but disappeared from the state file.\n\
1897 Check its logs for errors."
1898 ));
1899 }
1900
1901 tokio::time::sleep(poll_interval).await;
1902 }
1903}
1904
1905async fn resolve_registry_target(subdomain: &str) -> ResolveResult {
1910 let registry = get_cached_host_registry().await;
1911 if !crate::proxy::hostname::hostname_fits(subdomain) {
1912 return ResolveResult::Unknown {
1914 heading: "Host name too long".to_string(),
1915 known: registry.project_labels(),
1916 };
1917 }
1918 match registry.resolve(subdomain, settings().proxy.wildcard) {
1919 crate::proxy::hostname::HostTarget::Daemon {
1920 ref dir,
1921 ref namespace,
1922 ref daemon,
1923 ..
1924 } => {
1925 let per_checkout = registry.shares_daemon_id(namespace, daemon);
1930 resolve_registry_daemon(subdomain, dir, namespace, daemon, per_checkout).await
1931 }
1932 crate::proxy::hostname::HostTarget::ProjectPage { project } => {
1933 let daemons = registry
1934 .projects
1935 .get(&project)
1936 .map(|p| p.primary.labels())
1937 .unwrap_or_default();
1938 ResolveResult::Page {
1939 project,
1940 worktree: None,
1941 daemons,
1942 }
1943 }
1944 crate::proxy::hostname::HostTarget::WorktreePage { project, worktree } => {
1945 let daemons = registry
1946 .projects
1947 .get(&project)
1948 .and_then(|p| p.worktrees.get(&worktree))
1949 .map(|c| c.labels())
1950 .unwrap_or_default();
1951 ResolveResult::Page {
1952 project,
1953 worktree: Some(worktree),
1954 daemons,
1955 }
1956 }
1957 crate::proxy::hostname::HostTarget::UnknownProject { known } => ResolveResult::Unknown {
1958 heading: "Unknown project".to_string(),
1959 known,
1960 },
1961 crate::proxy::hostname::HostTarget::UnknownDaemon {
1962 project,
1963 worktree,
1964 known,
1965 } => ResolveResult::Unknown {
1966 heading: match worktree {
1967 Some(wt) => format!("Unknown daemon in '{wt}' of project '{project}'"),
1968 None => format!("Unknown daemon in project '{project}'"),
1969 },
1970 known,
1971 },
1972 }
1973}
1974
1975async fn resolve_registry_daemon(
1982 host: &str,
1983 dir: &std::path::Path,
1984 namespace: &str,
1985 daemon: &str,
1986 per_checkout: bool,
1987) -> ResolveResult {
1988 let daemons = {
1989 let state_file = SUPERVISOR.state_file.lock().await;
1990 state_file.daemons.clone()
1991 };
1992
1993 let mut matches: Vec<crate::daemon::Daemon> = daemons
1994 .iter()
1995 .filter(|(id, d)| {
1996 id.name() == daemon && id.namespace() == namespace && d.status.is_running()
1997 })
1998 .map(|(_, d)| d.clone())
1999 .collect();
2000 matches = sort_by_checkout(matches, dir).await;
2003
2004 if let Some(d) = matches.first() {
2005 if per_checkout && !runs_in_checkout(d.clone(), dir).await {
2008 return ResolveResult::Error(format!(
2009 "'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
2010 running from {}.\n\
2011 These checkouts share the namespace '{namespace}', so pitchfork cannot run \
2012 both copies at once.\n\
2013 Give each checkout its own top-level `namespace`, or stop the other one first.",
2014 dir.display(),
2015 d.dir
2016 .as_deref()
2017 .map(|p| p.display().to_string())
2018 .unwrap_or_else(|| "an unknown directory".to_string()),
2019 ));
2020 }
2021 return match d.active_port.or_else(|| d.resolved_port.first().copied()) {
2022 Some(port) => ResolveResult::Ready(port),
2023 None => ResolveResult::NotFound,
2024 };
2025 }
2026
2027 let cached = CachedSlugEntry {
2028 slug: host.to_string(),
2029 namespace: Some(namespace.to_string()),
2030 daemon_name: daemon.to_string(),
2031 dir: dir.to_path_buf(),
2032 worktrees: vec![],
2033 rejected_worktree_prefixes: std::collections::HashSet::new(),
2034 };
2035 let result = try_auto_start(host, &cached, None, Some(namespace)).await;
2036
2037 if per_checkout && let ResolveResult::Ready(_) = result {
2041 let started = {
2042 let state_file = SUPERVISOR.state_file.lock().await;
2043 state_file
2044 .daemons
2045 .iter()
2046 .find(|(id, _)| id.name() == daemon && id.namespace() == namespace)
2047 .map(|(_, d)| d.clone())
2048 };
2049 if let Some(d) = started
2050 && !runs_in_checkout(d.clone(), dir).await
2051 {
2052 return ResolveResult::Error(format!(
2053 "'{host}' belongs to the checkout at {}, but daemon '{namespace}/{daemon}' is \
2054 running from {}.\n\
2055 These checkouts share the namespace '{namespace}', so pitchfork cannot run \
2056 both copies at once.\n\
2057 Give each checkout its own top-level `namespace`, or stop the other one first.",
2058 dir.display(),
2059 d.dir
2060 .as_deref()
2061 .map(|p| p.display().to_string())
2062 .unwrap_or_else(|| "an unknown directory".to_string()),
2063 ));
2064 }
2065 }
2066
2067 result
2068}
2069
2070fn is_local_client(req: &Request) -> bool {
2077 req.extensions()
2080 .get::<axum::extract::ConnectInfo<SocketAddr>>()
2081 .is_some_and(|ci| ci.0.ip().is_loopback())
2082}
2083
2084fn daemon_runs_in(daemon: &crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
2093 daemon
2094 .dir
2095 .as_deref()
2096 .is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
2097}
2098
2099async fn runs_in_checkout(daemon: crate::daemon::Daemon, checkout: &std::path::Path) -> bool {
2101 let checkout = checkout.to_path_buf();
2102 tokio::task::spawn_blocking(move || daemon_runs_in(&daemon, &checkout))
2103 .await
2104 .unwrap_or(false)
2105}
2106
2107async fn sort_by_checkout(
2113 daemons: Vec<crate::daemon::Daemon>,
2114 checkout: &std::path::Path,
2115) -> Vec<crate::daemon::Daemon> {
2116 if daemons.len() < 2 {
2117 return daemons;
2118 }
2119 let dirs: Vec<Option<std::path::PathBuf>> = daemons.iter().map(|d| d.dir.clone()).collect();
2120 let checkout = checkout.to_path_buf();
2121 let here = tokio::task::spawn_blocking(move || {
2122 dirs.iter()
2123 .map(|dir| {
2124 dir.as_deref()
2125 .is_some_and(|d| crate::proxy::hostname::checkout_root_of(d) == checkout)
2126 })
2127 .collect::<Vec<bool>>()
2128 })
2129 .await;
2130
2131 match here {
2132 Ok(here) => {
2133 let mut ordered: Vec<(bool, crate::daemon::Daemon)> =
2134 here.into_iter().zip(daemons).collect();
2135 ordered.sort_by_key(|(here, _)| !here);
2136 ordered.into_iter().map(|(_, d)| d).collect()
2137 }
2138 Err(e) => {
2139 log::warn!("Checkout attribution task failed: {e}");
2140 daemons
2141 }
2142 }
2143}
2144
2145fn escape_html(s: &str) -> String {
2147 s.replace('&', "&")
2148 .replace('<', "<")
2149 .replace('>', ">")
2150 .replace('"', """)
2151 .replace('\'', "'")
2152}
2153
2154fn html_page(status: StatusCode, title: &str, body: String) -> Response {
2156 let html = format!(
2157 r##"<!DOCTYPE html>
2158<html lang="en">
2159<head>
2160 <meta charset="UTF-8">
2161 <meta name="viewport" content="width=device-width, initial-scale=1">
2162 <title>{title} — pitchfork</title>
2163 <style>
2164 * {{ margin: 0; padding: 0; box-sizing: border-box; }}
2165 body {{
2166 font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
2167 background: #0f1117;
2168 color: #e1e4e8;
2169 display: flex;
2170 align-items: center;
2171 justify-content: center;
2172 min-height: 100vh;
2173 }}
2174 .container {{ max-width: 640px; padding: 2rem; }}
2175 h1 {{ font-size: 1.5rem; font-weight: 600; margin-bottom: 0.75rem; }}
2176 p {{ color: #8b949e; font-size: 0.9rem; margin-bottom: 0.75rem; }}
2177 ul {{ list-style: none; margin: 0.5rem 0 1rem; }}
2178 li {{ margin: 0.25rem 0; }}
2179 code, a {{
2180 color: #58a6ff;
2181 font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
2182 text-decoration: none;
2183 }}
2184 </style>
2185</head>
2186<body>
2187 <div class="container">{body}</div>
2188</body>
2189</html>"##
2190 );
2191 Response::builder()
2192 .status(status)
2193 .header("content-type", "text/html; charset=utf-8")
2194 .body(Body::from(html))
2195 .unwrap_or_else(|_| (status, title.to_string()).into_response())
2196}
2197
2198fn host_port_suffix(raw_host: &str) -> String {
2203 let port = if raw_host.starts_with('[') {
2204 raw_host.split_once("]:").map(|(_, port)| port)
2205 } else {
2206 raw_host.rsplit_once(':').map(|(_, port)| port)
2207 };
2208 port.filter(|p| p.chars().all(|c| c.is_ascii_digit()) && !p.is_empty())
2209 .map(|p| format!(":{p}"))
2210 .unwrap_or_default()
2211}
2212
2213fn page_placeholder_response(
2219 project: &str,
2220 worktree: Option<&str>,
2221 daemons: &[String],
2222 tld: &str,
2223 port_suffix: &str,
2224) -> Response {
2225 let heading = match worktree {
2226 Some(wt) => format!("{} · {}", escape_html(project), escape_html(wt)),
2227 None => escape_html(project),
2228 };
2229 let suffix = match worktree {
2230 Some(wt) => format!(
2231 "{}.{}.{}",
2232 escape_html(wt),
2233 escape_html(project),
2234 escape_html(tld)
2235 ),
2236 None => format!("{}.{}", escape_html(project), escape_html(tld)),
2237 };
2238 let list = if daemons.is_empty() {
2239 "<p>No daemon in this checkout has a port configured.</p>".to_string()
2240 } else {
2241 let items: String = daemons
2242 .iter()
2243 .map(|d| {
2244 let d = escape_html(d);
2245 format!("<li><a href=\"//{d}.{suffix}{port_suffix}\">{d}.{suffix}</a></li>")
2246 })
2247 .collect();
2248 format!("<p>Daemons here:</p><ul>{items}</ul>")
2249 };
2250 let body = format!(
2251 "<h1>{heading}</h1>\
2252 <p>This address is reserved for the {page} page, which is not built yet.</p>\
2253 {list}",
2254 page = if worktree.is_some() {
2255 "stack"
2256 } else {
2257 "project"
2258 },
2259 );
2260 html_page(StatusCode::OK, "pitchfork", body)
2261}
2262
2263fn unknown_host_response(host: &str, heading: &str, known: &[String]) -> Response {
2265 let list = if known.is_empty() {
2266 "<p>Nothing is registered under this name yet.</p>".to_string()
2267 } else {
2268 let items: String = known
2269 .iter()
2270 .map(|k| format!("<li><code>{}</code></li>", escape_html(k)))
2271 .collect();
2272 format!("<p>Known names:</p><ul>{items}</ul>")
2273 };
2274 let body = format!(
2275 "<h1>{heading}</h1><p>No route for <code>{host}</code>.</p>{list}",
2276 heading = escape_html(heading),
2277 host = escape_html(host),
2278 );
2279 html_page(StatusCode::NOT_FOUND, "Not found", body)
2280}
2281
2282fn strip_tld(host: &str, tld: &str) -> Option<String> {
2293 strip_dot_suffix_ignore_case(host, tld)
2294}
2295
2296fn bind_error_message(port: u16, err: &std::io::Error) -> String {
2298 if port < 1024 {
2299 format!(
2300 "Failed to bind proxy server to port {port}: {err}\n\
2301 Hint: ports below 1024 require elevated privileges. \
2302 Try: sudo pitchfork supervisor start"
2303 )
2304 } else {
2305 format!(
2306 "Failed to bind proxy server to port {port}: {err}\n\
2307 Hint: another process may already be using this port."
2308 )
2309 }
2310}
2311
2312fn starting_html_response(slug: &str, raw_host: &str) -> Response {
2317 let escaped_slug = slug
2318 .replace('&', "&")
2319 .replace('<', "<")
2320 .replace('>', ">")
2321 .replace('"', """)
2322 .replace('\'', "'");
2323 let escaped_host = raw_host
2324 .replace('&', "&")
2325 .replace('<', "<")
2326 .replace('>', ">")
2327 .replace('"', """)
2328 .replace('\'', "'");
2329
2330 let html = format!(
2331 r##"<!DOCTYPE html>
2332<html lang="en">
2333<head>
2334 <meta charset="UTF-8">
2335 <meta name="viewport" content="width=device-width, initial-scale=1">
2336 <meta http-equiv="refresh" content="2">
2337 <title>Starting {escaped_slug}… — pitchfork</title>
2338 <style>
2339 * {{ margin: 0; padding: 0; box-sizing: border-box; }}
2340 body {{
2341 font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
2342 background: #0f1117;
2343 color: #e1e4e8;
2344 display: flex;
2345 align-items: center;
2346 justify-content: center;
2347 min-height: 100vh;
2348 }}
2349 .container {{
2350 text-align: center;
2351 max-width: 480px;
2352 padding: 2rem;
2353 }}
2354 .spinner {{
2355 width: 48px;
2356 height: 48px;
2357 border: 4px solid rgba(255, 255, 255, 0.1);
2358 border-top-color: #58a6ff;
2359 border-radius: 50%;
2360 animation: spin 0.8s linear infinite;
2361 margin: 0 auto 1.5rem;
2362 }}
2363 @keyframes spin {{
2364 to {{ transform: rotate(360deg); }}
2365 }}
2366 h1 {{
2367 font-size: 1.5rem;
2368 font-weight: 600;
2369 margin-bottom: 0.5rem;
2370 }}
2371 .slug {{
2372 color: #58a6ff;
2373 font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace;
2374 }}
2375 .host {{
2376 color: #8b949e;
2377 font-size: 0.875rem;
2378 margin-top: 0.25rem;
2379 }}
2380 .hint {{
2381 color: #8b949e;
2382 font-size: 0.8rem;
2383 margin-top: 1.5rem;
2384 }}
2385 </style>
2386</head>
2387<body>
2388 <div class="container">
2389 <div class="spinner"></div>
2390 <h1>Starting <span class="slug">{escaped_slug}</span>…</h1>
2391 <p class="host">{escaped_host}</p>
2392 <p class="hint">This page will refresh automatically when the daemon is ready.</p>
2393 </div>
2394</body>
2395</html>"##
2396 );
2397
2398 Response::builder()
2399 .status(StatusCode::SERVICE_UNAVAILABLE)
2400 .header("content-type", "text/html; charset=utf-8")
2401 .header("retry-after", "2")
2402 .body(Body::from(html))
2403 .unwrap_or_else(|_| (StatusCode::SERVICE_UNAVAILABLE, "Starting…").into_response())
2404}
2405
2406async fn redirect_to_https_handler(req: Request) -> Response {
2415 if req.headers().contains_key("upgrade") {
2417 log::warn!("Dropping plain-HTTP WebSocket upgrade attempt — use wss:// instead of ws://");
2418 return (
2419 StatusCode::BAD_REQUEST,
2420 "WebSocket over plain HTTP is not supported on the HTTPS port. Use wss:// instead.",
2421 )
2422 .into_response();
2423 }
2424
2425 let raw_host = get_request_host(&req);
2426 let Some(raw_host) = raw_host else {
2427 return (StatusCode::BAD_REQUEST, "Missing Host header").into_response();
2428 };
2429
2430 let hostname = if raw_host.starts_with('[') {
2432 raw_host
2434 .split_once("]:")
2435 .map(|(host, _)| host)
2436 .unwrap_or(&raw_host)
2437 .trim_start_matches('[')
2438 .trim_end_matches(']')
2439 } else {
2440 let mut parts = raw_host.rsplitn(2, ':');
2442 let last = parts.next().unwrap_or(&raw_host);
2443 parts.next().unwrap_or(last)
2444 };
2445
2446 let path = req
2447 .uri()
2448 .path_and_query()
2449 .map(|pq| pq.as_str())
2450 .unwrap_or("/");
2451
2452 let https_port = match u16::try_from(settings().proxy.port).ok().filter(|&p| p > 0) {
2453 Some(443) | None => String::new(),
2454 Some(port) => format!(":{port}"),
2455 };
2456
2457 let host_for_url = if raw_host.starts_with('[') {
2458 format!("[{hostname}]")
2459 } else {
2460 hostname.to_string()
2461 };
2462
2463 let location = format!("https://{host_for_url}{https_port}{path}");
2464 (
2465 StatusCode::FOUND,
2466 [(axum::http::header::LOCATION, location)],
2467 )
2468 .into_response()
2469}
2470
2471fn error_response(status: StatusCode, message: &str) -> Response {
2473 (status, message.to_string()).into_response()
2474}
2475
2476#[cfg(test)]
2477mod tests {
2478 use super::*;
2479
2480 #[test]
2481 fn test_strip_tld() {
2482 assert_eq!(
2483 strip_tld("api.myproject.localhost", "localhost"),
2484 Some("api.myproject".to_string())
2485 );
2486 assert_eq!(
2488 strip_tld("API.MyProject.LOCALHOST", "localhost"),
2489 Some("API.MyProject".to_string())
2490 );
2491 assert_eq!(
2492 strip_tld("api.localhost", "LOCALHOST"),
2493 Some("api".to_string())
2494 );
2495 assert_eq!(
2496 strip_tld("api.localhost", "localhost"),
2497 Some("api".to_string())
2498 );
2499 assert_eq!(strip_tld("localhost", "localhost"), None);
2500 assert_eq!(
2501 strip_tld("api.myproject.test", "test"),
2502 Some("api.myproject".to_string())
2503 );
2504 assert_eq!(strip_tld("other.com", "localhost"), None);
2505 }
2506
2507 fn make_entry(name: &str) -> CachedSlugEntry {
2508 CachedSlugEntry {
2509 slug: name.to_string(),
2510 namespace: None,
2511 daemon_name: name.to_string(),
2512 dir: std::path::PathBuf::from(format!("/tmp/{name}")),
2513 worktrees: vec![],
2514 rejected_worktree_prefixes: std::collections::HashSet::new(),
2515 }
2516 }
2517
2518 #[test]
2519 fn test_wildcard_slug_lookup_exact_match() {
2520 let mut entries = std::collections::HashMap::new();
2521 entries.insert("myapp".to_string(), make_entry("myapp"));
2522 let result = wildcard_slug_lookup("myapp", &entries, true);
2524 assert!(result.is_some());
2525 assert_eq!(result.unwrap().daemon_name, "myapp");
2526 }
2527
2528 #[test]
2529 fn test_wildcard_slug_lookup_subdomain_fallback() {
2530 let mut entries = std::collections::HashMap::new();
2531 entries.insert("myapp".to_string(), make_entry("myapp"));
2532 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
2534 assert!(result.is_some());
2535 assert_eq!(result.unwrap().daemon_name, "myapp");
2536 }
2537
2538 #[test]
2539 fn test_wildcard_slug_lookup_nested_fallback() {
2540 let mut entries = std::collections::HashMap::new();
2541 entries.insert("myapp".to_string(), make_entry("myapp"));
2542 let result = wildcard_slug_lookup("a.b.myapp", &entries, true);
2544 assert!(result.is_some());
2545 assert_eq!(result.unwrap().daemon_name, "myapp");
2546 }
2547
2548 #[test]
2549 fn test_wildcard_slug_lookup_no_match() {
2550 let entries = std::collections::HashMap::new();
2551 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
2553 assert!(result.is_none());
2554 }
2555
2556 #[test]
2557 fn test_wildcard_slug_lookup_disabled() {
2558 let mut entries = std::collections::HashMap::new();
2559 entries.insert("myapp".to_string(), make_entry("myapp"));
2560 let result = wildcard_slug_lookup("tenant.myapp", &entries, false);
2562 assert!(result.is_none());
2563 let result = wildcard_slug_lookup("myapp", &entries, false);
2565 assert!(result.is_some());
2566 }
2567
2568 #[test]
2569 fn test_wildcard_slug_lookup_exact_beats_wildcard() {
2570 let mut entries = std::collections::HashMap::new();
2571 entries.insert("myapp".to_string(), make_entry("myapp"));
2572 let mut tenant_entry = make_entry("tenant-daemon");
2573 tenant_entry.slug = "tenant.myapp".to_string();
2574 entries.insert("tenant.myapp".to_string(), tenant_entry);
2575 let result = wildcard_slug_lookup("tenant.myapp", &entries, true);
2577 assert!(result.is_some());
2578 assert_eq!(result.unwrap().daemon_name, "tenant-daemon");
2579 }
2580
2581 #[test]
2582 fn test_wildcard_slug_lookup_ignores_case() {
2583 let mut entries = std::collections::HashMap::new();
2584 entries.insert("myapp".to_string(), make_entry("myapp"));
2585 for host in ["MyApp", "MYAPP", "myapp"] {
2587 let result = wildcard_slug_lookup(host, &entries, true);
2588 assert!(result.is_some(), "exact lookup failed for {host}");
2589 assert_eq!(result.unwrap().daemon_name, "myapp");
2590 }
2591 for host in ["Tenant.MyApp", "tenant.MYAPP", "A.B.MyApp"] {
2593 let result = wildcard_slug_lookup(host, &entries, true);
2594 assert!(result.is_some(), "wildcard lookup failed for {host}");
2595 assert_eq!(result.unwrap().daemon_name, "myapp");
2596 }
2597 }
2598
2599 #[test]
2600 fn test_wildcard_slug_lookup_case_insensitive_registration() {
2601 let mut entries = std::collections::HashMap::new();
2602 let mut entry = make_entry("upper");
2603 entry.slug = "MyApp".to_string();
2604 entries.insert("myapp".to_string(), entry);
2607 for host in ["myapp", "MyApp", "tenant.MYAPP"] {
2608 let result = wildcard_slug_lookup(host, &entries, true);
2609 assert!(result.is_some(), "lookup failed for {host}");
2610 assert_eq!(result.unwrap().daemon_name, "upper");
2611 }
2612 }
2613
2614 fn make_worktree(branch: &str, sanitized: &str) -> crate::proxy::worktree::WorktreeEntry {
2615 crate::proxy::worktree::WorktreeEntry {
2616 path: std::path::PathBuf::from(format!("/tmp/{sanitized}")),
2617 branch: branch.to_string(),
2618 sanitized_branch: sanitized.to_string(),
2619 namespace: Some(sanitized.to_string()),
2620 }
2621 }
2622
2623 #[test]
2624 fn test_reject_case_colliding_worktrees_drops_both_sides() {
2625 let wts = vec![
2626 make_worktree("Feature-A", "Feature-A"),
2627 make_worktree("feature-a", "feature-a"),
2628 make_worktree("main", "main"),
2629 ];
2630 let (kept, rejected) = reject_case_colliding_worktrees(wts);
2631 assert_eq!(kept.len(), 1);
2634 assert_eq!(kept[0].sanitized_branch, "main");
2635 assert!(rejected.contains("feature-a"));
2638 }
2639
2640 #[test]
2641 fn test_reject_case_colliding_worktrees_keeps_unambiguous() {
2642 let wts = vec![
2643 make_worktree("main", "main"),
2644 make_worktree("feature/a", "feature-a"),
2645 ];
2646 let (kept, rejected) = reject_case_colliding_worktrees(wts);
2647 assert_eq!(kept.len(), 2);
2648 assert!(rejected.is_empty());
2649 }
2650
2651 #[test]
2652 fn test_reject_case_colliding_worktrees_drops_sanitize_duplicates() {
2653 let wts = vec![
2656 make_worktree("feature/a", "feature-a"),
2657 make_worktree("feature.a", "feature-a"),
2658 ];
2659 let (kept, rejected) = reject_case_colliding_worktrees(wts);
2660 assert!(kept.is_empty());
2661 assert!(rejected.contains("feature-a"));
2662 }
2663
2664 #[test]
2665 fn test_match_worktree_prefix() {
2666 let mut entry = make_entry("myapp");
2667 entry.worktrees = vec![make_worktree("feature/b", "feature-b")];
2668 entry
2669 .rejected_worktree_prefixes
2670 .insert("feature-a".to_string());
2671
2672 assert!(matches!(
2673 match_worktree_prefix(&entry, "feature-b"),
2674 PrefixMatch::Worktree(_)
2675 ));
2676 assert!(matches!(
2678 match_worktree_prefix(&entry, "Feature-B"),
2679 PrefixMatch::Worktree(_)
2680 ));
2681 assert!(matches!(
2683 match_worktree_prefix(&entry, "feature-a"),
2684 PrefixMatch::Ambiguous
2685 ));
2686 assert!(matches!(
2687 match_worktree_prefix(&entry, "FEATURE-A"),
2688 PrefixMatch::Ambiguous
2689 ));
2690 assert!(matches!(
2692 match_worktree_prefix(&entry, "tenant"),
2693 PrefixMatch::Unknown
2694 ));
2695 }
2696
2697 #[test]
2698 fn test_strip_dot_suffix_ignore_case() {
2699 assert_eq!(
2700 strip_dot_suffix_ignore_case("feature-a.myapp", "myapp"),
2701 Some("feature-a".to_string())
2702 );
2703 assert_eq!(
2704 strip_dot_suffix_ignore_case("Feature-A.MyApp", "myapp"),
2705 Some("Feature-A".to_string())
2706 );
2707 assert_eq!(
2708 strip_dot_suffix_ignore_case("feature-a.myapp", "MYAPP"),
2709 Some("feature-a".to_string())
2710 );
2711 assert_eq!(strip_dot_suffix_ignore_case("xmyapp", "myapp"), None);
2713 assert_eq!(strip_dot_suffix_ignore_case(".myapp", "myapp"), None);
2714 assert_eq!(strip_dot_suffix_ignore_case("myapp", "myapp"), None);
2715 assert_eq!(
2716 strip_dot_suffix_ignore_case("feature-a.other", "myapp"),
2717 None
2718 );
2719 assert_eq!(
2721 strip_dot_suffix_ignore_case("café.myapp", "myapp"),
2722 Some("café".to_string())
2723 );
2724 assert_eq!(strip_dot_suffix_ignore_case("café", "afé"), None);
2725 }
2726
2727 #[cfg(feature = "proxy-tls")]
2728 #[test]
2729 fn test_generate_ca() {
2730 let dir = tempfile::tempdir().unwrap();
2731 let cert_path = dir.path().join("ca.pem");
2732 let key_path = dir.path().join("ca-key.pem");
2733
2734 generate_ca(&cert_path, &key_path).unwrap();
2735
2736 assert!(cert_path.exists(), "ca.pem should be created");
2737 assert!(key_path.exists(), "ca-key.pem should be created");
2738
2739 let cert_pem = std::fs::read_to_string(&cert_path).unwrap();
2740 let key_pem = std::fs::read_to_string(&key_path).unwrap();
2741
2742 assert!(cert_pem.contains("BEGIN CERTIFICATE"), "should be PEM cert");
2743 assert!(
2744 key_pem.contains("BEGIN") && key_pem.contains("PRIVATE KEY"),
2745 "should be PEM key"
2746 );
2747 }
2748
2749 fn cookie_fields(headers: &HeaderMap) -> Vec<&[u8]> {
2751 headers
2752 .get_all(COOKIE)
2753 .iter()
2754 .map(HeaderValue::as_bytes)
2755 .collect()
2756 }
2757
2758 #[test]
2761 fn test_join_cookie_fields_joins_with_semicolon_space() {
2762 let mut headers = HeaderMap::new();
2763 headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
2764 headers.append(COOKIE, HeaderValue::from_static("consent=ads,stats"));
2765 headers.append(COOKIE, HeaderValue::from_static("theme=dark"));
2766
2767 join_cookie_fields(&mut headers);
2768
2769 assert_eq!(
2770 cookie_fields(&headers),
2771 vec![&b"_session=abc123; consent=ads,stats; theme=dark"[..]]
2772 );
2773 }
2774
2775 #[test]
2778 fn test_join_cookie_fields_joins_bytes_outside_ascii() {
2779 let mut headers = HeaderMap::new();
2780 headers.append(COOKIE, HeaderValue::from_static("_session=abc123"));
2781 headers.append(
2782 COOKIE,
2783 HeaderValue::from_bytes(b"name=Jos\xc3\xa9").unwrap(),
2784 );
2785
2786 join_cookie_fields(&mut headers);
2787
2788 assert_eq!(
2789 cookie_fields(&headers),
2790 vec![&b"_session=abc123; name=Jos\xc3\xa9"[..]]
2791 );
2792 }
2793
2794 #[test]
2796 fn test_join_cookie_fields_without_cookies() {
2797 let mut headers = HeaderMap::new();
2798 headers.insert(HOST, HeaderValue::from_static("app.localhost"));
2799
2800 join_cookie_fields(&mut headers);
2801
2802 assert!(headers.get(COOKIE).is_none());
2803 }
2804
2805 #[test]
2806 fn test_host_port_suffix() {
2807 assert_eq!(host_port_suffix("api.myproj.localhost:8088"), ":8088");
2808 assert_eq!(host_port_suffix("api.myproj.localhost"), "");
2809 assert_eq!(host_port_suffix("[::1]:8088"), ":8088");
2810 assert_eq!(host_port_suffix("[::1]"), "");
2811 assert_eq!(host_port_suffix("host:notaport"), "");
2813 }
2814
2815 #[test]
2819 fn test_daemon_runs_in() {
2820 let temp = tempfile::tempdir().unwrap();
2821 let repo = temp.path().join("my-repo");
2822 std::fs::create_dir_all(repo.join(".git/worktrees/feature")).unwrap();
2823 std::fs::create_dir_all(repo.join("sub")).unwrap();
2824 let nested = repo.join(".worktrees/feature");
2826 std::fs::create_dir_all(&nested).unwrap();
2827 std::fs::write(
2828 nested.join(".git"),
2829 format!(
2830 "gitdir: {}\n",
2831 repo.join(".git/worktrees/feature").display()
2832 ),
2833 )
2834 .unwrap();
2835
2836 let root = |p: &std::path::Path| crate::proxy::hostname::checkout_root_of(p);
2837 let repo_root = root(&repo);
2838 let nested_root = root(&nested);
2839
2840 let mut daemon = crate::daemon::Daemon {
2841 dir: Some(repo.join("sub")),
2842 ..Default::default()
2843 };
2844 assert!(daemon_runs_in(&daemon, &repo_root));
2845 assert!(!daemon_runs_in(&daemon, &nested_root));
2846
2847 daemon.dir = Some(nested.clone());
2850 assert!(daemon_runs_in(&daemon, &nested_root));
2851 assert!(!daemon_runs_in(&daemon, &repo_root));
2852
2853 daemon.dir = Some(temp.path().join("elsewhere"));
2855 assert!(!daemon_runs_in(&daemon, &repo_root));
2856
2857 daemon.dir = None;
2858 assert!(!daemon_runs_in(&daemon, &repo_root));
2859 }
2860
2861 #[test]
2864 fn test_is_local_client() {
2865 let build = |info: Option<SocketAddr>| {
2866 let mut req = Request::new(Body::empty());
2867 if let Some(addr) = info {
2868 req.extensions_mut()
2869 .insert(axum::extract::ConnectInfo(addr));
2870 }
2871 req
2872 };
2873
2874 assert!(is_local_client(&build(Some(
2875 "127.0.0.1:5000".parse().unwrap()
2876 ))));
2877 assert!(is_local_client(&build(Some("[::1]:5000".parse().unwrap()))));
2878 assert!(!is_local_client(&build(Some(
2879 "192.168.1.42:5000".parse().unwrap()
2880 ))));
2881 assert!(!is_local_client(&build(None)));
2882 }
2883}