specado_core/router/
primary_fallback.rs1use 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}