telos_agent/model/provider/
routed.rs1use std::collections::{HashMap, HashSet};
8use std::pin::Pin;
9
10use async_trait::async_trait;
11use futures_core::stream::Stream;
12
13use crate::error::AgentError;
14use crate::model::provider::ModelProvider;
15use crate::model::provider::deepseek::{DeepSeekConfig, DeepSeekProvider};
16use crate::model::provider::types::{
17 CompletionRequest, CompletionResponse, ModelHint, ProviderEvent,
18};
19
20#[derive(Clone)]
24pub struct RoutedModelConfig {
25 pub routes: HashMap<ModelHint, String>,
27 pub default_model: String,
29 pub api_key: String,
31 pub base_url: String,
33}
34
35impl std::fmt::Debug for RoutedModelConfig {
36 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37 f.debug_struct("RoutedModelConfig")
38 .field("routes", &self.routes)
39 .field("default_model", &self.default_model)
40 .field("api_key", &"[REDACTED]")
41 .field("base_url", &self.base_url)
42 .finish()
43 }
44}
45
46impl RoutedModelConfig {
47 pub fn resolve(&self, hint: Option<ModelHint>) -> &str {
50 hint.and_then(|h| self.routes.get(&h).map(|s| s.as_str())).unwrap_or(&self.default_model)
51 }
52
53 pub fn dual(api_key: String, thinking: String, execution: String) -> Self {
71 let mut routes = HashMap::new();
72 routes.insert(ModelHint::Thinking, thinking.clone());
73 routes.insert(ModelHint::Recovery, thinking.clone());
74 routes.insert(ModelHint::Execution, execution.clone());
75 routes.insert(ModelHint::Summarization, execution.clone());
76 Self {
77 routes,
78 default_model: execution,
79 api_key,
80 base_url: "https://api.deepseek.com".into(),
81 }
82 }
83
84 pub fn with_base_url(mut self, url: String) -> Self {
86 self.base_url = url;
87 self
88 }
89
90 fn all_models(&self) -> HashSet<&str> {
92 let mut models: HashSet<&str> = self.routes.values().map(|s| s.as_str()).collect();
93 models.insert(&self.default_model);
94 models
95 }
96}
97
98pub struct RoutedProvider {
105 config: RoutedModelConfig,
106 providers: HashMap<String, DeepSeekProvider>,
108}
109
110impl RoutedProvider {
111 pub fn new(config: RoutedModelConfig) -> Self {
112 let mut providers = HashMap::new();
113 for model in config.all_models() {
114 let provider_config = DeepSeekConfig {
115 api_key: config.api_key.clone(),
116 model: model.to_string(),
117 base_url: config.base_url.clone(),
118 };
119 providers.insert(model.to_string(), DeepSeekProvider::new(provider_config));
120 }
121 Self { config, providers }
122 }
123
124 fn resolve(&self, hint: Option<ModelHint>) -> &DeepSeekProvider {
126 let model = self.config.resolve(hint);
127 tracing::debug!(
128 hint = ?hint,
129 model = %model,
130 "model route"
131 );
132 &self.providers[model]
134 }
135}
136
137#[async_trait]
138impl ModelProvider for RoutedProvider {
139 async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, AgentError> {
140 let hint = request.model_hint;
141 let provider = self.resolve(hint);
142 let result = provider.complete(request).await;
143 if let Ok(ref resp) = result {
144 tracing::debug!(
145 hint = ?hint,
146 input_tokens = resp.usage.map(|u| u.input_tokens).unwrap_or(0),
147 output_tokens = resp.usage.map(|u| u.output_tokens).unwrap_or(0),
148 "routed complete"
149 );
150 }
151 result
152 }
153
154 fn stream_complete<'a>(
155 &'a self,
156 request: CompletionRequest,
157 ) -> Pin<Box<dyn Stream<Item = Result<ProviderEvent, AgentError>> + Send + 'a>> {
158 let hint = request.model_hint;
159 let provider = self.resolve(hint);
160 provider.stream_complete(request)
162 }
163
164 fn estimate_tokens(&self, text: &str) -> usize {
165 self.resolve(None).estimate_tokens(text)
166 }
167}
168
169#[cfg(test)]
170mod tests {
171 use super::*;
172 use std::collections::HashMap;
173
174 fn test_config() -> RoutedModelConfig {
175 let mut routes = HashMap::new();
176 routes.insert(ModelHint::Thinking, "deepseek-v4-pro".into());
177 routes.insert(ModelHint::Execution, "deepseek-v4-flash".into());
178 routes.insert(ModelHint::Recovery, "deepseek-v4-pro".into());
179 routes.insert(ModelHint::Summarization, "deepseek-v4-flash".into());
180 RoutedModelConfig {
181 routes,
182 default_model: "deepseek-v4-flash".into(),
183 api_key: "test-key".into(),
184 base_url: "https://api.deepseek.com".into(),
185 }
186 }
187
188 #[test]
189 fn resolve_known_hint_returns_correct_model() {
190 let config = test_config();
191 assert_eq!(config.resolve(Some(ModelHint::Thinking)), "deepseek-v4-pro");
192 assert_eq!(config.resolve(Some(ModelHint::Execution)), "deepseek-v4-flash");
193 assert_eq!(config.resolve(Some(ModelHint::Recovery)), "deepseek-v4-pro");
194 }
195
196 #[test]
197 fn resolve_none_returns_default() {
198 let config = test_config();
199 assert_eq!(config.resolve(None), "deepseek-v4-flash");
200 }
201
202 #[test]
203 fn dual_constructor_maps_correctly() {
204 let config =
205 RoutedModelConfig::dual("key".into(), "pro-model".into(), "flash-model".into());
206 assert_eq!(config.resolve(Some(ModelHint::Thinking)), "pro-model");
207 assert_eq!(config.resolve(Some(ModelHint::Recovery)), "pro-model");
208 assert_eq!(config.resolve(Some(ModelHint::Execution)), "flash-model");
209 assert_eq!(config.resolve(Some(ModelHint::Summarization)), "flash-model");
210 assert_eq!(config.resolve(None), "flash-model");
211 }
212
213 #[test]
214 fn all_models_collects_unique_names() {
215 let config = test_config();
216 let models = config.all_models();
217 assert_eq!(models.len(), 2);
218 assert!(models.contains("deepseek-v4-pro"));
219 assert!(models.contains("deepseek-v4-flash"));
220 }
221
222 #[test]
223 fn routed_provider_constructs_without_error() {
224 let config = test_config();
225 let provider = RoutedProvider::new(config);
226 assert_eq!(provider.providers.len(), 2);
228 }
229
230 #[test]
231 fn estimate_tokens_delegates_to_default() {
232 let config = test_config();
233 let provider = RoutedProvider::new(config);
234 let tokens = provider.estimate_tokens("hello world");
237 assert!(tokens > 0);
238 }
239}