1use std::{future::Future, sync::atomic::Ordering};
6
7use moka::sync::Cache as MokaCache;
8use reqwest::Version;
9
10use crate::{
11 agent::Agent,
12 error::{FaithError, FaithErrorKind},
13 warm_up::{extract_host, origin_key, reduce_to_origin},
14};
15
16impl Agent {
17 pub fn prefetch_dns(&self, host: &str) -> Result<impl Future<Output = ()> + use<>, FaithError> {
25 if self.is_closed() {
26 return Err(FaithErrorKind::Closed.into());
27 }
28
29 let Some(host) = extract_host(host) else {
30 return Err(FaithErrorKind::AddressParse.into());
31 };
32
33 #[cfg(feature = "dns")]
34 let resolver = self.dns_resolver_inner();
35 Ok(async move {
36 #[cfg(feature = "dns")]
38 if let Some(resolver) = resolver {
39 resolver.prefetch(&host).await;
40 }
41 #[cfg(not(feature = "dns"))]
42 let _ = host;
43 })
44 }
45
46 pub fn preconnect(&self, origin: &str) -> Result<impl Future<Output = ()> + use<>, FaithError> {
55 let Some(raw_client) = self.raw_client() else {
56 return Err(FaithErrorKind::Closed.into());
57 };
58
59 let Some(url) = reduce_to_origin(origin) else {
60 return Err(FaithErrorKind::AddressParse.into());
61 };
62 let key = origin_key(&url);
63
64 let redundant = self.warmed.contains_key(&key)
67 || !self.warming.entry(key.clone()).or_insert(()).is_fresh();
68
69 #[cfg(feature = "http3")]
76 let h3_port = self
77 .alt_svc_cache()
78 .filter(|_| self.h3_upgrade_enabled)
79 .and_then(|cache| {
80 if self.h3_prober().is_some() {
81 cache.confirmed_port(&url)
82 } else {
83 cache.should_use_h3(&url)
84 }
85 });
86 #[cfg(not(feature = "http3"))]
87 let h3_port: Option<u16> = None;
88
89 #[cfg(feature = "connection-tracking")]
90 let conn_tracker = self.conn_tracker.clone();
91 let warmed = self.warmed.clone();
92 let warming = self.warming.clone();
93 let warm_generation = self.warm_generation.clone();
95 let generation = warm_generation.load(Ordering::Relaxed);
96
97 Ok(async move {
98 if redundant {
99 return;
100 }
101
102 struct ReleaseClaim {
105 warming: MokaCache<String, ()>,
106 key: String,
107 }
108 impl Drop for ReleaseClaim {
109 fn drop(&mut self) {
110 self.warming.invalidate(&self.key);
111 }
112 }
113 let _release = ReleaseClaim {
114 warming,
115 key: key.clone(),
116 };
117
118 let request = match h3_port {
119 Some(port) => {
120 let mut h3_url = url.clone();
121 if Some(port) != h3_url.port_or_known_default() {
125 let _ = h3_url.set_port(Some(port));
126 }
127 raw_client.head(h3_url).version(Version::HTTP_3)
128 }
129 None => raw_client.head(url.clone()),
130 };
131
132 let outcome = request.send().await;
133
134 #[cfg(feature = "connection-tracking")]
137 if h3_port.is_none()
138 && let Ok(response) = &outcome
139 && let Some(info) = response
140 .extensions()
141 .get::<hyper_util::client::legacy::connect::HttpInfo>()
142 {
143 conn_tracker.track_warmup(info.local_addr(), info.remote_addr());
144 }
145
146 if outcome.is_ok() && warm_generation.load(Ordering::Relaxed) == generation {
151 warmed.insert(key, ());
152 }
153 })
154 }
155}