tokio_prompt_orchestrator/
model_fallback.rs1use std::collections::HashMap;
4use std::sync::{Arc, Mutex};
5use std::time::{Duration, Instant};
6
7#[derive(Debug, Clone, PartialEq)]
9pub enum ModelError {
10 RateLimit,
12 ContextTooLong,
14 InvalidRequest(String),
16 ServerError(String),
18 Timeout,
20 Unavailable,
22}
23
24impl ModelError {
25 pub fn is_retryable(&self) -> bool {
27 matches!(self, Self::Timeout | Self::ServerError(_))
28 }
29
30 pub fn suggests_fallback(&self) -> bool {
32 matches!(self, Self::RateLimit | Self::ContextTooLong | Self::Unavailable)
33 }
34}
35
36#[derive(Debug, Clone)]
38pub struct FallbackModel {
39 pub model_id: String,
41 pub priority: u8,
43 pub max_context_tokens: usize,
45 pub cost_per_1k_tokens: f64,
47 pub supports_tools: bool,
49 pub supports_streaming: bool,
51}
52
53#[derive(Debug, Clone)]
55pub struct FallbackReport {
56 pub model_id: String,
58 pub priority: u8,
60 pub failures: u32,
62 pub in_cooldown: bool,
64 pub cooldown_remaining_secs: Option<u64>,
66}
67
68pub 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 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 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 if let Some(&until) = self.cooldown_until.get(&model.model_id) {
100 if now < until {
101 continue;
102 }
103 }
104 if self.failure_counts.get(&model.model_id).copied().unwrap_or(0) > 3 {
106 continue;
107 }
108 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 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 pub fn record_success(&mut self, model_id: &str) {
135 self.failure_counts.remove(model_id);
136 self.cooldown_until.remove(model_id);
137 if let Some(idx) = self.models.iter().position(|m| m.model_id == model_id) {
139 self.current_idx = idx;
140 }
141 }
142
143 pub fn reset_cooldown(&mut self, model_id: &str) {
145 self.cooldown_until.remove(model_id);
146 }
147
148 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 pub fn len(&self) -> usize {
177 self.models.len()
178 }
179
180 pub fn is_empty(&self) -> bool {
182 self.models.is_empty()
183 }
184}
185
186pub struct FallbackManager {
189 chain: Arc<Mutex<FallbackChain>>,
190}
191
192impl FallbackManager {
193 pub fn new(chain: FallbackChain) -> Self {
195 Self {
196 chain: Arc::new(Mutex::new(chain)),
197 }
198 }
199
200 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 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 }
259 }
260 }
261
262 Err("All models in fallback chain failed".to_string())
263 }
264
265 pub fn chain(&self) -> Arc<Mutex<FallbackChain>> {
267 Arc::clone(&self.chain)
268 }
269}