1use std::error::Error as StdError;
25use std::fmt;
26use std::io;
27use std::ops::ControlFlow;
28use std::time::Duration;
29
30#[derive(Debug, Clone, Copy)]
32pub struct RetryPolicy {
33 pub max_attempts: u32,
47 pub base_delay: Duration,
49 pub max_delay: Duration,
51}
52
53impl RetryPolicy {
54 pub const UPLOAD: RetryPolicy = RetryPolicy {
57 max_attempts: 10,
58 base_delay: Duration::from_millis(50),
59 max_delay: Duration::from_secs(30),
60 };
61
62 pub fn delay_for(&self, next_attempt: u32) -> Duration {
63 let exp = next_attempt.saturating_sub(2);
67 let mult = 1u64.checked_shl(exp).unwrap_or(u64::MAX);
68 let ms = (self.base_delay.as_millis() as u64).saturating_mul(mult);
69 std::cmp::min(Duration::from_millis(ms), self.max_delay)
70 }
71}
72
73pub fn retry_sync<T, E, F>(policy: &RetryPolicy, mut op: F) -> Result<T, E>
82where
83 F: FnMut(u32) -> Result<T, ControlFlow<E, E>>,
84{
85 let max = policy.max_attempts.max(1);
86 let mut attempt: u32 = 1;
87 loop {
88 if attempt > 1 {
89 std::thread::sleep(policy.delay_for(attempt));
90 }
91 match op(attempt) {
92 Ok(v) => return Ok(v),
93 Err(ControlFlow::Break(e)) => return Err(e),
94 Err(ControlFlow::Continue(e)) => {
95 if attempt >= max {
96 return Err(e);
97 }
98 }
99 }
100 attempt += 1;
101 }
102}
103
104pub async fn retry_async<T, E, F, Fut>(policy: &RetryPolicy, mut op: F) -> Result<T, E>
108where
109 F: FnMut(u32) -> Fut,
110 Fut: std::future::Future<Output = Result<T, ControlFlow<E, E>>>,
111{
112 let max = policy.max_attempts.max(1);
113 let mut attempt: u32 = 1;
114 loop {
115 if attempt > 1 {
116 tokio::time::sleep(policy.delay_for(attempt)).await;
117 }
118 match op(attempt).await {
119 Ok(v) => return Ok(v),
120 Err(ControlFlow::Break(e)) => return Err(e),
121 Err(ControlFlow::Continue(e)) => {
122 if attempt >= max {
123 return Err(e);
124 }
125 }
126 }
127 attempt += 1;
128 }
129}
130
131#[derive(Debug, Clone, Copy, PartialEq, Eq)]
135pub enum SuccessClass {
136 Strict,
139 AllowRedirects,
143}
144
145pub fn retry_http_blocking<F, M>(
172 label: &str,
173 policy: &RetryPolicy,
174 success_class: SuccessClass,
175 mut send: F,
176 error_msg: M,
177) -> anyhow::Result<(reqwest::StatusCode, String)>
178where
179 F: FnMut(u32) -> Result<reqwest::blocking::Response, reqwest::Error>,
180 M: Fn(reqwest::StatusCode, &str) -> String,
181{
182 use anyhow::Context as _;
183 retry_sync(policy, |attempt| {
184 match send(attempt) {
185 Ok(resp) => {
186 let status = resp.status();
187 let succeeded = match success_class {
188 SuccessClass::Strict => status.is_success(),
189 SuccessClass::AllowRedirects => status.is_success() || status.is_redirection(),
190 };
191 let body = resp
192 .text()
193 .unwrap_or_else(|e| format!("<failed to read body: {e}>"));
194 if succeeded {
195 Ok((status, body))
196 } else {
197 let msg = error_msg(status, &body);
198 let inner = anyhow::anyhow!("{msg}");
199 let wrapped = anyhow::Error::new(HttpError::new(
200 std::io::Error::other(inner.to_string()),
201 status.as_u16(),
202 ))
203 .context(inner);
204 if is_retriable(wrapped.as_ref()) {
210 Err(ControlFlow::Continue(wrapped))
211 } else {
212 Err(ControlFlow::Break(wrapped))
213 }
214 }
215 }
216 Err(e) => {
217 let err = anyhow::Error::new(HttpError::from_response(e, None))
221 .context(format!("{label}: HTTP transport error"));
222 if is_retriable(err.as_ref()) {
223 Err(ControlFlow::Continue(err))
224 } else {
225 Err(ControlFlow::Break(err))
226 }
227 }
228 }
229 })
230 .with_context(|| format!("{label}: exhausted retry attempts"))
231}
232
233pub async fn retry_http_async<F, Fut, M>(
254 label: &str,
255 policy: &RetryPolicy,
256 success_class: SuccessClass,
257 mut send: F,
258 error_msg: M,
259) -> anyhow::Result<reqwest::Response>
260where
261 F: FnMut(u32) -> Fut,
262 Fut: std::future::Future<Output = Result<reqwest::Response, reqwest::Error>>,
263 M: Fn(reqwest::StatusCode, &str) -> String,
264{
265 use anyhow::Context as _;
266 retry_async(policy, |attempt| {
267 let fut = send(attempt);
268 let error_msg = &error_msg;
269 async move {
270 match fut.await {
271 Ok(resp) => {
272 let status = resp.status();
273 let succeeded = match success_class {
274 SuccessClass::Strict => status.is_success(),
275 SuccessClass::AllowRedirects => {
276 status.is_success() || status.is_redirection()
277 }
278 };
279 if succeeded {
280 Ok(resp)
281 } else {
282 let body = resp
283 .text()
284 .await
285 .unwrap_or_else(|e| format!("<failed to read body: {e}>"));
286 let msg = error_msg(status, &body);
287 let inner = anyhow::anyhow!("{msg}");
288 let wrapped = anyhow::Error::new(HttpError::new(
289 std::io::Error::other(inner.to_string()),
290 status.as_u16(),
291 ))
292 .context(inner);
293 if is_retriable(wrapped.as_ref()) {
299 Err(ControlFlow::Continue(wrapped))
300 } else {
301 Err(ControlFlow::Break(wrapped))
302 }
303 }
304 }
305 Err(e) => {
306 let err = anyhow::Error::new(HttpError::from_response(e, None))
310 .context(format!("{label}: HTTP transport error"));
311 if is_retriable(err.as_ref()) {
312 Err(ControlFlow::Continue(err))
313 } else {
314 Err(ControlFlow::Break(err))
315 }
316 }
317 }
318 }
319 })
320 .await
321 .with_context(|| format!("{label}: exhausted retry attempts"))
322}
323
324pub fn classify_http_sync(
332 result: reqwest::Result<reqwest::blocking::Response>,
333) -> Result<reqwest::blocking::Response, ControlFlow<anyhow::Error, anyhow::Error>> {
334 use anyhow::anyhow;
335 match result {
336 Ok(resp) => {
337 let status = resp.status();
338 if status.is_success() || status.is_redirection() {
339 Ok(resp)
340 } else if status.is_server_error() {
341 Err(ControlFlow::Continue(anyhow!(
342 "HTTP {} {}",
343 status.as_u16(),
344 status.canonical_reason().unwrap_or("server error")
345 )))
346 } else {
347 Err(ControlFlow::Break(anyhow!(
349 "HTTP {} {}",
350 status.as_u16(),
351 status.canonical_reason().unwrap_or("client error")
352 )))
353 }
354 }
355 Err(e) => Err(ControlFlow::Continue(anyhow!(e))),
357 }
358}
359
360#[derive(Debug)]
376pub struct HttpError {
377 source: Box<dyn StdError + Send + Sync + 'static>,
380 pub status: u16,
382}
383
384impl HttpError {
385 pub fn new<E>(source: E, status: u16) -> Self
388 where
389 E: StdError + Send + Sync + 'static,
390 {
391 Self {
392 source: Box::new(source),
393 status,
394 }
395 }
396
397 pub fn from_response<E>(err: E, resp: Option<&reqwest::Response>) -> Self
401 where
402 E: StdError + Send + Sync + 'static,
403 {
404 Self::new(err, resp.map(|r| r.status().as_u16()).unwrap_or(0))
405 }
406}
407
408impl fmt::Display for HttpError {
409 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
410 fmt::Display::fmt(&self.source, f)
413 }
414}
415
416impl StdError for HttpError {
417 fn source(&self) -> Option<&(dyn StdError + 'static)> {
418 Some(&*self.source)
419 }
420}
421
422#[derive(Debug)]
428pub struct Retriable(Box<dyn StdError + Send + Sync + 'static>);
429
430impl Retriable {
431 pub fn new<E>(source: E) -> Self
437 where
438 E: StdError + Send + Sync + 'static,
439 {
440 Self(Box::new(source))
441 }
442}
443
444impl fmt::Display for Retriable {
445 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
446 fmt::Display::fmt(&self.0, f)
447 }
448}
449
450impl StdError for Retriable {
451 fn source(&self) -> Option<&(dyn StdError + 'static)> {
452 Some(&*self.0)
453 }
454}
455
456pub fn is_network_error(err: &(dyn StdError + 'static)) -> bool {
485 let mut cur: Option<&(dyn StdError + 'static)> = Some(err);
486 while let Some(e) = cur {
487 if let Some(io_err) = e.downcast_ref::<io::Error>() {
490 match io_err.kind() {
491 io::ErrorKind::UnexpectedEof
492 | io::ErrorKind::TimedOut
493 | io::ErrorKind::ConnectionRefused
494 | io::ErrorKind::ConnectionReset
495 | io::ErrorKind::ConnectionAborted
496 | io::ErrorKind::BrokenPipe => return true,
497 _ => {}
498 }
499 let m = io_err.to_string().to_lowercase();
500 if m == "eof" || m == "unexpected eof" {
501 return true;
502 }
503 }
504
505 let s = e.to_string().to_lowercase();
509 if NETWORK_ERROR_NEEDLES.iter().any(|n| s.contains(n)) {
510 return true;
511 }
512
513 cur = e.source();
514 }
515 false
516}
517
518const NETWORK_ERROR_NEEDLES: &[&str] = &[
528 "connection reset",
529 "network is unreachable",
530 "connection closed",
531 "connection refused",
532 "tls handshake timeout",
533 "i/o timeout",
534 "broken pipe",
535 "timeout awaiting response headers",
536 "context deadline exceeded",
537 "operation timed out",
539 "the network connection was aborted",
541 "an existing connection was forcibly closed",
543 "dns error",
550 "failed to lookup address",
553 "no such host is known",
555];
556
557pub fn is_retriable(err: &(dyn StdError + 'static)) -> bool {
569 let mut cur: Option<&(dyn StdError + 'static)> = Some(err);
571 while let Some(e) = cur {
572 if e.is::<Retriable>() {
573 return true;
574 }
575 if let Some(http) = e.downcast_ref::<HttpError>()
576 && (http.status >= 500 || http.status == 429)
577 {
578 return true;
579 }
580 cur = e.source();
581 }
582
583 is_network_error(err)
585}
586
587pub fn is_retriable_opt(err: Option<&(dyn StdError + 'static)>) -> bool {
590 err.is_some_and(is_retriable)
591}
592
593pub fn jitter_duration(base: Duration) -> Duration {
604 let nanos = base.as_nanos() as u64;
605 let window = nanos / 5;
607 if window == 0 {
608 return base;
609 }
610 let seed = crate::sde::resolve_now().timestamp_subsec_nanos() as u64;
617 let offset = seed % (window * 2);
618 let jittered = nanos.saturating_sub(window).saturating_add(offset);
620 Duration::from_nanos(jittered)
621}
622
623#[cfg(test)]
624mod tests {
625 use super::*;
626 use std::sync::atomic::{AtomicU32, Ordering};
627
628 fn fast_policy() -> RetryPolicy {
629 RetryPolicy {
630 max_attempts: 4,
631 base_delay: Duration::from_millis(1),
632 max_delay: Duration::from_millis(5),
633 }
634 }
635
636 #[test]
637 fn jitter_returns_base_when_window_rounds_to_zero() {
638 for n in 0..5u64 {
642 let base = Duration::from_nanos(n);
643 assert_eq!(
644 jitter_duration(base),
645 base,
646 "sub-5ns base {n} must pass through unjittered"
647 );
648 }
649 }
650
651 #[test]
652 fn jitter_stays_within_plus_minus_twenty_percent() {
653 let base = Duration::from_millis(100);
656 let jittered = jitter_duration(base);
657 let lo = base.mul_f64(0.8);
658 let hi = base.mul_f64(1.2);
659 assert!(
660 jittered >= lo && jittered < hi,
661 "jittered {jittered:?} outside [{lo:?}, {hi:?})"
662 );
663 }
664
665 #[test]
666 fn delay_progression_caps_at_max() {
667 let p = RetryPolicy {
668 max_attempts: 10,
669 base_delay: Duration::from_millis(100),
670 max_delay: Duration::from_millis(500),
671 };
672 assert_eq!(p.delay_for(2), Duration::from_millis(100));
673 assert_eq!(p.delay_for(3), Duration::from_millis(200));
674 assert_eq!(p.delay_for(4), Duration::from_millis(400));
675 assert_eq!(p.delay_for(5), Duration::from_millis(500)); assert_eq!(p.delay_for(8), Duration::from_millis(500)); }
678
679 #[test]
680 fn sync_succeeds_on_first_attempt() {
681 let calls = AtomicU32::new(0);
682 let result: Result<&str, ()> = retry_sync(&fast_policy(), |_| {
683 calls.fetch_add(1, Ordering::SeqCst);
684 Ok("ok")
685 });
686 assert_eq!(result, Ok("ok"));
687 assert_eq!(calls.load(Ordering::SeqCst), 1);
688 }
689
690 #[test]
691 fn sync_retries_until_success() {
692 let calls = AtomicU32::new(0);
693 let result: Result<u32, &str> = retry_sync(&fast_policy(), |attempt| {
694 calls.fetch_add(1, Ordering::SeqCst);
695 if attempt < 3 {
696 Err(ControlFlow::Continue("transient"))
697 } else {
698 Ok(attempt)
699 }
700 });
701 assert_eq!(result, Ok(3));
702 assert_eq!(calls.load(Ordering::SeqCst), 3);
703 }
704
705 #[test]
706 fn sync_break_stops_immediately() {
707 let calls = AtomicU32::new(0);
708 let result: Result<(), &str> = retry_sync(&fast_policy(), |_| {
709 calls.fetch_add(1, Ordering::SeqCst);
710 Err(ControlFlow::Break("fatal"))
711 });
712 assert_eq!(result, Err("fatal"));
713 assert_eq!(calls.load(Ordering::SeqCst), 1);
714 }
715
716 #[test]
717 fn sync_returns_last_error_after_exhaustion() {
718 let calls = AtomicU32::new(0);
719 let result: Result<(), String> = retry_sync(&fast_policy(), |attempt| {
720 calls.fetch_add(1, Ordering::SeqCst);
721 Err(ControlFlow::Continue(format!("fail {attempt}")))
722 });
723 assert_eq!(result, Err("fail 4".to_string()));
724 assert_eq!(calls.load(Ordering::SeqCst), 4);
725 }
726
727 #[tokio::test]
728 async fn async_retries_until_success() {
729 let calls = std::sync::Arc::new(AtomicU32::new(0));
730 let calls_inner = calls.clone();
731 let result: Result<u32, &str> = retry_async(&fast_policy(), move |attempt| {
732 let c = calls_inner.clone();
733 async move {
734 c.fetch_add(1, Ordering::SeqCst);
735 if attempt < 2 {
736 Err(ControlFlow::Continue("transient"))
737 } else {
738 Ok(attempt)
739 }
740 }
741 })
742 .await;
743 assert_eq!(result, Ok(2));
744 assert_eq!(calls.load(Ordering::SeqCst), 2);
745 }
746
747 #[derive(Debug)]
755 struct StrErr(&'static str);
756 impl fmt::Display for StrErr {
757 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
758 f.write_str(self.0)
759 }
760 }
761 impl StdError for StrErr {}
762
763 #[derive(Debug)]
764 struct OwnedErr(String);
765 impl fmt::Display for OwnedErr {
766 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
767 f.write_str(&self.0)
768 }
769 }
770 impl StdError for OwnedErr {}
771
772 #[test]
773 fn network_error_substrings_match() {
774 for s in [
775 "connection reset by peer",
776 "network is unreachable",
777 "connection closed unexpectedly",
778 "connection refused",
779 "tls handshake timeout",
780 "i/o timeout",
781 "CONNECTION RESET",
782 "TLS Handshake Timeout",
783 "write: broken pipe",
784 "net/http: timeout awaiting response headers",
785 "context deadline exceeded",
786 "client error (Connect): dns error: failed to lookup address information: Name or service not known",
791 "dns error: nodename nor servname provided, or not known",
792 "dns error: No such host is known. (os error 11001)",
793 ] {
794 let e = OwnedErr(s.to_string());
795 assert!(is_network_error(&e), "expected network error: {s:?}");
796 }
797 }
798
799 #[test]
800 fn network_error_io_eof_kinds() {
801 let e = io::Error::from(io::ErrorKind::UnexpectedEof);
802 assert!(is_network_error(&e));
803
804 let e2 = io::Error::other("EOF");
806 assert!(is_network_error(&e2));
807 }
808
809 #[test]
816 fn is_network_error_classifies_io_timedout() {
817 let e = io::Error::from(io::ErrorKind::TimedOut);
818 assert!(is_network_error(&e));
819 assert!(is_retriable(&e));
820 }
821
822 #[test]
823 fn is_network_error_classifies_io_connection_refused() {
824 let e = io::Error::from(io::ErrorKind::ConnectionRefused);
825 assert!(is_network_error(&e));
826 assert!(is_retriable(&e));
827 }
828
829 #[test]
830 fn is_network_error_classifies_io_connection_reset() {
831 let e = io::Error::from(io::ErrorKind::ConnectionReset);
832 assert!(is_network_error(&e));
833 assert!(is_retriable(&e));
834 }
835
836 #[test]
837 fn is_network_error_classifies_io_connection_aborted() {
838 let e = io::Error::from(io::ErrorKind::ConnectionAborted);
839 assert!(is_network_error(&e));
840 assert!(is_retriable(&e));
841 }
842
843 #[test]
844 fn is_network_error_classifies_io_broken_pipe() {
845 let e = io::Error::from(io::ErrorKind::BrokenPipe);
846 assert!(is_network_error(&e));
847 assert!(is_retriable(&e));
848 }
849
850 #[test]
851 fn is_network_error_classifies_operation_timed_out_substring() {
852 let other_kind = io::Error::other("operation timed out");
857 assert!(is_network_error(&other_kind));
858 assert!(is_retriable(&other_kind));
859
860 let kind_only = io::Error::from(io::ErrorKind::TimedOut);
861 assert!(is_network_error(&kind_only));
862 assert!(is_retriable(&kind_only));
863 }
864
865 #[test]
866 fn network_error_wrapped_unexpected_eof() {
867 #[derive(Debug)]
869 struct Wrap(io::Error);
870 impl fmt::Display for Wrap {
871 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
872 write!(f, "read failed")
873 }
874 }
875 impl StdError for Wrap {
876 fn source(&self) -> Option<&(dyn StdError + 'static)> {
877 Some(&self.0)
878 }
879 }
880 let inner = io::Error::from(io::ErrorKind::UnexpectedEof);
881 let outer = Wrap(inner);
882 assert!(is_network_error(&outer));
883 }
884
885 #[test]
886 fn network_error_non_network_strings_reject() {
887 for s in [
888 "file not found",
889 "permission denied",
890 "dial tcp: lookup example.com: no such host",
891 "",
892 ] {
893 let e = OwnedErr(s.to_string());
894 assert!(!is_network_error(&e), "expected NOT network error: {s:?}");
895 }
896 }
897
898 #[test]
899 fn retriable_opt_nil_passthrough() {
900 assert!(!is_retriable_opt(None));
901 }
902
903 #[test]
904 fn http_error_500_retriable() {
905 let e = HttpError::new(StrErr("internal server error"), 500);
906 assert!(is_retriable(&e));
907 }
908
909 #[test]
910 fn http_error_502_503_retriable() {
911 for s in [502u16, 503] {
912 let e = HttpError::new(StrErr("bad gateway"), s);
913 assert!(is_retriable(&e), "status {s} should be retriable");
914 }
915 }
916
917 #[test]
918 fn http_error_429_retriable() {
919 let e = HttpError::new(StrErr("rate limited"), 429);
920 assert!(is_retriable(&e));
921 }
922
923 #[test]
924 fn http_error_4xx_not_retriable() {
925 for s in [400u16, 401, 403, 404, 422] {
926 let e = HttpError::new(StrErr("client err"), s);
927 assert!(!is_retriable(&e), "status {s} should NOT be retriable");
928 }
929 }
930
931 #[test]
932 fn http_error_zero_status_routes_via_message() {
933 let net = HttpError::new(StrErr("connection reset"), 0);
936 assert!(is_retriable(&net));
937
938 let non_net = HttpError::new(StrErr("dial failed"), 0);
939 assert!(!is_retriable(&non_net));
940 }
941
942 #[test]
943 fn http_error_unwrap_chain_visible() {
944 let inner = StrErr("inner");
945 let e = HttpError::new(inner, 503);
946 assert!(e.source().is_some());
947 }
948
949 #[test]
950 fn from_response_nil_resp_yields_status_zero() {
951 let inner = io::Error::other("connect: dial tcp");
955 let e = HttpError::from_response(inner, None);
956 assert_eq!(e.status, 0);
957 }
958
959 #[test]
960 fn from_response_unwrap_chain_visible() {
961 let inner = io::Error::other("connection reset by peer");
964 let e = HttpError::from_response(inner, None);
965 assert!(
966 e.source().is_some(),
967 "inner error must be reachable via source()"
968 );
969 assert!(is_retriable(&e));
971 }
972
973 #[test]
974 fn retriable_wrapper_is_retriable() {
975 let e = Retriable::new(StrErr("retry me"));
976 assert!(is_retriable(&e));
977 }
978
979 #[test]
980 fn retriable_wrapper_overrides_4xx() {
981 let inner = HttpError::new(StrErr("exists"), 422);
983 let outer = Retriable::new(inner);
984 assert!(is_retriable(&outer));
985 }
986
987 #[test]
988 fn retriable_wrapper_unwrap_chain_visible() {
989 let inner = StrErr("inner");
990 let e = Retriable::new(inner);
991 assert!(e.source().is_some());
992 }
993
994 #[test]
995 fn plain_error_not_retriable() {
996 let e = StrErr("something");
997 assert!(!is_retriable(&e));
998 }
999
1000 #[test]
1001 fn anyhow_error_threadable() {
1002 let e: anyhow::Error = anyhow::anyhow!("connection refused");
1005 assert!(is_retriable(e.as_ref()));
1006
1007 let e2: anyhow::Error = anyhow::anyhow!("permission denied");
1008 assert!(!is_retriable(e2.as_ref()));
1009 }
1010
1011 #[test]
1012 fn is_retriable_chain_walks_to_http_error() {
1013 let inner = HttpError::new(StrErr("bad gateway"), 503);
1017 let wrapped: anyhow::Error = anyhow::Error::new(inner).context("publish failed");
1018 assert!(is_retriable(wrapped.as_ref()));
1019 }
1020
1021 #[test]
1032 fn classifier_5xx_via_anyhow_chain_uses_as_ref() {
1033 let wrapped: anyhow::Error =
1034 anyhow::Error::new(HttpError::new(std::io::Error::other("503"), 503))
1035 .context("publish");
1036 assert!(
1037 is_retriable(wrapped.as_ref()),
1038 "5xx HttpError reached via as_ref() must classify retriable"
1039 );
1040 }
1041
1042 #[test]
1043 fn classifier_root_cause_walks_past_http_error_drift_guard() {
1044 let wrapped: anyhow::Error =
1049 anyhow::Error::new(HttpError::new(std::io::Error::other("503"), 503))
1050 .context("publish");
1051 assert!(
1052 !is_retriable(wrapped.root_cause()),
1053 "root_cause() walks past HttpError; 5xx must NOT be detected via the leaf"
1054 );
1055 }
1056
1057 #[test]
1058 fn classifier_429_via_anyhow_chain_uses_as_ref() {
1059 let wrapped: anyhow::Error =
1062 anyhow::Error::new(HttpError::new(std::io::Error::other("429"), 429))
1063 .context("publish");
1064 assert!(is_retriable(wrapped.as_ref()));
1065 assert!(!is_retriable(wrapped.root_cause()));
1066 }
1067
1068 use crate::test_helpers::responder::spawn_oneshot_http_responder;
1077
1078 #[test]
1079 fn retry_http_blocking_success_returns_first_attempt() {
1080 let (addr, calls) =
1081 spawn_oneshot_http_responder(vec!["HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok"]);
1082 let client = reqwest::blocking::Client::builder()
1083 .timeout(Duration::from_secs(2))
1084 .build()
1085 .expect("client");
1086 let policy = RetryPolicy {
1087 max_attempts: 3,
1088 base_delay: Duration::from_millis(1),
1089 max_delay: Duration::from_millis(2),
1090 };
1091 let result = retry_http_blocking(
1092 "test",
1093 &policy,
1094 SuccessClass::Strict,
1095 |_| client.get(format!("http://{addr}/")).send(),
1096 |_, _| String::from("should not be called on success"),
1097 );
1098 let (status, body) = result.expect("success");
1099 assert_eq!(status.as_u16(), 200);
1100 assert_eq!(body, "ok");
1101 assert_eq!(calls.load(Ordering::SeqCst), 1, "single attempt");
1102 }
1103
1104 #[test]
1105 fn retry_http_blocking_retries_5xx_then_succeeds() {
1106 let (addr, calls) = spawn_oneshot_http_responder(vec![
1107 "HTTP/1.1 503 Service Unavailable\r\nContent-Length: 0\r\n\r\n",
1108 "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok",
1109 ]);
1110 let client = reqwest::blocking::Client::builder()
1111 .timeout(Duration::from_secs(2))
1112 .build()
1113 .expect("client");
1114 let policy = RetryPolicy {
1115 max_attempts: 3,
1116 base_delay: Duration::from_millis(1),
1117 max_delay: Duration::from_millis(2),
1118 };
1119 let result = retry_http_blocking(
1120 "test",
1121 &policy,
1122 SuccessClass::Strict,
1123 |_| client.get(format!("http://{addr}/")).send(),
1124 |status, body| format!("{status}: {body}"),
1125 );
1126 let (status, _) = result.expect("eventually succeeds");
1127 assert_eq!(status.as_u16(), 200);
1128 assert_eq!(calls.load(Ordering::SeqCst), 2, "one retry then success");
1129 }
1130
1131 #[test]
1132 fn retry_http_blocking_4xx_fast_fails_no_retry() {
1133 let (addr, calls) = spawn_oneshot_http_responder(vec![
1134 "HTTP/1.1 404 Not Found\r\nContent-Length: 9\r\n\r\nnot found",
1135 ]);
1136 let client = reqwest::blocking::Client::builder()
1137 .timeout(Duration::from_secs(2))
1138 .build()
1139 .expect("client");
1140 let policy = RetryPolicy {
1141 max_attempts: 5,
1142 base_delay: Duration::from_millis(1),
1143 max_delay: Duration::from_millis(2),
1144 };
1145 let result = retry_http_blocking(
1146 "myscope",
1147 &policy,
1148 SuccessClass::Strict,
1149 |_| client.get(format!("http://{addr}/")).send(),
1150 |status, body| format!("custom error: {status} body={body}"),
1151 );
1152 let err = result.expect_err("4xx must fast-fail");
1153 let chain = format!("{err:#}");
1154 assert!(
1155 chain.contains("custom error"),
1156 "error formatter must be invoked on non-success; got: {chain}"
1157 );
1158 assert!(chain.contains("404"), "status must be in chain: {chain}");
1159 assert_eq!(
1160 calls.load(Ordering::SeqCst),
1161 1,
1162 "4xx must NOT retry (only one connection accepted)"
1163 );
1164 }
1165
1166 #[test]
1167 fn retry_http_blocking_redirect_class_alters_success_predicate() {
1168 let (addr, _calls) = spawn_oneshot_http_responder(vec![
1169 "HTTP/1.1 307 Temporary Redirect\r\nLocation: /next\r\nContent-Length: 0\r\n\r\n",
1170 ]);
1171 let client = reqwest::blocking::Client::builder()
1172 .timeout(Duration::from_secs(2))
1173 .redirect(reqwest::redirect::Policy::none())
1175 .build()
1176 .expect("client");
1177 let policy = RetryPolicy {
1178 max_attempts: 3,
1179 base_delay: Duration::from_millis(1),
1180 max_delay: Duration::from_millis(2),
1181 };
1182 let result = retry_http_blocking(
1183 "test",
1184 &policy,
1185 SuccessClass::AllowRedirects,
1186 |_| client.get(format!("http://{addr}/")).send(),
1187 |_, _| String::from("should not be called on 3xx with AllowRedirects"),
1188 );
1189 let (status, _) = result.expect("3xx is success under AllowRedirects");
1190 assert_eq!(status.as_u16(), 307);
1191 }
1192
1193 #[tokio::test]
1204 async fn retry_http_async_success_returns_first_attempt() {
1205 let (addr, calls) =
1206 spawn_oneshot_http_responder(vec!["HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok"]);
1207 let client = reqwest::Client::builder()
1208 .timeout(Duration::from_secs(2))
1209 .build()
1210 .expect("client");
1211 let policy = RetryPolicy {
1212 max_attempts: 3,
1213 base_delay: Duration::from_millis(1),
1214 max_delay: Duration::from_millis(2),
1215 };
1216 let result = retry_http_async(
1217 "test",
1218 &policy,
1219 SuccessClass::Strict,
1220 |_| client.get(format!("http://{addr}/")).send(),
1221 |_, _| String::from("should not be called on success"),
1222 )
1223 .await;
1224 let resp = result.expect("success");
1225 assert_eq!(resp.status().as_u16(), 200);
1226 let body = resp.text().await.expect("body");
1227 assert_eq!(body, "ok");
1228 assert_eq!(calls.load(Ordering::SeqCst), 1, "single attempt");
1229 }
1230
1231 #[tokio::test]
1232 async fn retry_http_async_retries_5xx_then_succeeds() {
1233 let (addr, calls) = spawn_oneshot_http_responder(vec![
1234 "HTTP/1.1 503 Service Unavailable\r\nContent-Length: 0\r\n\r\n",
1235 "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok",
1236 ]);
1237 let client = reqwest::Client::builder()
1238 .timeout(Duration::from_secs(2))
1239 .build()
1240 .expect("client");
1241 let policy = RetryPolicy {
1242 max_attempts: 3,
1243 base_delay: Duration::from_millis(1),
1244 max_delay: Duration::from_millis(2),
1245 };
1246 let result = retry_http_async(
1247 "test",
1248 &policy,
1249 SuccessClass::Strict,
1250 |_| client.get(format!("http://{addr}/")).send(),
1251 |status, body| format!("{status}: {body}"),
1252 )
1253 .await;
1254 let resp = result.expect("eventually succeeds");
1255 assert_eq!(resp.status().as_u16(), 200);
1256 assert_eq!(calls.load(Ordering::SeqCst), 2, "one retry then success");
1257 }
1258
1259 #[tokio::test]
1260 async fn retry_http_async_4xx_fast_fails_no_retry() {
1261 let (addr, calls) = spawn_oneshot_http_responder(vec![
1262 "HTTP/1.1 404 Not Found\r\nContent-Length: 9\r\n\r\nnot found",
1263 ]);
1264 let client = reqwest::Client::builder()
1265 .timeout(Duration::from_secs(2))
1266 .build()
1267 .expect("client");
1268 let policy = RetryPolicy {
1269 max_attempts: 5,
1270 base_delay: Duration::from_millis(1),
1271 max_delay: Duration::from_millis(2),
1272 };
1273 let result = retry_http_async(
1274 "myscope",
1275 &policy,
1276 SuccessClass::Strict,
1277 |_| client.get(format!("http://{addr}/")).send(),
1278 |status, body| format!("custom error: {status} body={body}"),
1279 )
1280 .await;
1281 let err = result.expect_err("4xx must fast-fail");
1282 let chain = format!("{err:#}");
1283 assert!(
1284 chain.contains("custom error"),
1285 "error formatter must be invoked on non-success; got: {chain}"
1286 );
1287 assert!(chain.contains("404"), "status must be in chain: {chain}");
1288 assert_eq!(
1289 calls.load(Ordering::SeqCst),
1290 1,
1291 "4xx must NOT retry (only one connection accepted)"
1292 );
1293 }
1294
1295 #[tokio::test]
1296 async fn retry_http_async_429_retries_then_succeeds() {
1297 let (addr, calls) = spawn_oneshot_http_responder(vec![
1302 "HTTP/1.1 429 Too Many Requests\r\nContent-Length: 0\r\n\r\n",
1303 "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok",
1304 ]);
1305 let client = reqwest::Client::builder()
1306 .timeout(Duration::from_secs(2))
1307 .build()
1308 .expect("client");
1309 let policy = RetryPolicy {
1310 max_attempts: 3,
1311 base_delay: Duration::from_millis(1),
1312 max_delay: Duration::from_millis(2),
1313 };
1314 let result = retry_http_async(
1315 "test",
1316 &policy,
1317 SuccessClass::Strict,
1318 |_| client.get(format!("http://{addr}/")).send(),
1319 |status, body| format!("{status}: {body}"),
1320 )
1321 .await;
1322 let resp = result.expect("429 retried then success");
1323 assert_eq!(resp.status().as_u16(), 200);
1324 assert_eq!(calls.load(Ordering::SeqCst), 2);
1325 }
1326
1327 const TRANSPORT_FAIL_URL: &str = "http://nonexistent.invalid/";
1351
1352 #[test]
1353 fn retry_http_blocking_transport_error_retries_then_fails() {
1354 let attempts = std::sync::Arc::new(AtomicU32::new(0));
1355 let attempts_inner = attempts.clone();
1356 let client = reqwest::blocking::Client::builder()
1357 .timeout(Duration::from_millis(500))
1358 .build()
1359 .expect("client");
1360 let policy = RetryPolicy {
1361 max_attempts: 3,
1362 base_delay: Duration::from_millis(1),
1363 max_delay: Duration::from_millis(2),
1364 };
1365 let result = retry_http_blocking(
1366 "test-transport",
1367 &policy,
1368 SuccessClass::Strict,
1369 |_| {
1370 attempts_inner.fetch_add(1, Ordering::SeqCst);
1371 client.get(TRANSPORT_FAIL_URL).send()
1372 },
1373 |_, _| String::from("non-success branch should not be reached"),
1374 );
1375 let err = result.expect_err("transport error must surface as Err");
1376 let chain = format!("{err:#}");
1377 assert!(
1378 attempts.load(Ordering::SeqCst) > 1,
1379 "transport error must be retried; got {} attempts; chain={chain}",
1380 attempts.load(Ordering::SeqCst)
1381 );
1382 assert!(
1383 chain.contains("test-transport"),
1384 "label must surface in error chain; got: {chain}"
1385 );
1386 }
1387
1388 #[tokio::test]
1389 async fn retry_http_async_transport_error_retries_then_fails() {
1390 let attempts = std::sync::Arc::new(AtomicU32::new(0));
1391 let attempts_inner = attempts.clone();
1392 let client = reqwest::Client::builder()
1393 .timeout(Duration::from_millis(500))
1394 .build()
1395 .expect("client");
1396 let policy = RetryPolicy {
1397 max_attempts: 3,
1398 base_delay: Duration::from_millis(1),
1399 max_delay: Duration::from_millis(2),
1400 };
1401 let result = retry_http_async(
1402 "test-transport-async",
1403 &policy,
1404 SuccessClass::Strict,
1405 |_| {
1406 attempts_inner.fetch_add(1, Ordering::SeqCst);
1407 client.get(TRANSPORT_FAIL_URL).send()
1408 },
1409 |_, _| String::from("non-success branch should not be reached"),
1410 )
1411 .await;
1412 let err = result.expect_err("transport error must surface as Err");
1413 assert!(
1414 attempts.load(Ordering::SeqCst) > 1,
1415 "transport error must be retried; got {} attempts",
1416 attempts.load(Ordering::SeqCst)
1417 );
1418 let chain = format!("{err:#}");
1419 assert!(
1420 chain.contains("test-transport-async"),
1421 "label must surface in error chain; got: {chain}"
1422 );
1423 }
1424}