Skip to main content

weft_core/defaults/
router.rs

1use crate::config::ProviderConfig;
2use crate::layers::RouterLayer;
3use crate::types::ChatRequest;
4use anyhow::{bail, Result};
5use async_trait::async_trait;
6
7/// Routes to the provider that lists the requested model.
8/// Falls back to default_provider if no match.
9pub struct DefaultRouter {
10    pub default_provider: String,
11}
12
13#[async_trait]
14impl RouterLayer for DefaultRouter {
15    async fn route(&self, request: &ChatRequest, providers: &[ProviderConfig]) -> Result<String> {
16        // If x_provider is set, use it directly
17        if let Some(ref p) = request.x_provider {
18            if providers.iter().any(|prov| prov.name == *p) {
19                return Ok(p.clone());
20            }
21            bail!("Requested provider '{}' not configured", p);
22        }
23
24        // Find provider that has the requested model
25        for prov in providers {
26            if prov.models.contains(&request.model) {
27                return Ok(prov.name.clone());
28            }
29        }
30
31        // Fallback to default
32        if providers.iter().any(|p| p.name == self.default_provider) {
33            return Ok(self.default_provider.clone());
34        }
35
36        // Last resort: first provider
37        providers
38            .first()
39            .map(|p| p.name.clone())
40            .ok_or_else(|| anyhow::anyhow!("No providers configured"))
41    }
42}
43
44#[cfg(test)]
45mod tests {
46    use super::*;
47    use crate::config::{ProviderApi, ProviderConfig};
48    use crate::types::{ChatMessage, ChatRequest};
49
50    fn make_request(model: &str) -> ChatRequest {
51        ChatRequest {
52            model: model.into(),
53            messages: vec![ChatMessage {
54                role: "user".into(),
55                content: "hi".into(),
56                tool_calls: None,
57                tool_call_id: None,
58            }],
59            stream: false,
60            temperature: None,
61            max_tokens: None,
62            top_p: None,
63            tools: None,
64            tool_choice: None,
65            response_format: None,
66            x_provider: None,
67        }
68    }
69
70    fn make_providers() -> Vec<ProviderConfig> {
71        vec![
72            ProviderConfig {
73                name: "openrouter".into(),
74                base_url: "https://openrouter.ai/api/v1".into(),
75                format: "openai".into(),
76                api: ProviderApi::ChatCompletions,
77                keys: vec![],
78                models: vec!["claude-sonnet-4".into(), "gpt-4o".into()],
79            },
80            ProviderConfig {
81                name: "anthropic".into(),
82                base_url: "https://api.anthropic.com/v1".into(),
83                format: "anthropic".into(),
84                api: ProviderApi::ChatCompletions,
85                keys: vec![],
86                models: vec!["claude-sonnet-4".into()],
87            },
88        ]
89    }
90
91    #[tokio::test]
92    async fn test_route_by_model() {
93        let router = DefaultRouter {
94            default_provider: "openrouter".into(),
95        };
96        let req = make_request("claude-sonnet-4");
97        let result = router.route(&req, &make_providers()).await.unwrap();
98        // First provider with the model wins
99        assert_eq!(result, "openrouter");
100    }
101
102    #[tokio::test]
103    async fn test_route_fallback_to_default() {
104        let router = DefaultRouter {
105            default_provider: "openrouter".into(),
106        };
107        let req = make_request("unknown-model");
108        let result = router.route(&req, &make_providers()).await.unwrap();
109        assert_eq!(result, "openrouter");
110    }
111
112    #[tokio::test]
113    async fn test_route_x_provider() {
114        let router = DefaultRouter {
115            default_provider: "openrouter".into(),
116        };
117        let mut req = make_request("claude-sonnet-4");
118        req.x_provider = Some("anthropic".into());
119        let result = router.route(&req, &make_providers()).await.unwrap();
120        assert_eq!(result, "anthropic");
121    }
122}