weft_core/defaults/
router.rs1use crate::config::ProviderConfig;
2use crate::layers::RouterLayer;
3use crate::types::ChatRequest;
4use anyhow::{bail, Result};
5use async_trait::async_trait;
6
7pub 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 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 for prov in providers {
26 if prov.models.contains(&request.model) {
27 return Ok(prov.name.clone());
28 }
29 }
30
31 if providers.iter().any(|p| p.name == self.default_provider) {
33 return Ok(self.default_provider.clone());
34 }
35
36 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 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}