Skip to main content

turnframe_provider/
fallback.rs

1//! Retry and fallback across candidates (spec §20.7, invariant I17).
2//!
3//! [`execute_with_fallback`] walks the candidates a router produced, retries within one where
4//! the error class allows it, moves on where it does not, and records every attempt. Where in
5//! the turn the call happens decides whether moving on is safe, and nothing in a
6//! [`ModelRequest`] reveals it, so [`FallbackStage`] is a required positional argument:
7//!
8//! * [`FallbackStage::PreCommit`]: understanding a message, before any command executes;
9//!   retries and fallback are free, because nothing has happened yet.
10//! * [`FallbackStage::PostCommitNarration`]: the reply, after the commit, whose facts the
11//!   events already fix. A critical purpose passed here is refused outright (I17).
12//!
13//! A [`FallbackOutcome`] carries one [`ModelResponse`] from one provider: nothing is ever
14//! merged across candidates, because an answer half from one model and half from another is
15//! an answer nobody reviewed.
16
17use 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/// Where in the turn a call is being made (invariant I17).
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
32#[serde(rename_all = "snake_case")]
33pub enum FallbackStage {
34    /// Before any command executes.
35    PreCommit,
36    /// After events are committed, for answering and narration only.
37    PostCommitNarration,
38}
39
40impl FallbackStage {
41    /// Stable snake-case label.
42    #[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    /// Returns `true` when `purpose` may run at this stage.
51    ///
52    /// A [critical](ModelPurpose::is_critical) purpose — one whose output can
53    /// become commands or reads — may only run before commit.
54    #[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/// Waits between attempts. Injectable so tests do not sleep.
70#[async_trait::async_trait]
71pub trait Sleeper: Send + Sync + fmt::Debug {
72    /// Waits for `duration`.
73    async fn sleep(&self, duration: Duration);
74}
75
76/// Sleeps on the Tokio timer.
77#[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/// How hard to try one candidate before moving on (spec §20.7).
90#[derive(Debug, Clone, PartialEq, Eq)]
91pub struct RetryPolicy {
92    /// Attempts per candidate, including the first. `1` disables retries.
93    pub max_attempts: u32,
94    /// Delay before the second attempt.
95    pub initial_backoff: Duration,
96    /// Ceiling on the computed delay.
97    pub max_backoff: Duration,
98    /// Multiplier applied per additional attempt.
99    pub backoff_multiplier: u32,
100    /// Whether to spread delays. **Off by default**: a fixed schedule makes a
101    /// replayed turn reproduce the same timing, and the jitter that is
102    /// available is derived from the request id rather than from randomness,
103    /// so it stays reproducible too.
104    pub jitter: bool,
105    /// Whether a provider's own `Retry-After` overrides the computed backoff.
106    pub honour_retry_after: bool,
107    /// Ceiling on a honoured `Retry-After`, so a provider cannot park a turn.
108    pub max_retry_after: Duration,
109}
110
111impl RetryPolicy {
112    /// Three attempts, 200 ms growing by four, capped at five seconds, no
113    /// jitter, honouring `Retry-After` up to thirty seconds.
114    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    /// One attempt per candidate: fall back rather than retry.
125    pub const NO_RETRY: Self = Self {
126        max_attempts: 1,
127        ..Self::DEFAULT
128    };
129
130    /// The delay before attempt number `attempt` of a candidate.
131    ///
132    /// `attempt` is 1-based, so [`AttemptNumber::FIRST`] has no delay.
133    #[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    /// The delay to wait before `attempt`, given what went wrong last time.
146    ///
147    /// A [`RetryAfter`](RetryClass::RetryAfter) failure that carried a delay
148    /// wins over the computed backoff when
149    /// [`honour_retry_after`](Self::honour_retry_after) is set, capped by
150    /// [`max_retry_after`](Self::max_retry_after). Jitter, when enabled, is
151    /// derived from `request_id` so the same turn replays with the same
152    /// timing.
153    #[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        // Deterministic spread in [50%, 100%] of the computed backoff, seeded
170        // by the stable request id and the attempt number.
171        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/// How one attempt ended.
187#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
188#[serde(tag = "kind", rename_all = "snake_case")]
189#[non_exhaustive]
190pub enum AttemptOutcome {
191    /// The provider answered.
192    Succeeded,
193    /// It failed and the same provider was tried again.
194    Retried {
195        /// The failure family.
196        code: String,
197    },
198    /// It failed and routing moved to the next candidate.
199    FellBack {
200        /// The failure family.
201        code: String,
202    },
203    /// It failed and the stage was abandoned.
204    Failed {
205        /// The failure family.
206        code: String,
207    },
208    /// The caller cancelled.
209    Cancelled,
210}
211
212impl AttemptOutcome {
213    /// Stable snake-case label.
214    #[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    /// Maps onto the outcome the core replay record stores.
226    #[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/// One recorded model call (spec §20.7: "record every provider attempt").
239///
240/// `Eq` is deliberately absent: the record carries the requested temperature,
241/// and a float has no total equality.
242#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
243pub struct ProviderAttempt {
244    /// 1-based position within the whole stage, across candidates.
245    pub attempt: AttemptNumber,
246    /// The stable id sent on every attempt of the logical call.
247    pub request_id: RequestId,
248    /// Why the call was made.
249    pub purpose: ModelPurpose,
250    /// Where in the turn it happened.
251    pub stage: FallbackStage,
252    /// Which profile was called.
253    pub model: ModelRef,
254    /// How it ended.
255    pub outcome: AttemptOutcome,
256    /// How the failure was classified, when it failed.
257    #[serde(default, skip_serializing_if = "Option::is_none")]
258    pub class: Option<RetryClass>,
259    /// Wall-clock duration of the attempt.
260    pub latency: Duration,
261    /// Tokens reported, when the attempt succeeded and the provider reported.
262    #[serde(default, skip_serializing_if = "Option::is_none")]
263    pub input_tokens: Option<u64>,
264    /// Tokens generated, when reported.
265    #[serde(default, skip_serializing_if = "Option::is_none")]
266    pub output_tokens: Option<u64>,
267    /// Sampling temperature the request asked for, when it set one.
268    ///
269    /// Recorded as sent, so an audit reads the value that produced the answer
270    /// rather than a default filled in later.
271    #[serde(default, skip_serializing_if = "Option::is_none")]
272    pub temperature: Option<f32>,
273    /// Finish reasons the provider reported, verbatim.
274    ///
275    /// Not normalized: providers spell them differently, and the difference is
276    /// often exactly what the audit is being read for.
277    #[serde(default, skip_serializing_if = "Vec::is_empty")]
278    pub finish_reasons: Vec<String>,
279}
280
281impl ProviderAttempt {
282    /// Converts into the record the core replay log stores (I20).
283    #[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            // Set by the runtime, which is the layer that knows which prompt
293            // source produced the instructions this call carried.
294            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/// A successful stage, with the trail that got there.
306#[derive(Debug, Clone, PartialEq)]
307pub struct FallbackOutcome {
308    /// The answer, from exactly one provider. Never assembled from several.
309    pub response: ModelResponse,
310    /// Every attempt, in order, the successful one last.
311    pub attempts: Vec<ProviderAttempt>,
312}
313
314impl FallbackOutcome {
315    /// The profile that answered.
316    #[must_use]
317    pub fn served_by(&self) -> ModelRef {
318        self.response.reference()
319    }
320
321    /// Returns `true` when more than one profile was called.
322    #[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/// A stage that ran out of candidates.
331#[derive(Debug, Clone, PartialEq, thiserror::Error)]
332#[error("{error} after {} attempt(s)", attempts.len())]
333pub struct FallbackFailure {
334    /// The failure of the last attempt, or the reason nothing was attempted.
335    pub error: ProviderError,
336    /// Every attempt, in order.
337    pub attempts: Vec<ProviderAttempt>,
338}
339
340impl FallbackFailure {
341    /// A failure with no attempt behind it (an empty candidate list, a stage
342    /// that refused the purpose).
343    #[must_use]
344    pub fn unattempted(error: ProviderError) -> Self {
345        Self {
346            error,
347            attempts: Vec::new(),
348        }
349    }
350}
351
352/// Knobs that are not the stage.
353///
354/// The stage is deliberately absent: it is a positional argument of
355/// [`execute_with_fallback`] so it cannot be defaulted away.
356#[derive(Clone)]
357pub struct FallbackOptions {
358    /// How hard to try each candidate.
359    pub retry: RetryPolicy,
360    /// Whether the walk may move past the first candidate at all. Setting it
361    /// to `false` restricts a stage to one profile without changing the
362    /// router's output.
363    pub allow_provider_fallback: bool,
364    /// Who waits between attempts.
365    pub sleeper: Arc<dyn Sleeper>,
366    /// Where attempt latencies come from.
367    pub clock: Arc<dyn Clock>,
368}
369
370impl FallbackOptions {
371    /// The default policy, sleeping on the Tokio timer and timing on the wall
372    /// clock.
373    #[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    /// Sets the retry policy.
384    #[must_use]
385    pub fn with_retry(mut self, retry: RetryPolicy) -> Self {
386        self.retry = retry;
387        self
388    }
389
390    /// Restricts the stage to the first candidate.
391    #[must_use]
392    pub fn without_provider_fallback(mut self) -> Self {
393        self.allow_provider_fallback = false;
394        self
395    }
396
397    /// Injects a sleeper.
398    #[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    /// Injects a clock.
405    #[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/// Runs `request` against `candidates` until one answers.
428///
429/// Each candidate gets up to [`RetryPolicy::max_attempts`]: `Retry` and `RetryAfter` wait, a
430/// `Fallback` moves on, `Fatal` abandons the stage; every attempt reuses the request id.
431///
432/// # Errors
433///
434/// [`FallbackFailure`], with the last error and the whole trail, when there is no
435/// candidate, the stage refuses the request's purpose (I17), or every candidate failed.
436///
437/// ```
438/// use std::sync::Arc;
439/// use turnframe_provider::prelude::*;
440/// use turnframe_provider::testing::{ImmediateSleeper, StaticProvider};
441///
442/// # futures::executor::block_on(async {
443/// let flaky = Arc::new(
444///     StaticProvider::new("a", "m").failing_once(ProviderError::server(Some(503))),
445/// );
446/// let candidates = vec![StaticProvider::candidate(Arc::clone(&flaky))];
447/// let options = FallbackOptions::new().with_sleeper(Arc::new(ImmediateSleeper::new()));
448///
449/// let outcome = execute_with_fallback(
450///     &candidates,
451///     &ModelRequest::new(ModelPurpose::Acknowledge),
452///     FallbackStage::PreCommit,
453///     &options,
454/// )
455/// .await
456/// .unwrap();
457///
458/// assert_eq!(outcome.attempts.len(), 2, "one failure, then one success");
459/// assert!(!outcome.fell_back(), "the same provider answered");
460/// # });
461/// ```
462// The failure deliberately carries the whole attempt trail (spec §20.7:
463// "record every provider attempt"), which puts the `Err` variant over clippy's
464// size threshold. Boxing it would move an allocation onto a path that runs once
465// per stage, to save copying a struct that is returned once per stage.
466#[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        // I17: re-running a mutation plan after effects may have committed is
475        // the failure this whole module exists to prevent.
476        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                        // The endpoint's own sentence, sanitized by the
563                        // adapter. Without it a malformed request this library
564                        // sent and a provider that is genuinely down are the
565                        // same line, and only one of them is anybody's bug to
566                        // fix.
567                        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                        // Neither another try nor another provider can help.
595                        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
612/// Duration between two clock readings, floored at zero.
613fn 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        // Narration is exactly what the post-commit stage is for.
782        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        // `a` produces half a plan before failing; `b` produces a whole one.
797        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        // Drain `a`'s single good answer so the next call fails.
807        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        // Every record names the same logical call.
905        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}