1use std::fmt;
18use std::sync::Arc;
19use std::time::Duration;
20
21use serde::{Deserialize, Serialize};
22
23use crate::error::{ProviderDetail, ProviderError, ProviderErrorKind, RetryClass};
24use crate::ids::{AttemptNumber, ModelRef, RequestId};
25use crate::purpose::ModelPurpose;
26use crate::request::ModelRequest;
27use crate::response::ModelResponse;
28use crate::router::{Clock, ProviderCandidate, SystemClock};
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
32#[serde(rename_all = "snake_case")]
33pub enum FallbackStage {
34 PreCommit,
36 PostCommitNarration,
38}
39
40impl FallbackStage {
41 #[must_use]
43 pub const fn as_str(self) -> &'static str {
44 match self {
45 Self::PreCommit => "pre_commit",
46 Self::PostCommitNarration => "post_commit_narration",
47 }
48 }
49
50 #[must_use]
55 pub const fn admits(self, purpose: ModelPurpose) -> bool {
56 match self {
57 Self::PreCommit => true,
58 Self::PostCommitNarration => !purpose.is_critical(),
59 }
60 }
61}
62
63impl fmt::Display for FallbackStage {
64 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65 f.write_str(self.as_str())
66 }
67}
68
69#[async_trait::async_trait]
71pub trait Sleeper: Send + Sync + fmt::Debug {
72 async fn sleep(&self, duration: Duration);
74}
75
76#[derive(Debug, Clone, Copy, Default)]
78pub struct TokioSleeper;
79
80#[async_trait::async_trait]
81impl Sleeper for TokioSleeper {
82 async fn sleep(&self, duration: Duration) {
83 if !duration.is_zero() {
84 tokio::time::sleep(duration).await;
85 }
86 }
87}
88
89#[derive(Debug, Clone, PartialEq, Eq)]
91pub struct RetryPolicy {
92 pub max_attempts: u32,
94 pub initial_backoff: Duration,
96 pub max_backoff: Duration,
98 pub backoff_multiplier: u32,
100 pub jitter: bool,
105 pub honour_retry_after: bool,
107 pub max_retry_after: Duration,
109}
110
111impl RetryPolicy {
112 pub const DEFAULT: Self = Self {
115 max_attempts: 3,
116 initial_backoff: Duration::from_millis(200),
117 max_backoff: Duration::from_secs(5),
118 backoff_multiplier: 4,
119 jitter: false,
120 honour_retry_after: true,
121 max_retry_after: Duration::from_secs(30),
122 };
123
124 pub const NO_RETRY: Self = Self {
126 max_attempts: 1,
127 ..Self::DEFAULT
128 };
129
130 #[must_use]
134 pub fn backoff_for(&self, attempt: AttemptNumber) -> Duration {
135 let step = attempt.get().saturating_sub(1);
136 if step == 0 {
137 return Duration::ZERO;
138 }
139 let factor = self.backoff_multiplier.saturating_pow(step - 1);
140 self.initial_backoff
141 .saturating_mul(factor)
142 .min(self.max_backoff)
143 }
144
145 #[must_use]
154 pub fn delay_for(
155 &self,
156 attempt: AttemptNumber,
157 previous: &ProviderError,
158 request_id: RequestId,
159 ) -> Duration {
160 if self.honour_retry_after
161 && let Some(asked) = previous.retry_after()
162 {
163 return asked.min(self.max_retry_after);
164 }
165 let base = self.backoff_for(attempt);
166 if !self.jitter || base.is_zero() {
167 return base;
168 }
169 let seed = request_id.as_uuid().as_u128() ^ u128::from(attempt.get());
172 let permille = 500 + u64::try_from(seed % 501).unwrap_or(0);
173 Duration::from_nanos(
174 u64::try_from(base.as_nanos().saturating_mul(u128::from(permille)) / 1000)
175 .unwrap_or(u64::MAX),
176 )
177 }
178}
179
180impl Default for RetryPolicy {
181 fn default() -> Self {
182 Self::DEFAULT
183 }
184}
185
186#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
188#[serde(tag = "kind", rename_all = "snake_case")]
189#[non_exhaustive]
190pub enum AttemptOutcome {
191 Succeeded,
193 Retried {
195 code: String,
197 },
198 FellBack {
200 code: String,
202 },
203 Failed {
205 code: String,
207 },
208 Cancelled,
210}
211
212impl AttemptOutcome {
213 #[must_use]
215 pub const fn as_str(&self) -> &'static str {
216 match self {
217 Self::Succeeded => "succeeded",
218 Self::Retried { .. } => "retried",
219 Self::FellBack { .. } => "fell_back",
220 Self::Failed { .. } => "failed",
221 Self::Cancelled => "cancelled",
222 }
223 }
224
225 #[must_use]
227 pub fn to_core_outcome(&self) -> turnframe_core::replay::ProviderAttemptOutcome {
228 use turnframe_core::replay::ProviderAttemptOutcome as Core;
229 match self {
230 Self::Succeeded => Core::Succeeded,
231 Self::Retried { code } | Self::Failed { code } => Core::Failed { code: code.clone() },
232 Self::FellBack { code } => Core::FellBack { code: code.clone() },
233 Self::Cancelled => Core::Cancelled,
234 }
235 }
236}
237
238#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
243pub struct ProviderAttempt {
244 pub attempt: AttemptNumber,
246 pub request_id: RequestId,
248 pub purpose: ModelPurpose,
250 pub stage: FallbackStage,
252 pub model: ModelRef,
254 pub outcome: AttemptOutcome,
256 #[serde(default, skip_serializing_if = "Option::is_none")]
258 pub class: Option<RetryClass>,
259 pub latency: Duration,
261 #[serde(default, skip_serializing_if = "Option::is_none")]
263 pub input_tokens: Option<u64>,
264 #[serde(default, skip_serializing_if = "Option::is_none")]
266 pub output_tokens: Option<u64>,
267 #[serde(default, skip_serializing_if = "Option::is_none")]
272 pub temperature: Option<f32>,
273 #[serde(default, skip_serializing_if = "Vec::is_empty")]
278 pub finish_reasons: Vec<String>,
279}
280
281impl ProviderAttempt {
282 #[must_use]
284 pub fn to_core_record(&self) -> turnframe_core::replay::ProviderAttemptRecord {
285 turnframe_core::replay::ProviderAttemptRecord {
286 attempt: self.attempt.get(),
287 purpose: self.purpose.as_str().to_owned(),
288 provider_key: self.model.provider.clone(),
289 model_key: self.model.model.clone(),
290 request_id: self.request_id.to_string(),
291 prompt_version: None,
292 prompt_ref: None,
295 outcome: self.outcome.to_core_outcome(),
296 latency_ms: u64::try_from(self.latency.as_millis()).ok(),
297 input_tokens: self.input_tokens,
298 output_tokens: self.output_tokens,
299 temperature: self.temperature,
300 finish_reasons: self.finish_reasons.clone(),
301 }
302 }
303}
304
305#[derive(Debug, Clone, PartialEq)]
307pub struct FallbackOutcome {
308 pub response: ModelResponse,
310 pub attempts: Vec<ProviderAttempt>,
312}
313
314impl FallbackOutcome {
315 #[must_use]
317 pub fn served_by(&self) -> ModelRef {
318 self.response.reference()
319 }
320
321 #[must_use]
323 pub fn fell_back(&self) -> bool {
324 self.attempts
325 .iter()
326 .any(|attempt| matches!(attempt.outcome, AttemptOutcome::FellBack { .. }))
327 }
328}
329
330#[derive(Debug, Clone, PartialEq, thiserror::Error)]
332#[error("{error} after {} attempt(s)", attempts.len())]
333pub struct FallbackFailure {
334 pub error: ProviderError,
336 pub attempts: Vec<ProviderAttempt>,
338}
339
340impl FallbackFailure {
341 #[must_use]
344 pub fn unattempted(error: ProviderError) -> Self {
345 Self {
346 error,
347 attempts: Vec::new(),
348 }
349 }
350}
351
352#[derive(Clone)]
357pub struct FallbackOptions {
358 pub retry: RetryPolicy,
360 pub allow_provider_fallback: bool,
364 pub sleeper: Arc<dyn Sleeper>,
366 pub clock: Arc<dyn Clock>,
368}
369
370impl FallbackOptions {
371 #[must_use]
374 pub fn new() -> Self {
375 Self {
376 retry: RetryPolicy::DEFAULT,
377 allow_provider_fallback: true,
378 sleeper: Arc::new(TokioSleeper),
379 clock: Arc::new(SystemClock),
380 }
381 }
382
383 #[must_use]
385 pub fn with_retry(mut self, retry: RetryPolicy) -> Self {
386 self.retry = retry;
387 self
388 }
389
390 #[must_use]
392 pub fn without_provider_fallback(mut self) -> Self {
393 self.allow_provider_fallback = false;
394 self
395 }
396
397 #[must_use]
399 pub fn with_sleeper<S: Sleeper + 'static>(mut self, sleeper: Arc<S>) -> Self {
400 self.sleeper = sleeper;
401 self
402 }
403
404 #[must_use]
406 pub fn with_clock<C: Clock + 'static>(mut self, clock: Arc<C>) -> Self {
407 self.clock = clock;
408 self
409 }
410}
411
412impl Default for FallbackOptions {
413 fn default() -> Self {
414 Self::new()
415 }
416}
417
418impl fmt::Debug for FallbackOptions {
419 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
420 f.debug_struct("FallbackOptions")
421 .field("retry", &self.retry)
422 .field("allow_provider_fallback", &self.allow_provider_fallback)
423 .finish_non_exhaustive()
424 }
425}
426
427#[allow(clippy::result_large_err)]
467pub async fn execute_with_fallback(
468 candidates: &[ProviderCandidate],
469 request: &ModelRequest,
470 stage: FallbackStage,
471 options: &FallbackOptions,
472) -> Result<FallbackOutcome, FallbackFailure> {
473 if !stage.admits(request.purpose) {
474 return Err(FallbackFailure::unattempted(
477 ProviderError::invalid_request("critical_purpose_after_commit"),
478 ));
479 }
480 if candidates.is_empty() {
481 return Err(FallbackFailure::unattempted(ProviderError::other(
482 "no_candidate",
483 )));
484 }
485
486 let usable = if options.allow_provider_fallback {
487 candidates
488 } else {
489 &candidates[..1]
490 };
491
492 let mut attempts: Vec<ProviderAttempt> = Vec::new();
493 let mut counter = AttemptNumber::FIRST;
494 let mut last_error = ProviderError::other("no_candidate");
495
496 for (index, candidate) in usable.iter().enumerate() {
497 let model = candidate.reference();
498 let is_last_candidate = index + 1 == usable.len();
499 let mut within = AttemptNumber::FIRST;
500
501 loop {
502 let started = options.clock.now();
503 let result = candidate.provider.generate(request.clone()).await;
504 let latency = elapsed_since(options.clock.now(), started);
505
506 match result {
507 Ok(response) => {
508 tracing::debug!(
509 request_id = %request.request_id,
510 purpose = request.purpose.as_str(),
511 stage = stage.as_str(),
512 provider = model.provider.as_str(),
513 model = model.model.as_str(),
514 attempt = counter.get(),
515 latency_ms = u64::try_from(latency.as_millis()).unwrap_or(u64::MAX),
516 "provider attempt succeeded"
517 );
518 attempts.push(ProviderAttempt {
519 attempt: counter,
520 request_id: request.request_id,
521 purpose: request.purpose,
522 stage,
523 model,
524 outcome: AttemptOutcome::Succeeded,
525 class: None,
526 latency,
527 input_tokens: Some(response.usage.input),
528 output_tokens: Some(response.usage.output),
529 temperature: request.temperature,
530 finish_reasons: vec![response.finish.as_str().to_owned()],
531 });
532 return Ok(FallbackOutcome { response, attempts });
533 }
534 Err(error) => {
535 let error = error.with_model(&model);
536 let class = error.retry_class();
537 let code = error.kind().as_str().to_owned();
538 let has_attempts_left = within.get() < options.retry.max_attempts;
539 let will_retry = class.allows_same_provider() && has_attempts_left;
540 let will_fall_back =
541 !will_retry && class.allows_another_candidate() && !is_last_candidate;
542
543 let outcome = if matches!(error.kind(), ProviderErrorKind::Cancelled) {
544 AttemptOutcome::Cancelled
545 } else if will_retry {
546 AttemptOutcome::Retried { code }
547 } else if will_fall_back {
548 AttemptOutcome::FellBack { code }
549 } else {
550 AttemptOutcome::Failed { code }
551 };
552 tracing::warn!(
553 request_id = %request.request_id,
554 purpose = request.purpose.as_str(),
555 stage = stage.as_str(),
556 provider = model.provider.as_str(),
557 model = model.model.as_str(),
558 attempt = counter.get(),
559 error = error.kind().as_str(),
560 class = class.as_str(),
561 outcome = outcome.as_str(),
562 detail = error.detail().map(ProviderDetail::as_str),
568 "provider attempt failed"
569 );
570 attempts.push(ProviderAttempt {
571 attempt: counter,
572 request_id: request.request_id,
573 purpose: request.purpose,
574 stage,
575 model: model.clone(),
576 outcome,
577 class: Some(class),
578 latency,
579 input_tokens: None,
580 output_tokens: None,
581 temperature: request.temperature,
582 finish_reasons: Vec::new(),
583 });
584 counter = counter.next();
585
586 if will_retry {
587 within = within.next();
588 let delay = options.retry.delay_for(within, &error, request.request_id);
589 options.sleeper.sleep(delay).await;
590 continue;
591 }
592 last_error = error;
593 if class == RetryClass::Fatal {
594 return Err(FallbackFailure {
596 error: last_error,
597 attempts,
598 });
599 }
600 break;
601 }
602 }
603 }
604 }
605
606 Err(FallbackFailure {
607 error: last_error,
608 attempts,
609 })
610}
611
612fn elapsed_since(
614 now: chrono::DateTime<chrono::Utc>,
615 started: chrono::DateTime<chrono::Utc>,
616) -> Duration {
617 (now - started).to_std().unwrap_or(Duration::ZERO)
618}
619
620#[cfg(test)]
621mod tests {
622 use super::*;
623 use crate::capabilities::ProviderCapabilities;
624 use crate::error::RetryClass;
625 use crate::provider::ModelProvider;
626 use crate::response::FinishReason;
627 use crate::testing::{ImmediateSleeper, ManualClock, StaticProvider};
628 use serde_json::json;
629
630 fn options(sleeper: &Arc<ImmediateSleeper>) -> FallbackOptions {
631 FallbackOptions::new()
632 .with_sleeper(Arc::clone(sleeper))
633 .with_clock(Arc::new(ManualClock::at_epoch()))
634 }
635
636 fn request(purpose: ModelPurpose) -> ModelRequest {
637 ModelRequest::new(purpose).with_request_id(RequestId::nil())
638 }
639
640 #[tokio::test]
641 async fn a_retryable_failure_retries_the_same_provider() {
642 let provider = Arc::new(
643 StaticProvider::new("a", "m")
644 .failing(vec![
645 ProviderError::server(Some(503)),
646 ProviderError::timeout(),
647 ])
648 .answering_text("ok"),
649 );
650 let sleeper = Arc::new(ImmediateSleeper::new());
651 let outcome = execute_with_fallback(
652 &[StaticProvider::candidate(Arc::clone(&provider))],
653 &request(ModelPurpose::Acknowledge),
654 FallbackStage::PreCommit,
655 &options(&sleeper),
656 )
657 .await
658 .unwrap();
659
660 assert_eq!(outcome.attempts.len(), 3);
661 assert!(matches!(
662 outcome.attempts[0].outcome,
663 AttemptOutcome::Retried { .. }
664 ));
665 assert!(matches!(
666 outcome.attempts[2].outcome,
667 AttemptOutcome::Succeeded
668 ));
669 assert!(!outcome.fell_back());
670 assert_eq!(outcome.served_by(), ModelRef::new("a", "m"));
671 assert_eq!(sleeper.slept().len(), 2, "one wait per retry");
672 }
673
674 #[tokio::test]
675 async fn exhausting_a_candidate_moves_to_the_next_one() {
676 let broken =
677 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::timeout()));
678 let good = Arc::new(StaticProvider::new("b", "m").answering_text("ciao"));
679 let sleeper = Arc::new(ImmediateSleeper::new());
680 let outcome = execute_with_fallback(
681 &[
682 StaticProvider::candidate(Arc::clone(&broken)),
683 StaticProvider::candidate(Arc::clone(&good)),
684 ],
685 &request(ModelPurpose::Acknowledge),
686 FallbackStage::PreCommit,
687 &options(&sleeper),
688 )
689 .await
690 .unwrap();
691
692 assert_eq!(
693 broken.call_count(),
694 3,
695 "max_attempts on the first candidate"
696 );
697 assert_eq!(good.call_count(), 1);
698 assert_eq!(outcome.attempts.len(), 4);
699 assert!(matches!(
700 outcome.attempts[2].outcome,
701 AttemptOutcome::FellBack { .. }
702 ));
703 assert!(outcome.fell_back());
704 assert_eq!(outcome.response.text(), "ciao");
705 assert_eq!(outcome.served_by(), ModelRef::new("b", "m"));
706 }
707
708 #[tokio::test]
709 async fn a_fallback_class_moves_on_without_retrying() {
710 let unauthorized =
711 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::authentication()));
712 let good = Arc::new(StaticProvider::new("b", "m").answering_text("ok"));
713 let sleeper = Arc::new(ImmediateSleeper::new());
714 let outcome = execute_with_fallback(
715 &[
716 StaticProvider::candidate(Arc::clone(&unauthorized)),
717 StaticProvider::candidate(Arc::clone(&good)),
718 ],
719 &request(ModelPurpose::Acknowledge),
720 FallbackStage::PreCommit,
721 &options(&sleeper),
722 )
723 .await
724 .unwrap();
725 assert_eq!(
726 unauthorized.call_count(),
727 1,
728 "no point retrying bad credentials"
729 );
730 assert_eq!(outcome.attempts.len(), 2);
731 assert!(sleeper.slept().is_empty());
732 }
733
734 #[tokio::test]
735 async fn a_fatal_class_abandons_the_stage_at_once() {
736 let refusing =
737 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::refusal()));
738 let good = Arc::new(StaticProvider::new("b", "m").answering_text("ok"));
739 let sleeper = Arc::new(ImmediateSleeper::new());
740 let failure = execute_with_fallback(
741 &[
742 StaticProvider::candidate(Arc::clone(&refusing)),
743 StaticProvider::candidate(Arc::clone(&good)),
744 ],
745 &request(ModelPurpose::Acknowledge),
746 FallbackStage::PreCommit,
747 &options(&sleeper),
748 )
749 .await
750 .unwrap_err();
751
752 assert_eq!(refusing.call_count(), 1);
753 assert_eq!(good.call_count(), 0, "a refusal is not shopped around");
754 assert_eq!(failure.attempts.len(), 1);
755 assert_eq!(failure.error.retry_class(), RetryClass::Fatal);
756 assert!(failure.to_string().contains("1 attempt"), "{failure}");
757 }
758
759 #[tokio::test]
760 async fn a_critical_purpose_is_refused_after_commit() {
761 let provider = Arc::new(StaticProvider::new("a", "m").answering_text("ok"));
762 let sleeper = Arc::new(ImmediateSleeper::new());
763 for purpose in [ModelPurpose::Extract, ModelPurpose::Investigate] {
764 let failure = execute_with_fallback(
765 &[StaticProvider::candidate(Arc::clone(&provider))],
766 &request(purpose),
767 FallbackStage::PostCommitNarration,
768 &options(&sleeper),
769 )
770 .await
771 .unwrap_err();
772 assert!(failure.attempts.is_empty());
773 assert_eq!(provider.call_count(), 0, "the model was never called");
774 assert!(
775 failure
776 .error
777 .code()
778 .is_some_and(|code| code.as_str() == "critical_purpose_after_commit")
779 );
780 }
781 assert!(
783 execute_with_fallback(
784 &[StaticProvider::candidate(Arc::clone(&provider))],
785 &request(ModelPurpose::Acknowledge),
786 FallbackStage::PostCommitNarration,
787 &options(&sleeper),
788 )
789 .await
790 .is_ok()
791 );
792 }
793
794 #[tokio::test]
795 async fn partial_outputs_from_different_providers_are_never_merged() {
796 let half = Arc::new(
798 StaticProvider::new("a", "m")
799 .replying_once_json(json!({"acts": ["only_from_a"]}))
800 .always_failing(ProviderError::server(None)),
801 );
802 let whole =
803 Arc::new(StaticProvider::new("b", "m").answering_json(json!({"acts": ["from_b"]})));
804 let sleeper = Arc::new(ImmediateSleeper::new());
805
806 let _ = half.generate(request(ModelPurpose::Extract)).await.unwrap();
808
809 let outcome = execute_with_fallback(
810 &[
811 StaticProvider::candidate(Arc::clone(&half)),
812 StaticProvider::candidate(Arc::clone(&whole)),
813 ],
814 &request(ModelPurpose::Extract),
815 FallbackStage::PreCommit,
816 &options(&sleeper),
817 )
818 .await
819 .unwrap();
820
821 assert_eq!(outcome.response.provider.as_str(), "b");
822 assert_eq!(outcome.response.content.len(), 1);
823 assert_eq!(
824 outcome.response.text(),
825 json!({"acts": ["from_b"]}).to_string()
826 );
827 assert!(
828 !outcome.response.text().contains("only_from_a"),
829 "no content crossed providers"
830 );
831 }
832
833 #[tokio::test]
834 async fn the_request_id_is_stable_across_every_attempt() {
835 let broken =
836 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::timeout()));
837 let good = Arc::new(StaticProvider::new("b", "m").answering_text("ok"));
838 let sleeper = Arc::new(ImmediateSleeper::new());
839 let request = request(ModelPurpose::Acknowledge);
840 let outcome = execute_with_fallback(
841 &[
842 StaticProvider::candidate(Arc::clone(&broken)),
843 StaticProvider::candidate(Arc::clone(&good)),
844 ],
845 &request,
846 FallbackStage::PreCommit,
847 &options(&sleeper),
848 )
849 .await
850 .unwrap();
851
852 for attempt in &outcome.attempts {
853 assert_eq!(attempt.request_id, request.request_id);
854 }
855 for seen in broken.calls() {
856 assert_eq!(seen.request_id, request.request_id);
857 }
858 assert_eq!(outcome.response.request_id, request.request_id);
859 }
860
861 #[tokio::test]
862 async fn attempts_convert_into_core_replay_records() {
863 let broken =
864 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::timeout()));
865 let good = Arc::new(
866 StaticProvider::new("b", "m")
867 .answering_text("ok")
868 .with_usage(crate::response::TokenUsage::new(12, 3)),
869 );
870 let sleeper = Arc::new(ImmediateSleeper::new());
871 let outcome = execute_with_fallback(
872 &[
873 StaticProvider::candidate(broken),
874 StaticProvider::candidate(good),
875 ],
876 &request(ModelPurpose::Extract),
877 FallbackStage::PreCommit,
878 &options(&sleeper),
879 )
880 .await
881 .unwrap();
882
883 let records: Vec<_> = outcome
884 .attempts
885 .iter()
886 .map(ProviderAttempt::to_core_record)
887 .collect();
888 assert_eq!(records.len(), 4);
889 assert_eq!(records[0].attempt, 1);
890 assert_eq!(records[0].provider_key.as_str(), "a");
891 assert_eq!(records[0].purpose, "extract");
892 assert_eq!(
893 records[2].outcome,
894 turnframe_core::replay::ProviderAttemptOutcome::FellBack {
895 code: "timeout".to_owned()
896 }
897 );
898 assert_eq!(
899 records[3].outcome,
900 turnframe_core::replay::ProviderAttemptOutcome::Succeeded
901 );
902 assert_eq!(records[3].input_tokens, Some(12));
903 assert_eq!(records[3].output_tokens, Some(3));
904 assert!(
906 records
907 .iter()
908 .all(|r| r.request_id == RequestId::nil().to_string())
909 );
910 }
911
912 #[tokio::test]
913 async fn an_empty_candidate_list_fails_without_attempting() {
914 let sleeper = Arc::new(ImmediateSleeper::new());
915 let failure = execute_with_fallback(
916 &[],
917 &request(ModelPurpose::Acknowledge),
918 FallbackStage::PreCommit,
919 &options(&sleeper),
920 )
921 .await
922 .unwrap_err();
923 assert!(failure.attempts.is_empty());
924 assert_eq!(
925 failure.error.code().map(|c| c.as_str().to_owned()),
926 Some("no_candidate".to_owned())
927 );
928 }
929
930 #[tokio::test]
931 async fn provider_fallback_can_be_switched_off() {
932 let broken =
933 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::authentication()));
934 let good = Arc::new(StaticProvider::new("b", "m").answering_text("ok"));
935 let sleeper = Arc::new(ImmediateSleeper::new());
936 let failure = execute_with_fallback(
937 &[
938 StaticProvider::candidate(Arc::clone(&broken)),
939 StaticProvider::candidate(Arc::clone(&good)),
940 ],
941 &request(ModelPurpose::Acknowledge),
942 FallbackStage::PreCommit,
943 &options(&sleeper).without_provider_fallback(),
944 )
945 .await
946 .unwrap_err();
947 assert_eq!(good.call_count(), 0);
948 assert_eq!(failure.attempts.len(), 1);
949 }
950
951 #[tokio::test]
952 async fn retry_after_is_honoured_and_capped() {
953 let provider = Arc::new(
954 StaticProvider::new("a", "m")
955 .failing_once(ProviderError::rate_limited(Some(Duration::from_secs(2))))
956 .answering_text("ok"),
957 );
958 let sleeper = Arc::new(ImmediateSleeper::new());
959 let outcome = execute_with_fallback(
960 &[StaticProvider::candidate(Arc::clone(&provider))],
961 &request(ModelPurpose::Acknowledge),
962 FallbackStage::PreCommit,
963 &options(&sleeper),
964 )
965 .await
966 .unwrap();
967 assert!(outcome.attempts[0].class == Some(RetryClass::RetryAfter));
968 assert_eq!(sleeper.slept(), vec![Duration::from_secs(2)]);
969
970 let capped = RetryPolicy {
971 max_retry_after: Duration::from_millis(500),
972 ..RetryPolicy::DEFAULT
973 };
974 assert_eq!(
975 capped.delay_for(
976 AttemptNumber(2),
977 &ProviderError::rate_limited(Some(Duration::from_secs(90))),
978 RequestId::nil()
979 ),
980 Duration::from_millis(500)
981 );
982 }
983
984 #[test]
985 fn backoff_grows_and_is_capped_with_jitter_off_by_default() {
986 let policy = RetryPolicy::DEFAULT;
987 assert!(!policy.jitter, "determinism is the default");
988 assert_eq!(policy.backoff_for(AttemptNumber::FIRST), Duration::ZERO);
989 assert_eq!(
990 policy.backoff_for(AttemptNumber(2)),
991 Duration::from_millis(200)
992 );
993 assert_eq!(
994 policy.backoff_for(AttemptNumber(3)),
995 Duration::from_millis(800)
996 );
997 assert_eq!(policy.backoff_for(AttemptNumber(9)), policy.max_backoff);
998 assert_eq!(RetryPolicy::NO_RETRY.max_attempts, 1);
999 assert_eq!(RetryPolicy::default(), RetryPolicy::DEFAULT);
1000 }
1001
1002 #[test]
1003 fn jitter_is_derived_from_the_request_id_so_a_replay_repeats_it() {
1004 let policy = RetryPolicy {
1005 jitter: true,
1006 honour_retry_after: false,
1007 ..RetryPolicy::DEFAULT
1008 };
1009 let error = ProviderError::timeout();
1010 let id = RequestId::nil();
1011 let first = policy.delay_for(AttemptNumber(2), &error, id);
1012 let again = policy.delay_for(AttemptNumber(2), &error, id);
1013 assert_eq!(first, again, "same call, same schedule");
1014 let base = policy.backoff_for(AttemptNumber(2));
1015 assert!(first >= base / 2 && first <= base, "{first:?} vs {base:?}");
1016 }
1017
1018 #[test]
1019 fn the_stage_decides_which_purposes_may_run() {
1020 assert!(FallbackStage::PreCommit.admits(ModelPurpose::Extract));
1021 assert!(FallbackStage::PreCommit.admits(ModelPurpose::Acknowledge));
1022 assert!(!FallbackStage::PostCommitNarration.admits(ModelPurpose::Extract));
1023 assert!(!FallbackStage::PostCommitNarration.admits(ModelPurpose::Investigate));
1024 assert!(FallbackStage::PostCommitNarration.admits(ModelPurpose::Answer));
1025 assert_eq!(FallbackStage::PreCommit.to_string(), "pre_commit");
1026 }
1027
1028 #[tokio::test]
1029 async fn a_cancelled_attempt_is_recorded_as_cancelled() {
1030 let provider =
1031 Arc::new(StaticProvider::new("a", "m").always_failing(ProviderError::cancelled()));
1032 let sleeper = Arc::new(ImmediateSleeper::new());
1033 let failure = execute_with_fallback(
1034 &[StaticProvider::candidate(provider)],
1035 &request(ModelPurpose::Acknowledge),
1036 FallbackStage::PreCommit,
1037 &options(&sleeper),
1038 )
1039 .await
1040 .unwrap_err();
1041 assert_eq!(failure.attempts[0].outcome, AttemptOutcome::Cancelled);
1042 assert_eq!(
1043 failure.attempts[0].to_core_record().outcome,
1044 turnframe_core::replay::ProviderAttemptOutcome::Cancelled
1045 );
1046 }
1047
1048 #[test]
1049 fn options_render_without_leaking_their_collaborators() {
1050 let rendered = format!("{:?}", FallbackOptions::new());
1051 assert!(rendered.contains("FallbackOptions"));
1052 assert!(rendered.contains("allow_provider_fallback"));
1053 }
1054
1055 #[tokio::test]
1056 async fn capabilities_of_a_candidate_are_the_profiles_own() {
1057 let provider = Arc::new(
1058 StaticProvider::new("a", "m")
1059 .with_capabilities(ProviderCapabilities::minimal().with_streaming(true))
1060 .answering_text("ok"),
1061 );
1062 let candidate = StaticProvider::candidate(provider);
1063 assert!(candidate.profile.capabilities.streaming);
1064 assert!(candidate.healthy);
1065 let response = candidate
1066 .provider
1067 .generate(request(ModelPurpose::Acknowledge))
1068 .await
1069 .unwrap();
1070 assert_eq!(response.finish, FinishReason::Stop);
1071 }
1072}