Skip to main content

telos_agent/model/provider/
routed.rs

1//! Hint-based model routing provider.
2//!
3//! [`RoutedModelConfig`] maps [`ModelHint`] values to concrete model names.
4//! [`RoutedProvider`] implements [`ModelProvider`] by resolving the hint on
5//! each request and delegating to a pre-created [`DeepSeekProvider`].
6
7use 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/// Maps [`ModelHint`] values to concrete model names.
21///
22/// Hints not present in the map fall back to `default_model`.
23#[derive(Clone)]
24pub struct RoutedModelConfig {
25    /// hint → model_name mapping
26    pub routes: HashMap<ModelHint, String>,
27    /// Model used when no hint matches or hint is `None`
28    pub default_model: String,
29    /// API key shared across all routed models
30    pub api_key: String,
31    /// Base URL shared across all routed models
32    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    /// Resolve a hint to a concrete model name.
48    /// Returns `default_model` when hint is `None` or not in the routes map.
49    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    /// Convenience constructor for the common two-model case.
54    ///
55    /// Routing strategy:
56    /// - Thinking → thinking model (pro): first-iteration understanding,
57    ///   stuck-detection re-planning, complex reasoning.
58    /// - Recovery → thinking model (pro): after a tool error, the model
59    ///   needs to diagnose what went wrong — real reasoning work.
60    /// - Execution → execution model (flash): the steady-state of most
61    ///   turns — processing tool results, issuing shell/edit commands.
62    ///   Flash is fast and cheap; if it makes a mistake, pro catches it
63    ///   on the next iteration via Recovery.
64    /// - Summarization → execution model (flash): compressing already-
65    ///   resolved history is purely mechanical.
66    /// - Default → flash: when no model_hint is set, use the fast path.
67    ///
68    /// If a single model is specified (`--model deepseek-v4-pro`) the
69    /// [`RoutedProvider`] is skipped entirely and all calls go to that model.
70    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    /// Set a custom base URL (e.g. for self-hosted or proxy endpoints).
85    pub fn with_base_url(mut self, url: String) -> Self {
86        self.base_url = url;
87        self
88    }
89
90    /// Collect all unique model names referenced in this config.
91    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
98/// A [`ModelProvider`] that routes requests to different models based on
99/// [`ModelHint`].
100///
101/// Providers are pre-created at construction time — one per unique model name
102/// in the config. Provider selection is a simple HashMap lookup with no
103/// allocation on the hot path.
104pub struct RoutedProvider {
105    config: RoutedModelConfig,
106    /// model_name → provider (pre-created)
107    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    /// Look up the provider for a given hint.
125    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        // Safety: new() pre-creates providers for every model in config
133        &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 borrows from self.providers (lifetime = 'a) ✅
161        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        // Just verify construction succeeds and providers map is populated
227        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        // estimate_tokens uses the default provider; actual value depends on
235        // tiktoken-rs but should be > 0 for non-empty text
236        let tokens = provider.estimate_tokens("hello world");
237        assert!(tokens > 0);
238    }
239}