Skip to main content

tokio_prompt_orchestrator/
failover.rs

1//! # Provider Failover Chain
2//!
3//! Wraps multiple [`ModelWorker`] implementations in a priority-ordered failover
4//! chain.  On failure the chain automatically retries with the next provider,
5//! skipping any provider that the health monitor currently marks as unusable.
6//!
7//! ## Example
8//!
9//! ```no_run
10//! use std::sync::Arc;
11//! use tokio_prompt_orchestrator::{EchoWorker, ModelWorker};
12//! use tokio_prompt_orchestrator::provider_health::ProviderHealthMonitor;
13//! use tokio_prompt_orchestrator::failover::FailoverChain;
14//!
15//! #[tokio::main]
16//! async fn main() {
17//!     let health = ProviderHealthMonitor::new(50);
18//!     let chain = FailoverChain::new(health)
19//!         .add_worker("primary", Arc::new(EchoWorker::new()))
20//!         .add_worker("backup", Arc::new(EchoWorker::new()))
21//!         .with_max_attempts(3);
22//!
23//!     let (provider_used, response) = chain.infer("Hello!").await.unwrap();
24//!     println!("Answered by {}: {}", provider_used, response);
25//! }
26//! ```
27
28use std::sync::Arc;
29use std::time::Instant;
30
31use crate::provider_health::ProviderHealthMonitor;
32use crate::worker::ModelWorker;
33use crate::OrchestratorError;
34
35/// A priority-ordered chain of model workers with automatic failover.
36///
37/// Workers are tried in insertion order.  Unhealthy workers (as determined by
38/// the [`ProviderHealthMonitor`]) are skipped without consuming an attempt
39/// slot.  The `max_attempts` limit counts only *actual* inference attempts
40/// (i.e. workers that were healthy and returned a result — success or failure).
41pub struct FailoverChain {
42    /// Ordered list of `(provider_id, worker)` pairs.
43    workers: Vec<(String, Arc<dyn ModelWorker>)>,
44    health: ProviderHealthMonitor,
45    /// Maximum number of inference attempts before giving up.
46    max_attempts: usize,
47}
48
49impl FailoverChain {
50    /// Create a new, empty failover chain.
51    ///
52    /// # Panics
53    ///
54    /// Never panics.
55    pub fn new(health: ProviderHealthMonitor) -> Self {
56        Self {
57            workers: Vec::new(),
58            health,
59            max_attempts: 3,
60        }
61    }
62
63    /// Append a worker to the end of the chain.
64    ///
65    /// Workers are tried in the order they are added.
66    ///
67    /// # Panics
68    ///
69    /// Never panics.
70    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    /// Override the maximum number of *attempted* inferences before the chain
80    /// gives up.  Defaults to `3`.  A value of `0` is silently raised to `1`.
81    ///
82    /// # Panics
83    ///
84    /// Never panics.
85    pub fn with_max_attempts(mut self, n: usize) -> Self {
86        self.max_attempts = n.max(1);
87        self
88    }
89
90    /// Run inference through the chain.
91    ///
92    /// Workers are tried in priority order.  For each candidate:
93    ///
94    /// 1. The health monitor is consulted; unhealthy providers are skipped
95    ///    (they do **not** count against `max_attempts`).
96    /// 2. The worker's [`ModelWorker::infer`] method is called.
97    /// 3. On success the outcome is recorded with the health monitor and the
98    ///    `(provider_id, response_text)` tuple is returned.
99    /// 4. On failure the outcome is recorded and the next candidate is tried
100    ///    (unless `max_attempts` has been exhausted).
101    ///
102    /// Returns `Err(OrchestratorError::Inference)` if all candidates fail or
103    /// the chain is empty.
104    ///
105    /// # Panics
106    ///
107    /// Never panics.
108    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            // Skip providers currently deemed unhealthy.
124            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    /// Return the number of workers currently in the chain.
172    ///
173    /// # Panics
174    ///
175    /// Never panics.
176    pub fn len(&self) -> usize {
177        self.workers.len()
178    }
179
180    /// Return `true` if the chain contains no workers.
181    ///
182    /// # Panics
183    ///
184    /// Never panics.
185    pub fn is_empty(&self) -> bool {
186        self.workers.is_empty()
187    }
188}
189
190// ---------------------------------------------------------------------------
191// Unit tests
192// ---------------------------------------------------------------------------
193
194#[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    // A worker that always returns Err.
202    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    // A worker that returns a fixed string.
214    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        // Last error should be from w2.
266        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        // Three failing workers but max_attempts = 1.
273        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        // Only w1 is attempted (1 attempt), w3 is never reached.
280        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        // Poison p1 with enough consecutive failures to mark it unreachable.
288        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        // p1 should be skipped (unhealthy), p2 should answer.
298        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}