Skip to main content

specado_core/router/
primary_fallback.rs

1use crate::error::{Error, Result};
2use crate::router::traits::Router;
3use crate::types::{PromptSpec, UniformResponse};
4use async_trait::async_trait;
5use std::future::Future;
6use std::pin::Pin;
7use std::sync::Arc;
8
9pub type BoxedExecutor = Arc<dyn Fn(PromptSpec, String) -> ExecutorFuture + Send + Sync>;
10
11pub type ExecutorFuture = Pin<Box<dyn Future<Output = Result<UniformResponse>> + Send>>;
12
13pub struct PrimaryFallbackRouter {
14    primary: String,
15    fallbacks: Vec<String>,
16    executor: BoxedExecutor,
17}
18
19impl PrimaryFallbackRouter {
20    pub fn new(
21        primary: impl Into<String>,
22        fallbacks: Vec<String>,
23        executor: BoxedExecutor,
24    ) -> Self {
25        Self {
26            primary: primary.into(),
27            fallbacks,
28            executor,
29        }
30    }
31}
32
33#[async_trait]
34impl Router for PrimaryFallbackRouter {
35    async fn route(&self, prompt: PromptSpec) -> Result<UniformResponse> {
36        let executor = &self.executor;
37        let mut last_error: Option<Error> = None;
38
39        for provider in std::iter::once(&self.primary).chain(self.fallbacks.iter()) {
40            let provider_path = provider.clone();
41            match executor(prompt.clone(), provider_path).await {
42                Ok(response) => return Ok(response),
43                Err(err) => {
44                    last_error = Some(err);
45                }
46            }
47        }
48
49        Err(last_error.unwrap_or_else(|| {
50            Error::Config("No providers configured for PrimaryFallbackRouter".into())
51        }))
52    }
53}
54
55#[cfg(test)]
56mod tests {
57    use super::*;
58    use crate::error::Error;
59    use crate::types::{Extensions, FinishReason, LossinessReport, StrictMode};
60
61    fn sample_response(model: &str) -> UniformResponse {
62        UniformResponse {
63            content: "hi".into(),
64            tool_calls: Vec::new(),
65            finish_reason: FinishReason::Stop,
66            model: model.into(),
67            provider_used: model.into(),
68            usage: None,
69            extensions: Extensions {
70                lossiness: LossinessReport::new(StrictMode::Warn),
71                provider_capabilities: None,
72            },
73        }
74    }
75
76    fn make_prompt() -> PromptSpec {
77        PromptSpec {
78            version: "1".into(),
79            messages: Vec::new(),
80            sampling: Default::default(),
81            response: Default::default(),
82            tools: Vec::new(),
83            tool_choice: None,
84            strict_mode: StrictMode::Warn,
85            metadata: Default::default(),
86        }
87    }
88
89    #[tokio::test]
90    async fn uses_primary_when_successful() {
91        let executor: BoxedExecutor = Arc::new(|prompt, provider| {
92            Box::pin(async move {
93                let _ = prompt;
94                Ok(sample_response(&provider))
95            })
96        });
97
98        let router = PrimaryFallbackRouter::new("primary", vec!["fallback".into()], executor);
99        let response = router.route(make_prompt()).await.unwrap();
100        assert_eq!(response.provider_used, "primary");
101    }
102
103    #[tokio::test]
104    async fn falls_back_when_primary_fails() {
105        let failures = Arc::new(std::sync::Mutex::new(0));
106        let executor: BoxedExecutor = {
107            let failures = failures.clone();
108            Arc::new(move |prompt, provider| {
109                let failures = failures.clone();
110                Box::pin(async move {
111                    let _ = prompt;
112                    if provider == "primary" {
113                        *failures.lock().unwrap() += 1;
114                        Err(Error::Provider {
115                            provider: provider.clone(),
116                            kind: crate::error::ProviderErrorKind::ServerError,
117                        })
118                    } else {
119                        Ok(sample_response(&provider))
120                    }
121                })
122            })
123        };
124
125        let router = PrimaryFallbackRouter::new("primary", vec!["secondary".into()], executor);
126        let response = router.route(make_prompt()).await.unwrap();
127        assert_eq!(response.provider_used, "secondary");
128        assert_eq!(*failures.lock().unwrap(), 1);
129    }
130}