tokio_prompt_orchestrator/
failover.rs1use std::sync::Arc;
29use std::time::Instant;
30
31use crate::provider_health::ProviderHealthMonitor;
32use crate::worker::ModelWorker;
33use crate::OrchestratorError;
34
35pub struct FailoverChain {
42 workers: Vec<(String, Arc<dyn ModelWorker>)>,
44 health: ProviderHealthMonitor,
45 max_attempts: usize,
47}
48
49impl FailoverChain {
50 pub fn new(health: ProviderHealthMonitor) -> Self {
56 Self {
57 workers: Vec::new(),
58 health,
59 max_attempts: 3,
60 }
61 }
62
63 pub fn add_worker(
71 mut self,
72 provider_id: impl Into<String>,
73 worker: Arc<dyn ModelWorker>,
74 ) -> Self {
75 self.workers.push((provider_id.into(), worker));
76 self
77 }
78
79 pub fn with_max_attempts(mut self, n: usize) -> Self {
86 self.max_attempts = n.max(1);
87 self
88 }
89
90 pub async fn infer(&self, prompt: &str) -> Result<(String, String), OrchestratorError> {
109 if self.workers.is_empty() {
110 return Err(OrchestratorError::Inference(
111 "failover chain is empty — no workers configured".to_string(),
112 ));
113 }
114
115 let mut attempts = 0usize;
116 let mut last_error: Option<OrchestratorError> = None;
117
118 for (provider_id, worker) in &self.workers {
119 if attempts >= self.max_attempts {
120 break;
121 }
122
123 if !self.health.is_usable(provider_id).await {
125 tracing::debug!(
126 provider = %provider_id,
127 "failover chain skipping unhealthy provider"
128 );
129 continue;
130 }
131
132 attempts += 1;
133 let start = Instant::now();
134
135 match worker.infer(prompt).await {
136 Ok(tokens) => {
137 let latency_ms = start.elapsed().as_millis() as u64;
138 self.health.record(provider_id, latency_ms, true).await;
139
140 let response = tokens.join("");
141 tracing::debug!(
142 provider = %provider_id,
143 attempt = attempts,
144 latency_ms = latency_ms,
145 "failover chain: inference succeeded"
146 );
147 return Ok((provider_id.clone(), response));
148 }
149 Err(err) => {
150 let latency_ms = start.elapsed().as_millis() as u64;
151 self.health.record(provider_id, latency_ms, false).await;
152
153 tracing::warn!(
154 provider = %provider_id,
155 attempt = attempts,
156 error = %err,
157 "failover chain: inference failed, trying next provider"
158 );
159 last_error = Some(err);
160 }
161 }
162 }
163
164 Err(last_error.unwrap_or_else(|| {
165 OrchestratorError::Inference(
166 "all providers in failover chain were unhealthy or exhausted".to_string(),
167 )
168 }))
169 }
170
171 pub fn len(&self) -> usize {
177 self.workers.len()
178 }
179
180 pub fn is_empty(&self) -> bool {
186 self.workers.is_empty()
187 }
188}
189
190#[cfg(test)]
195mod tests {
196 use super::*;
197 use crate::provider_health::ProviderHealthMonitor;
198 use crate::worker::EchoWorker;
199 use async_trait::async_trait;
200
201 struct FailingWorker {
203 message: String,
204 }
205
206 #[async_trait]
207 impl ModelWorker for FailingWorker {
208 async fn infer(&self, _prompt: &str) -> Result<Vec<String>, OrchestratorError> {
209 Err(OrchestratorError::Inference(self.message.clone()))
210 }
211 }
212
213 struct FixedWorker {
215 response: String,
216 }
217
218 #[async_trait]
219 impl ModelWorker for FixedWorker {
220 async fn infer(&self, _prompt: &str) -> Result<Vec<String>, OrchestratorError> {
221 Ok(vec![self.response.clone()])
222 }
223 }
224
225 fn health() -> ProviderHealthMonitor {
226 ProviderHealthMonitor::new(20)
227 }
228
229 #[tokio::test]
230 async fn test_empty_chain_returns_error() {
231 let chain = FailoverChain::new(health());
232 let result = chain.infer("test").await;
233 assert!(result.is_err());
234 }
235
236 #[tokio::test]
237 async fn test_single_echo_worker_succeeds() {
238 let chain = FailoverChain::new(health())
239 .add_worker("echo", Arc::new(EchoWorker::new()));
240 let (provider, response) = chain.infer("hello").await.unwrap();
241 assert_eq!(provider, "echo");
242 assert!(!response.is_empty());
243 }
244
245 #[tokio::test]
246 async fn test_primary_failure_falls_back_to_secondary() {
247 let chain = FailoverChain::new(health())
248 .add_worker("bad", Arc::new(FailingWorker { message: "oops".into() }))
249 .add_worker("good", Arc::new(FixedWorker { response: "ok".into() }));
250
251 let (provider, response) = chain.infer("prompt").await.unwrap();
252 assert_eq!(provider, "good");
253 assert_eq!(response, "ok");
254 }
255
256 #[tokio::test]
257 async fn test_all_failing_returns_last_error() {
258 let chain = FailoverChain::new(health())
259 .add_worker("w1", Arc::new(FailingWorker { message: "err1".into() }))
260 .add_worker("w2", Arc::new(FailingWorker { message: "err2".into() }))
261 .with_max_attempts(5);
262
263 let result = chain.infer("prompt").await;
264 assert!(result.is_err());
265 let err_str = result.unwrap_err().to_string();
267 assert!(err_str.contains("err2"), "expected last error message, got: {err_str}");
268 }
269
270 #[tokio::test]
271 async fn test_max_attempts_limits_tries() {
272 let chain = FailoverChain::new(health())
274 .add_worker("w1", Arc::new(FailingWorker { message: "err1".into() }))
275 .add_worker("w2", Arc::new(FailingWorker { message: "err2".into() }))
276 .add_worker("w3", Arc::new(FixedWorker { response: "ok".into() }))
277 .with_max_attempts(1);
278
279 let result = chain.infer("prompt").await;
281 assert!(result.is_err(), "should fail because w3 was never reached");
282 }
283
284 #[tokio::test]
285 async fn test_unhealthy_provider_is_skipped() {
286 let h = health();
287 for _ in 0..5 {
289 h.record("p1", 0, false).await;
290 }
291
292 let chain = FailoverChain::new(h)
293 .add_worker("p1", Arc::new(FixedWorker { response: "from-p1".into() }))
294 .add_worker("p2", Arc::new(FixedWorker { response: "from-p2".into() }));
295
296 let (provider, response) = chain.infer("prompt").await.unwrap();
297 assert_eq!(provider, "p2");
299 assert_eq!(response, "from-p2");
300 }
301
302 #[tokio::test]
303 async fn test_health_monitor_updated_on_success() {
304 let h = health();
305 let chain = FailoverChain::new(h.clone())
306 .add_worker("p1", Arc::new(FixedWorker { response: "hi".into() }));
307
308 chain.infer("prompt").await.unwrap();
309
310 let snap = h.get_health("p1").await.unwrap();
311 assert_eq!(snap.total_requests, 1);
312 assert_eq!(snap.total_errors, 0);
313 }
314
315 #[tokio::test]
316 async fn test_health_monitor_updated_on_failure() {
317 let h = health();
318 let chain = FailoverChain::new(h.clone())
319 .add_worker("p1", Arc::new(FailingWorker { message: "boom".into() }));
320
321 let _ = chain.infer("prompt").await;
322
323 let snap = h.get_health("p1").await.unwrap();
324 assert_eq!(snap.total_requests, 1);
325 assert_eq!(snap.total_errors, 1);
326 assert_eq!(snap.consecutive_failures, 1);
327 }
328
329 #[test]
330 fn test_len_and_is_empty() {
331 let chain = FailoverChain::new(health());
332 assert!(chain.is_empty());
333 assert_eq!(chain.len(), 0);
334
335 let chain = chain.add_worker("p1", Arc::new(EchoWorker::new()));
336 assert!(!chain.is_empty());
337 assert_eq!(chain.len(), 1);
338 }
339}