Skip to main content

oxicode_ai/router/
fallback.rs

1//! Fallback chain for sequential model failover.
2
3#![allow(missing_docs)]
4
5use crate::ProviderEvent;
6use crate::context::Context;
7use crate::error::ProviderError;
8use crate::providers::ProviderRegistry;
9use crate::providers::StreamOptions;
10use crate::types::Model;
11use futures::Stream;
12use std::pin::Pin;
13use std::sync::Arc;
14
15/// Stream type alias.
16pub type BoxStream = Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>;
17
18/// Ordered fallback chain that tries models sequentially.
19#[derive(Debug, Clone, Default)]
20pub struct FallbackChain {
21    /// Ordered list of `"provider/model-id"` fallback targets.
22    pub models: Vec<String>,
23}
24
25impl FallbackChain {
26    /// Create a new fallback chain.
27    pub fn new(models: Vec<String>) -> Self {
28        Self { models }
29    }
30
31    /// Try each model in sequence using a ProviderRegistry.
32    pub async fn try_models(
33        &self,
34        registry: &Arc<ProviderRegistry>,
35        context: &Context,
36        options: Option<StreamOptions>,
37    ) -> Result<BoxStream, ProviderError> {
38        let mut last_err = ProviderError::StreamError("no fallback models configured".to_string());
39
40        for model_str in &self.models {
41            let Some((provider_name, model_id)) = Self::parse_model(model_str) else {
42                continue;
43            };
44
45            let Some(provider) = registry.get(&provider_name) else {
46                last_err = ProviderError::UnknownProvider(provider_name.clone());
47                continue;
48            };
49
50            let model = Self::build_model(&provider_name, &model_id);
51            match provider.stream(&model, context, options.clone()).await {
52                Ok(stream) => {
53                    tracing::info!(model = model_str, "Fallback model succeeded");
54                    return Ok(stream);
55                }
56                Err(e) => {
57                    tracing::warn!(model = model_str, error = %e, "Fallback model failed");
58                    last_err = e;
59                }
60            }
61        }
62        Err(last_err)
63    }
64
65    /// Try each model using a generic resolver closure.
66    pub async fn try_models_with_resolver<F>(
67        &self,
68        resolver: F,
69        context: &Context,
70        options: Option<StreamOptions>,
71    ) -> Result<BoxStream, ProviderError>
72    where
73        F: Fn(&str) -> Option<Arc<dyn crate::providers::Provider>> + Sync,
74    {
75        let mut last_err = ProviderError::StreamError("no fallback models configured".to_string());
76
77        for model_str in &self.models {
78            let Some((provider_name, model_id)) = Self::parse_model(model_str) else {
79                continue;
80            };
81
82            let Some(provider) = resolver(&provider_name) else {
83                last_err = ProviderError::UnknownProvider(provider_name.clone());
84                continue;
85            };
86
87            let model = Self::build_model(&provider_name, &model_id);
88            match provider.stream(&model, context, options.clone()).await {
89                Ok(stream) => {
90                    tracing::info!(model = model_str, "Fallback model succeeded");
91                    return Ok(stream);
92                }
93                Err(e) => {
94                    tracing::warn!(model = model_str, error = %e, "Fallback model failed");
95                    last_err = e;
96                }
97            }
98        }
99        Err(last_err)
100    }
101
102    fn parse_model(s: &str) -> Option<(String, String)> {
103        let (provider, model_id) = s.split_once('/')?;
104        let provider = provider.trim().to_string();
105        let model_id = model_id.trim().to_string();
106        if provider.is_empty() || model_id.is_empty() {
107            return None;
108        }
109        Some((provider, model_id))
110    }
111
112    fn build_model(provider: &str, model_id: &str) -> Model {
113        Model::new(
114            model_id,
115            model_id,
116            crate::Api::AnthropicMessages,
117            provider,
118            "",
119        )
120    }
121}