weft_core/package/
fallback.rs1use crate::config::ProviderConfig;
2use crate::layers::{
3 key_selector::ApiKeyState, ErrorAction, ErrorHandlerLayer, KeySelectorLayer, RequestError,
4 RouterLayer,
5};
6use crate::types::ChatRequest;
7use anyhow::Result;
8use async_trait::async_trait;
9use std::sync::Arc;
10
11pub struct FallbackRouter {
13 pub primary: Arc<dyn RouterLayer>,
14 pub fallback: Arc<dyn RouterLayer>,
15}
16
17#[async_trait]
18impl RouterLayer for FallbackRouter {
19 async fn route(&self, request: &ChatRequest, providers: &[ProviderConfig]) -> Result<String> {
20 match self.primary.route(request, providers).await {
21 Ok(result) => Ok(result),
22 Err(e) => {
23 tracing::warn!("Primary router failed, using fallback: {}", e);
24 self.fallback.route(request, providers).await
25 }
26 }
27 }
28}
29
30pub struct FallbackKeySelector {
32 pub primary: Arc<dyn KeySelectorLayer>,
33 pub fallback: Arc<dyn KeySelectorLayer>,
34}
35
36#[async_trait]
37impl KeySelectorLayer for FallbackKeySelector {
38 async fn select(&self, provider: &str, keys: &[ApiKeyState]) -> Result<usize> {
39 match self.primary.select(provider, keys).await {
40 Ok(idx) => Ok(idx),
41 Err(e) => {
42 tracing::warn!("Primary key selector failed, using fallback: {}", e);
43 self.fallback.select(provider, keys).await
44 }
45 }
46 }
47
48 fn mark_failed(&self, provider: &str, index: usize) {
49 self.fallback.mark_failed(provider, index);
50 }
51
52 fn mark_success(&self, provider: &str, index: usize) {
53 self.fallback.mark_success(provider, index);
54 }
55}
56
57pub struct FallbackErrorHandler {
59 pub primary: Arc<dyn ErrorHandlerLayer>,
60 pub fallback: Arc<dyn ErrorHandlerLayer>,
61}
62
63#[async_trait]
64impl ErrorHandlerLayer for FallbackErrorHandler {
65 async fn handle(&self, error: &RequestError) -> ErrorAction {
66 self.primary.handle(error).await
69 }
70}
71
72#[cfg(test)]
73mod tests {
74 use super::*;
75
76 use crate::types::{ChatMessage, ChatRequest};
77
78 fn test_request() -> ChatRequest {
79 ChatRequest {
80 model: "test".into(),
81 messages: vec![ChatMessage {
82 role: "user".into(),
83 content: "hi".into(),
84 tool_calls: None,
85 tool_call_id: None,
86 }],
87 stream: false,
88 temperature: None,
89 max_tokens: None,
90 top_p: None,
91 tools: None,
92 tool_choice: None,
93 response_format: None,
94 x_provider: None,
95 }
96 }
97
98 struct AlwaysFailRouter;
99
100 #[async_trait]
101 impl RouterLayer for AlwaysFailRouter {
102 async fn route(&self, _: &ChatRequest, _: &[ProviderConfig]) -> Result<String> {
103 anyhow::bail!("always fails")
104 }
105 }
106
107 struct FixedRouter(String);
108
109 #[async_trait]
110 impl RouterLayer for FixedRouter {
111 async fn route(&self, _: &ChatRequest, _: &[ProviderConfig]) -> Result<String> {
112 Ok(self.0.clone())
113 }
114 }
115
116 #[tokio::test]
117 async fn test_fallback_router_uses_primary() {
118 let router = FallbackRouter {
119 primary: Arc::new(FixedRouter("primary-provider".into())),
120 fallback: Arc::new(FixedRouter("fallback-provider".into())),
121 };
122 let req = test_request();
123 let result = router.route(&req, &[]).await.unwrap();
124 assert_eq!(result, "primary-provider");
125 }
126
127 #[tokio::test]
128 async fn test_fallback_router_falls_back() {
129 let router = FallbackRouter {
130 primary: Arc::new(AlwaysFailRouter),
131 fallback: Arc::new(FixedRouter("fallback-provider".into())),
132 };
133 let req = test_request();
134 let result = router.route(&req, &[]).await.unwrap();
135 assert_eq!(result, "fallback-provider");
136 }
137}