1use std::{
2 collections::{HashMap, HashSet},
3 error::Error as _,
4 num::NonZeroUsize,
5 sync::{
6 Arc, Mutex, Weak,
7 atomic::{AtomicBool, Ordering},
8 },
9 time::{Duration, Instant},
10};
11
12use futures_util::{StreamExt, stream};
13use reqwest::{Client, Method, Request, Response, StatusCode};
14use tokio::task::JoinHandle;
15
16use crate::{
17 CachePolicy, DEFAULT_MAX_SUBSCRIPTION_BYTES, Error, HealthPolicy, ProxyNode, RefreshReport,
18 Result, RotationPolicy, SourceOutcome, SourceReport, SubscriptionSource, SubscriptionUpdate,
19 SubscriptionValidators, cache::CacheStore, rotation::RotationState,
20};
21
22pub struct PoolBuilder {
25 sources: Vec<SubscriptionSource>,
26 nodes: Vec<ProxyNode>,
27 cache: Option<CachePolicy>,
28 rotation: RotationPolicy,
29 health: HealthPolicy,
30 subscription_client: Option<Client>,
31 request_timeout: Duration,
32 connect_timeout: Duration,
33 subscription_timeout: Duration,
34 refresh_interval: Duration,
35 health_interval: Duration,
36 max_subscription_bytes: usize,
37 max_in_flight_per_proxy: Option<NonZeroUsize>,
38 subscription_concurrency: NonZeroUsize,
39}
40
41impl Default for PoolBuilder {
42 fn default() -> Self {
43 Self {
44 sources: vec![],
45 nodes: vec![],
46 cache: None,
47 rotation: RotationPolicy::default(),
48 health: HealthPolicy::default(),
49 subscription_client: None,
50 request_timeout: Duration::from_secs(30),
51 connect_timeout: Duration::from_secs(10),
52 subscription_timeout: Duration::from_secs(30),
53 refresh_interval: Duration::from_secs(3600),
54 health_interval: Duration::from_secs(60),
55 max_subscription_bytes: DEFAULT_MAX_SUBSCRIPTION_BYTES,
56 max_in_flight_per_proxy: None,
57 subscription_concurrency: NonZeroUsize::new(4).unwrap(),
58 }
59 }
60}
61
62impl PoolBuilder {
63 pub fn new() -> Self {
65 Self::default()
66 }
67 pub fn subscription(mut self, source: SubscriptionSource) -> Self {
69 self.sources.push(source);
70 self
71 }
72 pub fn nodes(mut self, nodes: impl IntoIterator<Item = ProxyNode>) -> Self {
74 self.nodes.extend(nodes);
75 self
76 }
77 pub fn cache(mut self, policy: CachePolicy) -> Self {
79 self.cache = Some(policy);
80 self
81 }
82 pub fn rotation(mut self, policy: RotationPolicy) -> Self {
84 self.rotation = policy;
85 self
86 }
87 pub fn health(mut self, policy: HealthPolicy) -> Self {
89 self.health = policy;
90 self
91 }
92 pub fn subscription_client(mut self, client: Client) -> Self {
95 self.subscription_client = Some(client);
96 self
97 }
98 pub fn request_timeout(mut self, timeout: Duration) -> Self {
100 self.request_timeout = timeout;
101 self
102 }
103 pub fn connect_timeout(mut self, timeout: Duration) -> Self {
105 self.connect_timeout = timeout;
106 self
107 }
108 pub fn subscription_timeout(mut self, timeout: Duration) -> Self {
110 self.subscription_timeout = timeout;
111 self
112 }
113 pub fn refresh_interval(mut self, interval: Duration) -> Self {
115 self.refresh_interval = interval;
116 self
117 }
118 pub fn health_interval(mut self, interval: Duration) -> Self {
120 self.health_interval = interval;
121 self
122 }
123 pub fn max_subscription_bytes(mut self, bytes: usize) -> Self {
125 self.max_subscription_bytes = bytes;
126 self
127 }
128
129 pub fn max_in_flight_per_proxy(mut self, limit: NonZeroUsize) -> Self {
134 self.max_in_flight_per_proxy = Some(limit);
135 self
136 }
137
138 pub fn subscription_concurrency(mut self, limit: NonZeroUsize) -> Self {
140 self.subscription_concurrency = limit;
141 self
142 }
143
144 pub async fn build(mut self) -> Result<ProxyPool> {
147 self.rotation.validate()?;
148 self.health.validate()?;
149 for duration in [
150 self.request_timeout,
151 self.connect_timeout,
152 self.subscription_timeout,
153 self.refresh_interval,
154 self.health_interval,
155 self.health.timeout,
156 self.health.cooldown,
157 ] {
158 if duration.is_zero() || Instant::now().checked_add(duration).is_none() {
159 return Err(Error::Config(
160 "timeouts and intervals must be positive and representable",
161 ));
162 }
163 }
164 if self.max_subscription_bytes == 0 {
165 return Err(Error::Config("subscription size limit must be positive"));
166 }
167 if let Some(cache) = &self.cache {
168 cache.validate()?;
169 }
170 for node in &self.nodes {
171 node.validate()?;
172 }
173 let mut seen = HashSet::new();
174 self.sources.retain(|source| seen.insert(source.key()));
175 let subscription_client = match self.subscription_client {
176 Some(client) => client,
177 None => Client::builder()
178 .no_proxy()
179 .timeout(self.subscription_timeout)
180 .redirect(reqwest::redirect::Policy::limited(5))
181 .build()?,
182 };
183 let pool = ProxyPool {
184 inner: Arc::new(Inner {
185 state: Mutex::new(State::default()),
186 sources: self.sources,
187 static_nodes: self.nodes,
188 cache: self.cache.map(CacheStore::new),
189 rotation: self.rotation,
190 health: self.health,
191 subscription_client,
192 request_timeout: self.request_timeout,
193 connect_timeout: self.connect_timeout,
194 subscription_timeout: self.subscription_timeout,
195 refresh_interval: self.refresh_interval,
196 health_interval: self.health_interval,
197 max_subscription_bytes: self.max_subscription_bytes,
198 max_in_flight_per_proxy: self.max_in_flight_per_proxy,
199 subscription_concurrency: self.subscription_concurrency,
200 maintenance: Mutex::new(Weak::new()),
201 refresh_lock: tokio::sync::Mutex::new(()),
202 probe_lock: tokio::sync::Mutex::new(()),
203 }),
204 };
205 pool.refresh_inner(true).await?;
206 if pool.inner.health.check_on_build {
207 pool.check_health().await;
208 let any_healthy = pool
211 .inner
212 .state
213 .lock()
214 .unwrap_or_else(|error| error.into_inner())
215 .entries
216 .iter()
217 .any(|entry| !entry.pending_probe && entry.cooldown_until.is_none());
218 if !any_healthy {
219 return Err(Error::NoProxyAvailable);
220 }
221 }
222 Ok(pool)
223 }
224}
225
226struct Inner {
227 state: Mutex<State>,
228 sources: Vec<SubscriptionSource>,
229 static_nodes: Vec<ProxyNode>,
230 cache: Option<CacheStore>,
231 rotation: RotationPolicy,
232 health: HealthPolicy,
233 subscription_client: Client,
234 request_timeout: Duration,
235 connect_timeout: Duration,
236 subscription_timeout: Duration,
237 refresh_interval: Duration,
238 health_interval: Duration,
239 max_subscription_bytes: usize,
240 max_in_flight_per_proxy: Option<NonZeroUsize>,
241 subscription_concurrency: NonZeroUsize,
242 refresh_lock: tokio::sync::Mutex<()>,
243 probe_lock: tokio::sync::Mutex<()>,
244 maintenance: Mutex<Weak<MaintenanceGroup>>,
245}
246
247#[derive(Default)]
248struct State {
249 entries: Vec<Entry>,
250 by_source: HashMap<String, SourceState>,
251 rotation: RotationState,
252 session_allocator: RotationState,
253 index: HashMap<String, usize>,
254 candidates: Option<Arc<[String]>>,
255 candidates_expire: Option<Instant>,
256 refresh_revision: u64,
257 last_refresh: Option<RefreshReport>,
258}
259
260#[derive(Clone)]
261struct SourceState {
262 nodes: Vec<ProxyNode>,
263 validators: SubscriptionValidators,
264}
265
266struct Entry {
267 id: String,
268 node: ProxyNode,
269 client: Client,
270 failures: u32,
271 cooldown_until: Option<Instant>,
272 recovery_token: Option<Arc<()>>,
273 pending_probe: bool,
274 revision: u64,
275 generation: Arc<()>,
276 health_epoch: u64,
277 cooldown_round: u32,
278 in_flight: usize,
279}
280
281impl Entry {
282 fn eligible(&self, now: Instant) -> bool {
283 !self.pending_probe
284 && self.recovery_token.is_none()
285 && self.cooldown_until.is_none_or(|until| now >= until)
286 }
287
288 fn available(&self, now: Instant, limit: Option<NonZeroUsize>) -> bool {
289 self.eligible(now) && limit.is_none_or(|limit| self.in_flight < limit.get())
290 }
291
292 fn release_recovery(&mut self, token: Option<&Arc<()>>) {
293 if let Some(token) = token
294 && self
295 .recovery_token
296 .as_ref()
297 .is_some_and(|active| Arc::ptr_eq(active, token))
298 {
299 self.recovery_token = None;
300 }
301 }
302}
303
304#[derive(Clone)]
306pub struct ProxyPool {
307 inner: Arc<Inner>,
308}
309
310#[derive(Clone, Copy, Debug, Default)]
312pub struct PoolStats {
313 pub total: usize,
315 pub eligible: usize,
317 pub unavailable: usize,
319 pub in_flight: usize,
321 pub saturated: usize,
323}
324
325#[derive(Clone)]
329pub struct ProxySession {
330 pool: ProxyPool,
331 rotation: Arc<Mutex<RotationState>>,
332 policy: RotationPolicy,
333}
334
335impl ProxySession {
336 pub fn acquire(&self) -> Result<ProxyLease> {
338 self.pool.acquire_for(&HashSet::new(), Some(self))
339 }
340 pub fn rotate(&self) {
342 self.rotation
343 .lock()
344 .unwrap_or_else(|error| error.into_inner())
345 .force_rotate();
346 }
347 pub async fn execute(&self, request: Request) -> Result<Response> {
349 self.pool.execute_lease(request, self.acquire()?).await
350 }
351 pub async fn execute_with_failover(
353 &self,
354 request: Request,
355 max_attempts: NonZeroUsize,
356 ) -> Result<Response> {
357 self.pool
358 .failover_for(request, max_attempts, Some(self))
359 .await
360 }
361}
362
363impl ProxyPool {
364 pub fn builder() -> PoolBuilder {
366 PoolBuilder::new()
367 }
368
369 pub fn session(&self, policy: RotationPolicy) -> Result<ProxySession> {
372 policy.validate()?;
373 Ok(ProxySession {
374 pool: self.clone(),
375 rotation: Arc::new(Mutex::new(RotationState::default())),
376 policy,
377 })
378 }
379
380 pub fn last_refresh_report(&self) -> Option<RefreshReport> {
382 self.inner
383 .state
384 .lock()
385 .unwrap_or_else(|error| error.into_inner())
386 .last_refresh
387 .clone()
388 }
389
390 pub fn stats(&self) -> PoolStats {
392 let state = self
393 .inner
394 .state
395 .lock()
396 .unwrap_or_else(|error| error.into_inner());
397 let total = state.entries.len();
398 let now = Instant::now();
399 let eligible = state
400 .entries
401 .iter()
402 .filter(|entry| entry.available(now, self.inner.max_in_flight_per_proxy))
403 .count();
404 PoolStats {
405 total,
406 eligible,
407 unavailable: total - eligible,
408 in_flight: state.entries.iter().map(|entry| entry.in_flight).sum(),
409 saturated: state
410 .entries
411 .iter()
412 .filter(|entry| {
413 entry.eligible(now) && !entry.available(now, self.inner.max_in_flight_per_proxy)
414 })
415 .count(),
416 }
417 }
418
419 pub fn nodes(&self) -> Vec<ProxyNode> {
421 self.inner
422 .state
423 .lock()
424 .unwrap_or_else(|error| error.into_inner())
425 .entries
426 .iter()
427 .map(|entry| entry.node.clone())
428 .collect()
429 }
430
431 pub fn rotate(&self) {
433 self.inner
434 .state
435 .lock()
436 .unwrap_or_else(|error| error.into_inner())
437 .rotation
438 .force_rotate();
439 }
440
441 pub fn acquire(&self) -> Result<ProxyLease> {
444 self.acquire_for(&HashSet::new(), None)
445 }
446
447 fn acquire_for(
448 &self,
449 excluded: &HashSet<String>,
450 session: Option<&ProxySession>,
451 ) -> Result<ProxyLease> {
452 let mut state = self
453 .inner
454 .state
455 .lock()
456 .unwrap_or_else(|error| error.into_inner());
457 let now = Instant::now();
458 if state.candidates.is_none() || state.candidates_expire.is_some_and(|until| now >= until) {
459 state.candidates = Some(
460 state
461 .entries
462 .iter()
463 .filter(|entry| entry.available(now, self.inner.max_in_flight_per_proxy))
464 .map(|entry| entry.id.clone())
465 .collect(),
466 );
467 state.candidates_expire = state
468 .entries
469 .iter()
470 .filter_map(|entry| entry.cooldown_until.filter(|until| *until > now))
471 .min();
472 }
473 let all = state.candidates.as_ref().unwrap().clone();
474 let candidates = if excluded.is_empty() {
475 all
476 } else {
477 all.iter()
478 .filter(|id| !excluded.contains(*id))
479 .cloned()
480 .collect::<Arc<[String]>>()
481 };
482 if candidates.is_empty() {
483 return Err(
484 if state.entries.iter().any(|entry| {
485 !excluded.contains(&entry.id)
486 && entry.eligible(now)
487 && !entry.available(now, self.inner.max_in_flight_per_proxy)
488 }) {
489 Error::PoolSaturated
490 } else {
491 Error::NoProxyAvailable
492 },
493 );
494 }
495 let id = if let Some(session) = session {
496 let mut rotation = session
497 .rotation
498 .lock()
499 .unwrap_or_else(|error| error.into_inner());
500 if !rotation.has_current() {
501 let seed = state
502 .session_allocator
503 .select_shared(&candidates, &RotationPolicy::default(), now)
504 .ok_or(Error::NoProxyAvailable)?;
505 rotation.seed(&seed, now);
506 }
507 rotation.select_shared(&candidates, &session.policy, now)
508 } else {
509 state
510 .rotation
511 .select_shared(&candidates, &self.inner.rotation, now)
512 }
513 .ok_or(Error::NoProxyAvailable)?;
514 let index = *state.index.get(&id).ok_or(Error::NoProxyAvailable)?;
515 let entry = &mut state.entries[index];
516 entry.in_flight += 1;
517 if entry.cooldown_until.is_some() {
518 entry.recovery_token = Some(Arc::new(()));
519 }
520 let lease = ProxyLease {
521 node: entry.node.clone(),
522 client: entry.client.clone(),
523 id,
524 owner: Arc::downgrade(&self.inner),
525 recovery_token: entry.recovery_token.clone(),
526 generation: entry.generation.clone(),
527 health_epoch: entry.health_epoch,
528 session_rotation: session.map(|session| Arc::downgrade(&session.rotation)),
529 completed: false,
530 };
531 if !entry.available(now, self.inner.max_in_flight_per_proxy) {
532 state.candidates = None;
533 }
534 Ok(lease)
535 }
536
537 pub async fn execute(&self, request: Request) -> Result<Response> {
543 self.execute_lease(request, self.acquire()?).await
544 }
545
546 async fn execute_lease(&self, request: Request, mut lease: ProxyLease) -> Result<Response> {
547 match lease.client.execute(request).await {
548 Ok(response) => {
549 if response.status() == StatusCode::PROXY_AUTHENTICATION_REQUIRED {
550 lease.report_failure();
551 } else {
552 lease.report_success();
553 }
554 Ok(response)
555 }
556 Err(error) => {
557 if is_passive_failure(&error) {
558 lease.report_failure();
559 }
560 Err(error.into())
561 }
562 }
563 }
564
565 pub async fn execute_with_failover(
569 &self,
570 request: Request,
571 max_attempts: NonZeroUsize,
572 ) -> Result<Response> {
573 self.failover_for(request, max_attempts, None).await
574 }
575
576 async fn failover_for(
577 &self,
578 request: Request,
579 max_attempts: NonZeroUsize,
580 session: Option<&ProxySession>,
581 ) -> Result<Response> {
582 if !matches!(
583 *request.method(),
584 Method::GET | Method::HEAD | Method::OPTIONS
585 ) {
586 return Err(Error::Config(
587 "automatic failover only supports GET, HEAD, and OPTIONS",
588 ));
589 }
590 if request.try_clone().is_none() {
591 return Err(Error::Config("request body cannot be replayed"));
592 }
593 let mut excluded = HashSet::new();
594 let mut last_error = None;
595 for _ in 0..max_attempts.get() {
596 let lease = match self.acquire_for(&excluded, session) {
597 Ok(lease) => lease,
598 Err(error) => return Err(last_error.unwrap_or(error)),
599 };
600 excluded.insert(lease.id.clone());
601 let attempt = request
602 .try_clone()
603 .ok_or(Error::Config("request body cannot be replayed"))?;
604 match self.execute_lease(attempt, lease).await {
605 Err(Error::Transport(error)) if error.is_connect() || error.is_timeout() => {
606 last_error = Some(Error::Transport(error));
607 }
608 result => return result,
609 }
610 }
611 Err(last_error.unwrap_or(Error::NoProxyAvailable))
612 }
613
614 pub async fn refresh(&self) -> Result<RefreshReport> {
617 let report = self.refresh_inner(false).await?;
618 if self.inner.health.check_on_build {
619 self.check_health().await;
620 }
621 Ok(report)
622 }
623
624 async fn refresh_inner(&self, initial: bool) -> Result<RefreshReport> {
625 let observed = self
626 .inner
627 .state
628 .lock()
629 .unwrap_or_else(|error| error.into_inner())
630 .refresh_revision;
631 let _guard = self.inner.refresh_lock.lock().await;
632 let mut by_source = {
633 let state = self
634 .inner
635 .state
636 .lock()
637 .unwrap_or_else(|error| error.into_inner());
638 if !initial
639 && state.refresh_revision != observed
640 && let Some(report) = &state.last_refresh
641 {
642 return Ok(report.clone());
643 }
644 state.by_source.clone()
645 };
646 let loaded = stream::iter(self.inner.sources.clone())
647 .map(|source| {
648 let pool = self.clone();
649 let previous = by_source.get(&source.key()).cloned();
650 async move { pool.load_source(&source, initial, previous).await }
651 })
652 .buffered(self.inner.subscription_concurrency.get())
653 .collect::<Vec<_>>()
654 .await;
655 let mut report = RefreshReport::default();
656 for (source, result) in self.inner.sources.iter().zip(loaded) {
657 report.updated_sources += usize::from(result.report.outcome == SourceOutcome::Updated);
658 report.not_modified_sources +=
659 usize::from(result.report.outcome == SourceOutcome::NotModified);
660 report.cached_sources += usize::from(matches!(
661 result.report.outcome,
662 SourceOutcome::FreshCache | SourceOutcome::StaleCache
663 ));
664 report.failed_sources += usize::from(result.report.error.is_some());
665 report.cache_read_failures += usize::from(result.cache_read_failed);
666 report.cache_write_failures += usize::from(result.cache_write_failed);
667 report.skipped_nodes += result.skipped;
668 if let Some(data) = result.data {
669 by_source.insert(source.key(), data);
670 }
671 report.sources.push(result.report);
672 }
673 let mut nodes = self.inner.static_nodes.clone();
674 for source in &self.inner.sources {
675 if let Some(data) = by_source.get(&source.key()) {
676 nodes.extend(data.nodes.iter().cloned());
677 }
678 }
679 let mut seen = HashSet::new();
680 nodes.retain(|node| seen.insert(node.id()));
681 if nodes.is_empty() {
682 let mut state = self
683 .inner
684 .state
685 .lock()
686 .unwrap_or_else(|error| error.into_inner());
687 state.last_refresh = Some(report.clone());
688 state.refresh_revision = state.refresh_revision.wrapping_add(1);
689 return Err(if report.sources.is_empty() {
690 Error::NoProxyAvailable
691 } else {
692 Error::Initialization(report.sources)
693 });
694 }
695 let existing: HashSet<_> = self
697 .inner
698 .state
699 .lock()
700 .unwrap_or_else(|error| error.into_inner())
701 .entries
702 .iter()
703 .map(|entry| entry.id.clone())
704 .collect();
705 let mut new_clients = HashMap::new();
706 for node in &nodes {
707 if !existing.contains(&node.id()) {
708 node.validate()?;
709 let client = Client::builder()
710 .no_proxy()
711 .proxy(reqwest::Proxy::all(node.url())?)
712 .timeout(self.inner.request_timeout)
713 .connect_timeout(self.inner.connect_timeout)
714 .redirect(reqwest::redirect::Policy::limited(5))
715 .build()?;
716 new_clients.insert(node.id(), client);
717 }
718 }
719 let mut state = self
720 .inner
721 .state
722 .lock()
723 .unwrap_or_else(|error| error.into_inner());
724 let mut old: HashMap<_, _> = std::mem::take(&mut state.entries)
725 .into_iter()
726 .map(|entry| (entry.id.clone(), entry))
727 .collect();
728 for node in nodes {
729 let id = node.id();
730 if let Some(mut entry) = old.remove(&id) {
731 entry.node = node;
732 state.entries.push(entry);
733 } else if let Some(client) = new_clients.remove(&id) {
734 state.entries.push(Entry {
735 id,
736 node,
737 client,
738 failures: 0,
739 cooldown_until: None,
740 recovery_token: None,
741 pending_probe: self.inner.health.check_on_build,
742 revision: 0,
743 generation: Arc::new(()),
744 health_epoch: 0,
745 cooldown_round: 0,
746 in_flight: 0,
747 });
748 }
749 }
750 state.index = state
751 .entries
752 .iter()
753 .enumerate()
754 .map(|(index, entry)| (entry.id.clone(), index))
755 .collect();
756 state.candidates = None;
757 state.candidates_expire = None;
758 state.by_source = by_source;
759 report.nodes = state.entries.len();
760 state.last_refresh = Some(report.clone());
761 state.refresh_revision = state.refresh_revision.wrapping_add(1);
762 Ok(report)
763 }
764
765 async fn load_source(
766 &self,
767 source: &SubscriptionSource,
768 initial: bool,
769 previous: Option<SourceState>,
770 ) -> LoadedSource {
771 let key = source.key();
772 let mut cache_error = None;
773 let cached = if let Some(cache) = &self.inner.cache {
774 match cache.load(&key).await {
775 Ok(cached) => cached,
776 Err(error) => {
777 cache_error = Some(error.to_string());
778 None
779 }
780 }
781 } else {
782 None
783 };
784 let cache_read_failed = cache_error.is_some();
785 if initial
786 && let Some(cached) = &cached
787 && cached.fresh
788 {
789 return LoadedSource {
790 data: Some(SourceState {
791 nodes: cached.nodes.clone(),
792 validators: cached.validators.clone(),
793 }),
794 report: SourceReport {
795 source_key: key,
796 outcome: SourceOutcome::FreshCache,
797 nodes: cached.nodes.len(),
798 error: None,
799 cache_error,
800 },
801 skipped: 0,
802 cache_read_failed,
803 cache_write_failed: false,
804 };
805 }
806 let had_memory = previous.is_some();
807 let previous = previous.or_else(|| {
808 cached.map(|cached| SourceState {
809 nodes: cached.nodes,
810 validators: cached.validators,
811 })
812 });
813 let result = tokio::time::timeout(
814 self.inner.subscription_timeout,
815 source.fetch_update(
816 &self.inner.subscription_client,
817 self.inner.max_subscription_bytes,
818 previous.as_ref().map(|data| &data.validators),
819 ),
820 )
821 .await;
822 let result = match result {
823 Ok(result) => result.map_err(|error| error.to_string()),
824 Err(_) => Err("subscription download timed out".to_owned()),
825 };
826 let (data, outcome, skipped, error) = match result {
827 Ok(SubscriptionUpdate::Modified { report, validators }) => (
828 Some(SourceState {
829 nodes: report.nodes,
830 validators,
831 }),
832 SourceOutcome::Updated,
833 report.skipped,
834 None,
835 ),
836 Ok(SubscriptionUpdate::NotModified) if previous.is_some() => {
837 (previous, SourceOutcome::NotModified, 0, None)
838 }
839 other => {
840 let error = match other {
841 Err(error) => error,
842 _ => "subscription returned 304 without stored nodes".to_owned(),
843 };
844 let outcome = if had_memory {
845 SourceOutcome::Retained
846 } else if previous.is_some() {
847 SourceOutcome::StaleCache
848 } else {
849 SourceOutcome::Failed
850 };
851 (previous, outcome, 0, Some(error))
852 }
853 };
854 let mut cache_write_failed = false;
855 if matches!(outcome, SourceOutcome::Updated | SourceOutcome::NotModified)
856 && let (Some(cache), Some(data)) = (&self.inner.cache, &data)
857 && let Err(error) = cache
858 .save_with_validators(&key, &data.nodes, &data.validators)
859 .await
860 {
861 cache_write_failed = true;
862 cache_error = Some(error.to_string());
863 }
864 let report = SourceReport {
865 source_key: key,
866 outcome,
867 nodes: data.as_ref().map_or(0, |data| data.nodes.len()),
868 error,
869 cache_error,
870 };
871 LoadedSource {
872 data,
873 report,
874 skipped,
875 cache_read_failed,
876 cache_write_failed,
877 }
878 }
879
880 pub async fn check_health(&self) -> PoolStats {
884 let Some(check_url) = self.inner.health.check_url.clone() else {
885 return self.stats();
886 };
887 let _guard = self.inner.probe_lock.lock().await;
888 let probes: Vec<_> = {
889 let mut state = self
890 .inner
891 .state
892 .lock()
893 .unwrap_or_else(|error| error.into_inner());
894 let now = Instant::now();
895 let probes = state
896 .entries
897 .iter_mut()
898 .filter_map(|entry| {
899 if entry.recovery_token.is_some()
900 || (!entry.pending_probe && !entry.eligible(now))
901 {
902 return None;
903 }
904 if entry.cooldown_until.is_some() || entry.pending_probe {
905 entry.recovery_token = Some(Arc::new(()));
906 }
907 Some((
908 ProbeLease {
909 id: entry.id.clone(),
910 revision: entry.revision,
911 generation: entry.generation.clone(),
912 health_epoch: entry.health_epoch,
913 recovery_token: entry.recovery_token.clone(),
914 owner: Arc::downgrade(&self.inner),
915 completed: false,
916 },
917 entry.client.clone(),
918 ))
919 })
920 .collect();
921 state.candidates = None;
922 probes
923 };
924 let timeout = self.inner.health.timeout;
925 stream::iter(probes)
926 .map(|(mut lease, client)| {
927 let check_url = check_url.clone();
928 async move {
929 let result = client.get(check_url).timeout(timeout).send().await;
932 let healthy = result.is_ok_and(|response| {
933 response.status().is_success() || response.status().is_redirection()
934 });
935 lease.complete(healthy);
936 }
937 })
938 .buffer_unordered(self.inner.health.concurrency.get())
939 .collect::<Vec<_>>()
940 .await;
941 self.stats()
942 }
943
944 pub fn spawn_maintenance(&self) -> MaintenanceTask {
947 let mut active = self
948 .inner
949 .maintenance
950 .lock()
951 .unwrap_or_else(|error| error.into_inner());
952 if let Some(group) = active.upgrade()
953 && group.running.load(Ordering::Acquire)
954 {
955 return MaintenanceTask { group };
956 }
957 let mut tasks = Vec::new();
958 for (period, refresh) in [
959 (self.inner.refresh_interval, true),
960 (self.inner.health_interval, false),
961 ] {
962 if (refresh && self.inner.sources.is_empty())
963 || (!refresh && self.inner.health.check_url.is_none())
964 {
965 continue;
966 }
967 let pool = self.clone();
968 tasks.push(tokio::spawn(async move {
969 let mut interval =
970 tokio::time::interval_at(tokio::time::Instant::now() + period, period);
971 interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
972 loop {
973 interval.tick().await;
974 if refresh {
975 let _ = pool.refresh_inner(false).await;
978 } else {
979 pool.check_health().await;
980 }
981 }
982 }));
983 }
984 let group = Arc::new(MaintenanceGroup {
985 tasks: tokio::sync::Mutex::new(tasks),
986 running: AtomicBool::new(true),
987 });
988 *active = Arc::downgrade(&group);
989 MaintenanceTask { group }
990 }
991}
992
993fn is_passive_failure(error: &reqwest::Error) -> bool {
994 let mut source = error.source();
995 let mut disconnected = false;
996 while let Some(cause) = source {
997 if let Some(error) = cause.downcast_ref::<hyper::Error>() {
998 if error.is_user() {
1001 return false;
1002 }
1003 disconnected |= error.is_incomplete_message() || error.is_closed();
1004 }
1005 if let Some(error) = cause.downcast_ref::<std::io::Error>() {
1006 disconnected |= matches!(
1007 error.kind(),
1008 std::io::ErrorKind::ConnectionReset
1009 | std::io::ErrorKind::ConnectionAborted
1010 | std::io::ErrorKind::BrokenPipe
1011 | std::io::ErrorKind::UnexpectedEof
1012 | std::io::ErrorKind::NotConnected
1013 );
1014 }
1015 source = cause.source();
1016 }
1017 error.is_connect() || error.is_timeout() || (error.is_request() && disconnected)
1018}
1019
1020struct LoadedSource {
1021 data: Option<SourceState>,
1022 report: SourceReport,
1023 skipped: usize,
1024 cache_read_failed: bool,
1025 cache_write_failed: bool,
1026}
1027
1028pub struct ProxyLease {
1032 node: ProxyNode,
1033 client: Client,
1034 id: String,
1035 owner: Weak<Inner>,
1036 recovery_token: Option<Arc<()>>,
1037 generation: Arc<()>,
1038 health_epoch: u64,
1039 session_rotation: Option<Weak<Mutex<RotationState>>>,
1040 completed: bool,
1041}
1042
1043impl ProxyLease {
1044 pub fn node(&self) -> &ProxyNode {
1046 &self.node
1047 }
1048 pub fn client(&self) -> &Client {
1050 &self.client
1051 }
1052 pub fn report_success(&mut self) {
1054 self.report(true);
1055 }
1056 pub fn report_failure(&mut self) {
1058 self.report(false);
1059 }
1060 fn report(&mut self, success: bool) {
1061 if !self.completed {
1062 let applied = self.owner.upgrade().is_some_and(|owner| {
1063 owner.finish(
1064 &self.id,
1065 Completion {
1066 generation: &self.generation,
1067 health_epoch: self.health_epoch,
1068 expected_revision: None,
1069 probe: false,
1070 token: self.recovery_token.as_ref(),
1071 outcome: Some(success),
1072 },
1073 )
1074 });
1075 if applied
1076 && !success
1077 && let Some(rotation) = self.session_rotation.as_ref().and_then(Weak::upgrade)
1078 {
1079 rotation
1080 .lock()
1081 .unwrap_or_else(|error| error.into_inner())
1082 .invalidate(&self.id);
1083 }
1084 self.completed = true;
1085 }
1086 }
1087}
1088
1089impl Drop for ProxyLease {
1090 fn drop(&mut self) {
1091 if !self.completed
1092 && let Some(owner) = self.owner.upgrade()
1093 {
1094 owner.finish(
1095 &self.id,
1096 Completion {
1097 generation: &self.generation,
1098 health_epoch: self.health_epoch,
1099 expected_revision: None,
1100 probe: false,
1101 token: self.recovery_token.as_ref(),
1102 outcome: None,
1103 },
1104 );
1105 }
1106 }
1107}
1108
1109struct Completion<'a> {
1110 generation: &'a Arc<()>,
1111 health_epoch: u64,
1112 expected_revision: Option<u64>,
1113 probe: bool,
1114 token: Option<&'a Arc<()>>,
1115 outcome: Option<bool>,
1116}
1117
1118impl Inner {
1119 fn finish(&self, id: &str, completion: Completion<'_>) -> bool {
1122 let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
1123 let Some(index) = state.index.get(id).copied() else {
1124 return false;
1125 };
1126 let entry = &mut state.entries[index];
1127 if !Arc::ptr_eq(&entry.generation, completion.generation) {
1128 return false;
1129 }
1130 let now = Instant::now();
1131 let before = entry.available(now, self.max_in_flight_per_proxy);
1132 let previous_deadline = entry.cooldown_until;
1133 if !completion.probe {
1134 entry.in_flight = entry.in_flight.saturating_sub(1);
1135 }
1136 entry.release_recovery(completion.token);
1137 let applicable = entry.health_epoch == completion.health_epoch
1138 && completion
1139 .expected_revision
1140 .is_none_or(|revision| entry.revision == revision);
1141 let mut failed = false;
1142 if applicable && let Some(success) = completion.outcome {
1143 entry.revision = entry.revision.wrapping_add(1);
1144 entry.pending_probe = false;
1145 if success {
1146 entry.failures = 0;
1147 entry.cooldown_round = 0;
1148 entry.cooldown_until = None;
1149 } else {
1150 failed = true;
1151 entry.failures = entry.failures.saturating_add(1);
1152 if completion.probe
1153 || entry.cooldown_until.is_some()
1154 || entry.failures >= self.health.failure_threshold.get()
1155 {
1156 entry.cooldown_round = entry.cooldown_round.saturating_add(1);
1157 entry.cooldown_until =
1158 now.checked_add(self.health.cooldown_for(entry.cooldown_round));
1159 entry.health_epoch = entry.health_epoch.wrapping_add(1);
1160 }
1161 }
1162 }
1163 if before != entry.available(now, self.max_in_flight_per_proxy)
1164 || previous_deadline != entry.cooldown_until
1165 {
1166 state.candidates = None;
1167 }
1168 if failed {
1169 state.rotation.invalidate(id);
1170 }
1171 applicable
1172 }
1173}
1174
1175struct ProbeLease {
1176 id: String,
1177 revision: u64,
1178 generation: Arc<()>,
1179 health_epoch: u64,
1180 recovery_token: Option<Arc<()>>,
1181 owner: Weak<Inner>,
1182 completed: bool,
1183}
1184impl ProbeLease {
1185 fn complete(&mut self, success: bool) {
1186 if let Some(owner) = self.owner.upgrade() {
1187 owner.finish(
1188 &self.id,
1189 Completion {
1190 generation: &self.generation,
1191 health_epoch: self.health_epoch,
1192 expected_revision: Some(self.revision),
1193 probe: true,
1194 token: self.recovery_token.as_ref(),
1195 outcome: Some(success),
1196 },
1197 );
1198 }
1199 self.completed = true;
1200 }
1201}
1202impl Drop for ProbeLease {
1203 fn drop(&mut self) {
1204 if !self.completed
1205 && let Some(owner) = self.owner.upgrade()
1206 {
1207 owner.finish(
1208 &self.id,
1209 Completion {
1210 generation: &self.generation,
1211 health_epoch: self.health_epoch,
1212 expected_revision: Some(self.revision),
1213 probe: true,
1214 token: self.recovery_token.as_ref(),
1215 outcome: None,
1216 },
1217 );
1218 }
1219 }
1220}
1221
1222#[derive(Clone)]
1225#[must_use = "keep this handle alive for background maintenance to continue"]
1226pub struct MaintenanceTask {
1227 group: Arc<MaintenanceGroup>,
1228}
1229struct MaintenanceGroup {
1230 tasks: tokio::sync::Mutex<Vec<JoinHandle<()>>>,
1231 running: AtomicBool,
1232}
1233impl MaintenanceTask {
1234 pub async fn shutdown(self) {
1236 let mut tasks = self.group.tasks.lock().await;
1237 self.group.running.store(false, Ordering::Release);
1238 for task in tasks.iter() {
1239 task.abort();
1240 }
1241 while let Some(task) = tasks.last_mut() {
1244 let _ = task.await;
1245 tasks.pop();
1246 }
1247 }
1248}
1249impl Drop for MaintenanceGroup {
1250 fn drop(&mut self) {
1251 for task in self.tasks.get_mut() {
1252 task.abort();
1253 }
1254 }
1255}
1256
1257#[cfg(test)]
1258mod tests {
1259 use super::*;
1260
1261 struct Cleanup(Arc<AtomicBool>);
1262 impl Drop for Cleanup {
1263 fn drop(&mut self) {
1264 self.0.store(true, Ordering::Release);
1265 }
1266 }
1267
1268 #[tokio::test]
1269 async fn cancelled_shutdown_can_be_resumed_and_still_waits_for_cleanup() {
1270 let cleaned_up = Arc::new(AtomicBool::new(false));
1271 let observed = cleaned_up.clone();
1272 let (started, ready) = tokio::sync::oneshot::channel();
1273 let task = tokio::spawn(async move {
1274 let _cleanup = Cleanup(observed);
1275 let _ = started.send(());
1276 std::future::pending::<()>().await;
1277 });
1278 ready.await.unwrap();
1279 let handle = MaintenanceTask {
1280 group: Arc::new(MaintenanceGroup {
1281 tasks: tokio::sync::Mutex::new(vec![task]),
1282 running: AtomicBool::new(true),
1283 }),
1284 };
1285 let mut cancelled = Box::pin(handle.clone().shutdown());
1286 assert!(futures_util::poll!(&mut cancelled).is_pending());
1287 drop(cancelled);
1288 assert!(!cleaned_up.load(Ordering::Acquire));
1290 handle.shutdown().await;
1291 assert!(cleaned_up.load(Ordering::Acquire));
1292 }
1293}