oxicode_ai/router/
fallback.rs1#![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
15pub type BoxStream = Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>;
17
18#[derive(Debug, Clone, Default)]
20pub struct FallbackChain {
21 pub models: Vec<String>,
23}
24
25impl FallbackChain {
26 pub fn new(models: Vec<String>) -> Self {
28 Self { models }
29 }
30
31 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 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}