Skip to main content

tokio_prompt_orchestrator/
model_fallback.rs

1//! Model fallback chains for resilience — automatic failover across model tiers.
2
3use std::collections::HashMap;
4use std::sync::{Arc, Mutex};
5use std::time::{Duration, Instant};
6
7/// Errors that can be returned by a model inference call.
8#[derive(Debug, Clone, PartialEq)]
9pub enum ModelError {
10    /// HTTP 429 / token-bucket exhausted.
11    RateLimit,
12    /// Prompt + completion would exceed the model's context window.
13    ContextTooLong,
14    /// The request itself was malformed (bad parameters, etc.).
15    InvalidRequest(String),
16    /// The provider returned a 5xx or equivalent error.
17    ServerError(String),
18    /// No response received within the allowed time window.
19    Timeout,
20    /// The model endpoint is temporarily unavailable.
21    Unavailable,
22}
23
24impl ModelError {
25    /// Returns `true` for errors that are worth retrying on the *same* model.
26    pub fn is_retryable(&self) -> bool {
27        matches!(self, Self::Timeout | Self::ServerError(_))
28    }
29
30    /// Returns `true` for errors that should trigger a move to the next model in the chain.
31    pub fn suggests_fallback(&self) -> bool {
32        matches!(self, Self::RateLimit | Self::ContextTooLong | Self::Unavailable)
33    }
34}
35
36/// A single model in a fallback chain.
37#[derive(Debug, Clone)]
38pub struct FallbackModel {
39    /// Provider-specific model identifier (e.g. `"gpt-4o"`, `"claude-3-5-sonnet-20241022"`).
40    pub model_id: String,
41    /// Lower priority value = preferred (0 is highest).
42    pub priority: u8,
43    /// Maximum context tokens supported by this model.
44    pub max_context_tokens: usize,
45    /// Estimated cost per 1 000 tokens (blended input+output).
46    pub cost_per_1k_tokens: f64,
47    /// Whether this model supports tool/function calling.
48    pub supports_tools: bool,
49    /// Whether this model supports streaming.
50    pub supports_streaming: bool,
51}
52
53/// Report about a single model in the chain.
54#[derive(Debug, Clone)]
55pub struct FallbackReport {
56    /// Model identifier.
57    pub model_id: String,
58    /// Chain priority.
59    pub priority: u8,
60    /// Consecutive failure count since last success.
61    pub failures: u32,
62    /// Whether the model is currently in its cooldown window.
63    pub in_cooldown: bool,
64    /// Seconds remaining on the cooldown, if applicable.
65    pub cooldown_remaining_secs: Option<u64>,
66}
67
68/// Ordered list of models with failure tracking and cooldown management.
69pub struct FallbackChain {
70    models: Vec<FallbackModel>,
71    current_idx: usize,
72    failure_counts: HashMap<String, u32>,
73    cooldown_until: HashMap<String, Instant>,
74}
75
76impl FallbackChain {
77    /// Build a new chain from the supplied models, sorted by ascending priority.
78    pub fn new(mut models: Vec<FallbackModel>) -> Self {
79        models.sort_by_key(|m| m.priority);
80        Self {
81            models,
82            current_idx: 0,
83            failure_counts: HashMap::new(),
84            cooldown_until: HashMap::new(),
85        }
86    }
87
88    /// Find the next available model that meets the given constraints.
89    ///
90    /// Skips models that:
91    /// - Are currently within their cooldown window
92    /// - Have accumulated more than 3 consecutive failures
93    /// - Do not support tools when `require_tools` is `true`
94    /// - Have a context window smaller than `min_context`
95    pub fn next_available(&self, require_tools: bool, min_context: usize) -> Option<&FallbackModel> {
96        let now = Instant::now();
97        for model in &self.models {
98            // Cooldown check
99            if let Some(&until) = self.cooldown_until.get(&model.model_id) {
100                if now < until {
101                    continue;
102                }
103            }
104            // Failure threshold
105            if self.failure_counts.get(&model.model_id).copied().unwrap_or(0) > 3 {
106                continue;
107            }
108            // Capability checks
109            if require_tools && !model.supports_tools {
110                continue;
111            }
112            if model.max_context_tokens < min_context {
113                continue;
114            }
115            return Some(model);
116        }
117        None
118    }
119
120    /// Record a failure for `model_id`.
121    ///
122    /// Applies a 60-second cooldown for [`ModelError::RateLimit`] errors.
123    pub fn record_failure(&mut self, model_id: &str, error: &ModelError) {
124        let count = self.failure_counts.entry(model_id.to_string()).or_insert(0);
125        *count += 1;
126
127        if matches!(error, ModelError::RateLimit) {
128            self.cooldown_until
129                .insert(model_id.to_string(), Instant::now() + Duration::from_secs(60));
130        }
131    }
132
133    /// Record a successful call for `model_id`, resetting its failure counter.
134    pub fn record_success(&mut self, model_id: &str) {
135        self.failure_counts.remove(model_id);
136        self.cooldown_until.remove(model_id);
137        // Advance current index to point at this model for future fast-path.
138        if let Some(idx) = self.models.iter().position(|m| m.model_id == model_id) {
139            self.current_idx = idx;
140        }
141    }
142
143    /// Manually clear the cooldown for `model_id`.
144    pub fn reset_cooldown(&mut self, model_id: &str) {
145        self.cooldown_until.remove(model_id);
146    }
147
148    /// Generate a status report for every model in the chain.
149    pub fn chain_report(&self) -> Vec<FallbackReport> {
150        let now = Instant::now();
151        self.models
152            .iter()
153            .map(|m| {
154                let failures = self.failure_counts.get(&m.model_id).copied().unwrap_or(0);
155                let cooldown_until = self.cooldown_until.get(&m.model_id).copied();
156                let in_cooldown = cooldown_until.map(|u| now < u).unwrap_or(false);
157                let cooldown_remaining_secs = cooldown_until.and_then(|u| {
158                    if now < u {
159                        Some(u.duration_since(now).as_secs())
160                    } else {
161                        None
162                    }
163                });
164                FallbackReport {
165                    model_id: m.model_id.clone(),
166                    priority: m.priority,
167                    failures,
168                    in_cooldown,
169                    cooldown_remaining_secs,
170                }
171            })
172            .collect()
173    }
174
175    /// Number of models in the chain.
176    pub fn len(&self) -> usize {
177        self.models.len()
178    }
179
180    /// Returns `true` if the chain contains no models.
181    pub fn is_empty(&self) -> bool {
182        self.models.is_empty()
183    }
184}
185
186/// Thread-safe wrapper around [`FallbackChain`] that executes closures against
187/// the best available model, retrying down the chain on failure.
188pub struct FallbackManager {
189    chain: Arc<Mutex<FallbackChain>>,
190}
191
192impl FallbackManager {
193    /// Create a new manager wrapping the given chain.
194    pub fn new(chain: FallbackChain) -> Self {
195        Self {
196            chain: Arc::new(Mutex::new(chain)),
197        }
198    }
199
200    /// Execute `f` against the best available model.
201    ///
202    /// Walks the chain until a model succeeds or all models are exhausted.
203    ///
204    /// - `require_tools`: skip models that do not support tool/function calling.
205    /// - `context_size`: skip models whose context window is too small.
206    pub fn execute<F, T>(
207        &self,
208        f: F,
209        require_tools: bool,
210        context_size: usize,
211    ) -> Result<T, String>
212    where
213        F: Fn(&str) -> Result<T, ModelError>,
214    {
215        // Collect candidate model IDs up-front to avoid holding the lock during `f`.
216        let candidates: Vec<String> = {
217            let chain = self.chain.lock().map_err(|e| format!("lock poisoned: {}", e))?;
218            chain
219                .models
220                .iter()
221                .filter(|m| {
222                    let now = Instant::now();
223                    let in_cooldown = chain
224                        .cooldown_until
225                        .get(&m.model_id)
226                        .map(|&u| now < u)
227                        .unwrap_or(false);
228                    let failures = chain.failure_counts.get(&m.model_id).copied().unwrap_or(0);
229                    !in_cooldown
230                        && failures <= 3
231                        && (!require_tools || m.supports_tools)
232                        && m.max_context_tokens >= context_size
233                })
234                .map(|m| m.model_id.clone())
235                .collect()
236        };
237
238        if candidates.is_empty() {
239            return Err("No available models in fallback chain".to_string());
240        }
241
242        for model_id in &candidates {
243            match f(model_id) {
244                Ok(result) => {
245                    if let Ok(mut chain) = self.chain.lock() {
246                        chain.record_success(model_id);
247                    }
248                    return Ok(result);
249                }
250                Err(err) => {
251                    if let Ok(mut chain) = self.chain.lock() {
252                        chain.record_failure(model_id, &err);
253                    }
254                    if !err.suggests_fallback() && !err.is_retryable() {
255                        return Err(format!("Non-recoverable error on {}: {:?}", model_id, err));
256                    }
257                    // Continue to next candidate
258                }
259            }
260        }
261
262        Err("All models in fallback chain failed".to_string())
263    }
264
265    /// Borrow a clone of the underlying chain handle for inspection.
266    pub fn chain(&self) -> Arc<Mutex<FallbackChain>> {
267        Arc::clone(&self.chain)
268    }
269}