Skip to main content

weft_core/package/
fallback.rs

1use 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
11/// Tries the primary (JS) router; on failure, falls back to the default.
12pub 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
30/// Tries the primary (JS) key selector; on failure, falls back to the default.
31pub 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
57/// Tries the primary (JS) error handler; on failure, falls back to the default.
58pub 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        // ErrorHandlerLayer doesn't return Result, so we just use the primary.
67        // The JS bridge already has retry+fallback-to-Fail logic internally.
68        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}