1use std::sync::{
7 Arc, LazyLock, RwLock,
8 atomic::{AtomicBool, Ordering},
9};
10use tracing::warn;
11
12use crate::{Client, ClientBuilder};
13
14static SHARED_FRONTING_POLICY: LazyLock<Arc<RwLock<FrontPolicy>>> =
15 LazyLock::new(|| Arc::new(RwLock::new(FrontPolicy::Off)));
16
17#[derive(Debug)]
19pub(crate) struct Front {
20 pub(crate) policy: Arc<RwLock<FrontPolicy>>,
21 enabled: AtomicBool,
22}
23
24impl Clone for Front {
25 fn clone(&self) -> Self {
26 Self {
27 policy: self.policy.clone(),
28 enabled: AtomicBool::new(false),
29 }
30 }
31}
32
33impl Front {
34 pub(crate) fn new(policy: FrontPolicy) -> Self {
35 Self {
36 enabled: AtomicBool::new(false),
37 policy: Arc::new(RwLock::new(policy)),
38 }
39 }
40
41 pub(crate) fn off() -> Self {
42 Self::new(FrontPolicy::Off)
43 }
44
45 pub(crate) fn shared() -> Self {
46 let policy = SHARED_FRONTING_POLICY.clone();
47 Self {
48 enabled: AtomicBool::new(false),
49 policy,
50 }
51 }
52
53 pub(crate) fn set_policy(&self, policy: FrontPolicy) {
54 *self.policy.write().unwrap() = policy;
55 self.enabled.store(false, Ordering::Relaxed);
56 }
57
58 pub(crate) fn is_enabled(&self) -> bool {
59 match *self.policy.read().unwrap() {
60 FrontPolicy::Off => false,
61 FrontPolicy::OnRetry => self.enabled.load(Ordering::Relaxed),
62 FrontPolicy::Always => true,
63 }
64 }
65
66 pub(crate) fn retry_enable(&self) {
69 if self.is_enabled() {
70 return;
71 }
72 if matches!(*self.policy.read().unwrap(), FrontPolicy::OnRetry) {
73 self.enabled.store(true, Ordering::Relaxed);
74 }
75 }
76}
77
78#[derive(Debug, Default, PartialEq, Clone)]
79pub enum FrontPolicy {
81 Always,
83 OnRetry,
85 #[default]
86 Off,
88}
89
90impl ClientBuilder {
91 pub fn with_fronting(mut self, policy: Option<FrontPolicy>) -> Self {
94 let front = if let Some(p) = policy {
95 Front::new(p)
96 } else {
97 Front::shared()
98 };
99
100 if !self.urls.iter().any(|url| url.has_front()) {
102 warn!(
103 "fronting is enabled, but none of the supplied urls have configured fronting domains: {:?}",
104 self.urls
105 );
106 }
107
108 self.front = front;
109
110 self
111 }
112}
113
114impl Client {
115 pub fn set_front_policy(&mut self, policy: FrontPolicy) {
122 self.front.set_policy(policy)
123 }
124
125 pub fn use_shared_front_policy(&mut self) {
127 self.front = Front::shared();
128 }
129
130 pub fn set_shared_front_policy(policy: FrontPolicy) {
138 *SHARED_FRONTING_POLICY.write().unwrap() = policy;
139 }
140}
141
142#[cfg(test)]
143mod tests {
144 use super::*;
145 use crate::{ApiClientCore, NO_PARAMS, Url};
146 use serial_test::serial;
147
148 impl Front {
149 pub(crate) fn policy(&self) -> FrontPolicy {
150 self.policy.read().unwrap().clone()
151 }
152 }
153
154 #[test]
156 fn set_policy_independent_client() {
157 let url1 = Url::new(
158 "https://validator.global.ssl.fastly.net",
159 Some(vec!["https://yelp.global.ssl.fastly.net"]),
160 )
161 .unwrap();
162
163 let mut client1 = ClientBuilder::new(url1.clone())
164 .unwrap()
165 .with_fronting(Some(FrontPolicy::Off))
166 .build()
167 .unwrap();
168 assert!(client1.front.policy() == FrontPolicy::Off);
169
170 let client2 = ClientBuilder::new(url1.clone())
171 .unwrap()
172 .with_fronting(Some(FrontPolicy::OnRetry))
173 .build()
174 .unwrap();
175
176 client1.set_front_policy(FrontPolicy::Always);
178 assert!(client1.front.policy() == FrontPolicy::Always);
179
180 assert!(client2.front.policy() == FrontPolicy::OnRetry);
183
184 let req = client1
187 .create_request(reqwest::Method::GET, &["/"], NO_PARAMS, None::<&()>)
188 .unwrap()
189 .build()
190 .unwrap();
191
192 let expected_host = url1.host_str().unwrap();
193 assert!(
194 req.headers()
195 .get(reqwest::header::HOST)
196 .is_some_and(|h| h.to_str().unwrap() == expected_host),
197 "{:?} != {:?}",
198 expected_host,
199 req,
200 );
201
202 let expected_front = url1.front_str().unwrap();
203 assert!(
204 req.url()
205 .host()
206 .is_some_and(|url| url.to_string() == expected_front),
207 "{:?} != {:?}",
208 expected_front,
209 req,
210 );
211 }
212
213 #[test]
217 #[serial]
218 fn set_policy_shared_client() {
219 struct RestoreFrontPolicy(FrontPolicy);
223 impl Drop for RestoreFrontPolicy {
224 fn drop(&mut self) {
225 *SHARED_FRONTING_POLICY.write().unwrap() = self.0.clone();
226 }
227 }
228 let _restore_policy = RestoreFrontPolicy(SHARED_FRONTING_POLICY.read().unwrap().clone());
229
230 let url1 = Url::new(
231 "https://validator.global.ssl.fastly.net",
232 Some(vec!["https://yelp.global.ssl.fastly.net"]),
233 )
234 .unwrap();
235
236 Client::set_shared_front_policy(FrontPolicy::Off);
237 assert!(*SHARED_FRONTING_POLICY.read().unwrap() == FrontPolicy::Off);
238
239 let client1 = ClientBuilder::new(url1.clone())
240 .unwrap()
241 .with_fronting(None)
242 .build()
243 .unwrap();
244 assert!(client1.front.policy() == FrontPolicy::Off);
245
246 let mut client2 = ClientBuilder::new(url1.clone())
247 .unwrap()
248 .with_fronting(Some(FrontPolicy::Off))
249 .build()
250 .unwrap();
251
252 Client::set_shared_front_policy(FrontPolicy::Always);
254 assert!(client1.front.policy() == FrontPolicy::Always);
255
256 assert!(client2.front.policy() == FrontPolicy::Off);
258
259 let req = client1
262 .create_request(reqwest::Method::GET, &["/"], NO_PARAMS, None::<&()>)
263 .unwrap()
264 .build()
265 .unwrap();
266
267 let expected_host = url1.host_str().unwrap();
268 assert!(
269 req.headers()
270 .get(reqwest::header::HOST)
271 .is_some_and(|h| h.to_str().unwrap() == expected_host),
272 "{:?} != {:?}",
273 expected_host,
274 req,
275 );
276
277 let expected_front = url1.front_str().unwrap();
278 assert!(
279 req.url()
280 .host()
281 .is_some_and(|url| url.to_string() == expected_front),
282 "{:?} != {:?}",
283 expected_front,
284 req,
285 );
286
287 client2.use_shared_front_policy();
289 assert!(client2.front.policy() == FrontPolicy::Always);
290
291 Client::set_shared_front_policy(FrontPolicy::OnRetry);
294 assert!(client1.front.policy() == FrontPolicy::OnRetry);
295 assert!(client2.front.policy() == FrontPolicy::OnRetry);
296
297 assert!(!client1.front.is_enabled());
298 assert!(!client2.front.is_enabled());
299
300 client1.front.retry_enable();
301 assert!(client1.front.is_enabled());
302 assert!(!client2.front.is_enabled());
303 }
304
305 #[tokio::test]
306 #[serial]
307 async fn nym_api_works() {
308 let url1 = Url::new(
312 "https://validator.global.ssl.fastly.net",
313 Some(vec!["https://yelp.global.ssl.fastly.net"]),
314 )
315 .unwrap(); let client = ClientBuilder::new(url1)
323 .expect("bad url")
324 .with_fronting(Some(FrontPolicy::Always))
325 .build()
326 .expect("failed to build client");
327
328 let response = client
329 .send_request::<_, (), &str, &str>(
330 reqwest::Method::GET,
331 &["api", "v1", "network", "details"],
332 NO_PARAMS,
333 None,
334 )
335 .await
336 .expect("failed get request");
337
338 assert_eq!(response.status(), 200);
340 }
341}
342
343#[cfg(test)]
344mod mocked_tests {
345 use super::*;
346 use crate::{ApiClientCore, HickoryDnsResolver, NO_PARAMS, Url};
347 use hickory_resolver::{
348 net::{DnsError, NetError, NoRecords},
349 proto::{
350 op::{Query, ResponseCode},
351 rr::{Name as HickoryName, RecordType},
352 },
353 };
354 use reqwest::dns::{Addrs, Name, Resolve, Resolving};
355 use serial_test::serial;
356 use std::{
357 collections::HashMap,
358 net::{IpAddr, SocketAddr},
359 str::FromStr,
360 };
361
362 #[tokio::test]
363 #[serial]
364 async fn fallback_on_failure() {
365 let url1 = Url::new(
370 "https://fake-domain.invalid",
371 Some(vec![
372 "https://fake-front-1.invalid",
373 "https://fake-front-2.invalid",
374 ]),
375 )
376 .unwrap();
377 let url2 = Url::new(
378 "https://validator.global.ssl.fastly.net",
379 Some(vec!["https://yelp.global.ssl.fastly.net"]),
380 )
381 .unwrap(); let mock_resolver = MockResolver::new()
384 .with_nxdomain("fake-front-1.invalid")
385 .with_servfail("fake-front-2.invalid");
386
387 let client = ClientBuilder::new_with_urls(vec![url1, url2])
388 .expect("bad url")
389 .with_fronting(Some(FrontPolicy::Always))
390 .dns_resolver(std::sync::Arc::new(mock_resolver))
391 .build()
392 .expect("failed to build client");
393
394 assert_eq!(
396 client.current_url().as_str(),
397 "https://fake-domain.invalid/",
398 );
399 assert_eq!(
400 client.current_url().front_str(),
401 Some("fake-front-1.invalid"),
402 );
403
404 let result = client
405 .send_request::<_, (), &str, &str>(
406 reqwest::Method::GET,
407 &["api", "v1", "network", "details"],
408 NO_PARAMS,
409 None,
410 )
411 .await;
412 assert!(result.is_err());
413
414 assert_eq!(
416 client.current_url().as_str(),
417 "https://fake-domain.invalid/",
418 );
419 assert_eq!(
420 client.current_url().front_str(),
421 Some("fake-front-2.invalid"),
422 );
423
424 let result = client
425 .send_request::<_, (), &str, &str>(
426 reqwest::Method::GET,
427 &["api", "v1", "network", "details"],
428 NO_PARAMS,
429 None,
430 )
431 .await;
432 assert!(result.is_err());
433
434 assert_eq!(
436 client.current_url().as_str(),
437 "https://validator.global.ssl.fastly.net/",
438 );
439 assert_eq!(
440 client.current_url().front_str(),
441 Some("yelp.global.ssl.fastly.net"),
442 );
443 }
444
445 #[derive(Clone)]
447 enum MockOutcome {
448 #[allow(dead_code)]
450 Addrs(Vec<IpAddr>),
451 NxDomain,
453 ServFail,
455 }
456
457 struct MockResolver {
465 outcomes: HashMap<String, MockOutcome>,
466 fallback: HickoryDnsResolver,
467 }
468
469 impl MockResolver {
470 fn new() -> Self {
471 Self {
472 outcomes: HashMap::new(),
473 fallback: HickoryDnsResolver::thread_resolver(),
474 }
475 }
476
477 fn with_nxdomain(mut self, host: &str) -> Self {
478 self.outcomes
479 .insert(host.to_string(), MockOutcome::NxDomain);
480 self
481 }
482
483 fn with_servfail(mut self, host: &str) -> Self {
484 self.outcomes
485 .insert(host.to_string(), MockOutcome::ServFail);
486 self
487 }
488
489 #[allow(dead_code)]
490 fn with_addrs(mut self, host: &str, addrs: Vec<IpAddr>) -> Self {
491 self.outcomes
492 .insert(host.to_string(), MockOutcome::Addrs(addrs));
493 self
494 }
495 }
496
497 fn mock_dns_error(net_err: NetError) -> Box<dyn std::error::Error + Send + Sync> {
502 Box::new(crate::ResolveError::ResolveError(net_err))
503 }
504
505 impl Resolve for MockResolver {
506 fn resolve(&self, name: Name) -> Resolving {
507 let host = name.as_str().to_string();
508 match self.outcomes.get(&host) {
509 Some(MockOutcome::Addrs(addrs)) => {
510 let addrs = addrs.clone();
511 Box::pin(async move {
512 let addrs: Addrs =
513 Box::new(addrs.into_iter().map(|ip| SocketAddr::new(ip, 0)));
514 Ok(addrs)
515 })
516 }
517 Some(MockOutcome::NxDomain) => {
518 let query = HickoryName::from_str(&host)
521 .map(|name| Query::query(name, RecordType::A))
522 .unwrap_or_default();
523 let err = mock_dns_error(NetError::Dns(DnsError::NoRecordsFound(
524 NoRecords::new(Box::new(query), ResponseCode::NXDomain),
525 )));
526 Box::pin(async move { Err(err) })
527 }
528 Some(MockOutcome::ServFail) => {
529 let err = mock_dns_error(NetError::Dns(DnsError::ResponseCode(
530 ResponseCode::ServFail,
531 )));
532 Box::pin(async move { Err(err) })
533 }
534 None => self.fallback.resolve(name),
535 }
536 }
537 }
538}