Skip to main content

behest_runtime/
router.rs

1//! Model router for runtime.
2//!
3//! Wraps `ProviderRegistry` with capability checking, retry logic,
4//! fallback strategies, and usage aggregation.
5
6use std::sync::Arc;
7use std::time::Duration;
8
9use tracing::{debug, warn};
10
11use behest_provider::{
12    ChatRequest, ChatResponse, EmbeddingRequest, EmbeddingResponse, ProviderCapabilities,
13    ProviderId, ProviderRegistry,
14};
15
16use super::error::{RuntimeError, RuntimeResult};
17use super::policy::RuntimePolicy;
18
19/// Routes model requests (chat and embedding) across providers with
20/// capability checking, exponential-backoff retry, and fallback chains.
21pub struct ModelRouter {
22    registry: Arc<ProviderRegistry>,
23    policy: RuntimePolicy,
24}
25
26impl ModelRouter {
27    /// Creates a new model router.
28    #[must_use]
29    pub fn new(registry: Arc<ProviderRegistry>, policy: RuntimePolicy) -> Self {
30        Self { registry, policy }
31    }
32
33    /// Returns the provider registry.
34    #[must_use]
35    pub fn registry(&self) -> &ProviderRegistry {
36        &self.registry
37    }
38
39    /// Returns the runtime policy.
40    #[must_use]
41    pub fn policy(&self) -> &RuntimePolicy {
42        &self.policy
43    }
44
45    /// Routes a chat request to a provider with capability checking and retry.
46    ///
47    /// # Errors
48    ///
49    /// Returns `RuntimeError` if provider not found, lacks capabilities, or all retries fail.
50    #[allow(clippy::too_many_lines)]
51    pub async fn route_chat(
52        &self,
53        provider_id: &ProviderId,
54        request: ChatRequest,
55        required_capabilities: Option<&ProviderCapabilities>,
56    ) -> RuntimeResult<ChatResponse> {
57        let provider = self
58            .registry
59            .chat(provider_id)
60            .ok_or_else(|| RuntimeError::ProviderNotFound(provider_id.to_string()))?;
61
62        if let Some(required) = required_capabilities {
63            let caps = provider.capabilities();
64            if !Self::supports_capabilities(&caps, required) {
65                return Err(RuntimeError::ProviderNotFound(format!(
66                    "provider {provider_id} lacks required capabilities",
67                )));
68            }
69        }
70
71        let mut last_error = None;
72        let max_attempts = if self.policy.retry_on_provider_error {
73            self.policy.max_retries + 1
74        } else {
75            1
76        };
77
78        for attempt in 1..=max_attempts {
79            match provider.complete(request.clone()).await {
80                Ok(response) => return Ok(response),
81                Err(e) => {
82                    if !e.is_retryable() || attempt == max_attempts {
83                        return Err(RuntimeError::from(e));
84                    }
85
86                    #[allow(clippy::cast_possible_truncation)]
87                    let backoff = Duration::from_millis(100 * 2u64.pow(attempt as u32 - 1));
88                    warn!(
89                        attempt,
90                        max_attempts,
91                        ?backoff,
92                        error = %e,
93                        "provider call failed, retrying"
94                    );
95                    tokio::time::sleep(backoff).await;
96                    last_error = Some(e);
97                }
98            }
99        }
100
101        Err(last_error
102            .unwrap_or_else(|| behest_core::error::ProviderError::Timeout {
103                provider: provider_id.clone(),
104            })
105            .into())
106    }
107
108    /// Routes a chat request across multiple providers with fallback ordering.
109    ///
110    /// Each provider is tried in order via [`Self::route_chat`] (which itself
111    /// applies per-provider retry logic). The first successful response is
112    /// returned. When all providers fail, the last error is propagated.
113    ///
114    /// # Errors
115    ///
116    /// Returns [`RuntimeError`] if every provider in the chain fails.
117    pub async fn route_chat_with_fallback(
118        &self,
119        provider_ids: &[ProviderId],
120        request: ChatRequest,
121        required_capabilities: Option<&ProviderCapabilities>,
122    ) -> RuntimeResult<ChatResponse> {
123        let mut last_error = None;
124
125        for provider_id in provider_ids {
126            match self
127                .route_chat(provider_id, request.clone(), required_capabilities)
128                .await
129            {
130                Ok(response) => return Ok(response),
131                Err(e) => {
132                    debug!(provider = %provider_id, error = %e, "provider failed, trying fallback");
133                    last_error = Some(e);
134                }
135            }
136        }
137
138        Err(last_error
139            .unwrap_or_else(|| RuntimeError::ProviderNotFound("no providers available".to_owned())))
140    }
141
142    /// Routes an embedding request with retry logic.
143    ///
144    /// # Errors
145    ///
146    /// Returns `RuntimeError` if provider not found or all retries fail.
147    pub async fn route_embedding(
148        &self,
149        provider_id: &ProviderId,
150        request: EmbeddingRequest,
151    ) -> RuntimeResult<EmbeddingResponse> {
152        let provider = self
153            .registry
154            .embedding(provider_id)
155            .ok_or_else(|| RuntimeError::ProviderNotFound(provider_id.to_string()))?;
156
157        let mut last_error = None;
158        let max_attempts = if self.policy.retry_on_provider_error {
159            self.policy.max_retries + 1
160        } else {
161            1
162        };
163
164        for attempt in 1..=max_attempts {
165            match provider.embed(request.clone()).await {
166                Ok(response) => return Ok(response),
167                Err(e) => {
168                    if !e.is_retryable() || attempt == max_attempts {
169                        return Err(RuntimeError::from(e));
170                    }
171
172                    #[allow(clippy::cast_possible_truncation)]
173                    let backoff = Duration::from_millis(100 * 2u64.pow(attempt as u32 - 1));
174                    warn!(
175                        attempt,
176                        max_attempts,
177                        ?backoff,
178                        error = %e,
179                        "embedding provider failed, retrying"
180                    );
181                    tokio::time::sleep(backoff).await;
182                    last_error = Some(e);
183                }
184            }
185        }
186
187        Err(last_error
188            .unwrap_or_else(|| behest_core::error::ProviderError::Timeout {
189                provider: provider_id.clone(),
190            })
191            .into())
192    }
193
194    /// Checks if provider capabilities support all required capabilities.
195    fn supports_capabilities(
196        available: &ProviderCapabilities,
197        required: &ProviderCapabilities,
198    ) -> bool {
199        (!required.chat || available.chat)
200            && (!required.chat_stream || available.chat_stream)
201            && (!required.tool_calling || available.tool_calling)
202            && (!required.parallel_tool_calls || available.parallel_tool_calls)
203            && (!required.json_schema_output || available.json_schema_output)
204            && (!required.vision || available.vision)
205            && (!required.embeddings || available.embeddings)
206    }
207}
208
209#[cfg(test)]
210#[allow(clippy::unwrap_used)]
211mod tests {
212    use super::*;
213    use async_trait::async_trait;
214    use behest_core::error::ProviderError;
215    use behest_provider::{ChatProvider, FinishReason, Message, ModelName, ProviderResult};
216    use std::sync::Arc;
217    use std::sync::atomic::{AtomicUsize, Ordering};
218
219    struct MockChatProvider {
220        id: ProviderId,
221        fail_count: Arc<AtomicUsize>,
222        caps: ProviderCapabilities,
223    }
224
225    impl MockChatProvider {
226        fn new(id: &str, fail_times: usize) -> Self {
227            Self {
228                id: ProviderId::new(id),
229                fail_count: Arc::new(AtomicUsize::new(fail_times)),
230                caps: ProviderCapabilities::chat(),
231            }
232        }
233
234        fn with_capabilities(id: &str, caps: ProviderCapabilities) -> Self {
235            Self {
236                id: ProviderId::new(id),
237                fail_count: Arc::new(AtomicUsize::new(0)),
238                caps,
239            }
240        }
241    }
242
243    #[async_trait]
244    impl ChatProvider for MockChatProvider {
245        fn id(&self) -> ProviderId {
246            self.id.clone()
247        }
248
249        fn capabilities(&self) -> ProviderCapabilities {
250            self.caps.clone()
251        }
252
253        async fn complete(&self, _request: ChatRequest) -> ProviderResult<ChatResponse> {
254            let remaining = self.fail_count.fetch_sub(1, Ordering::SeqCst);
255            if remaining > 0 {
256                return Err(ProviderError::Timeout {
257                    provider: self.id.clone(),
258                });
259            }
260
261            Ok(ChatResponse {
262                provider: self.id.clone(),
263                model: ModelName::new("test"),
264                message: Message::assistant_text("ok"),
265                finish_reason: FinishReason::Stop,
266                usage: None,
267                raw: None,
268            })
269        }
270    }
271
272    #[tokio::test]
273    async fn route_chat_should_succeed_on_first_try() {
274        let mut registry = ProviderRegistry::new();
275        registry.register_chat(MockChatProvider::new("test", 0));
276
277        let router = ModelRouter::new(Arc::new(registry), RuntimePolicy::new());
278        let request = ChatRequest::new(ModelName::new("test"));
279
280        let result = router
281            .route_chat(&ProviderId::new("test"), request, None)
282            .await;
283
284        assert!(result.is_ok());
285    }
286
287    #[tokio::test]
288    async fn route_chat_should_retry_on_retryable_error() {
289        let mut registry = ProviderRegistry::new();
290        registry.register_chat(MockChatProvider::new("test", 2));
291
292        let policy = RuntimePolicy::new().with_max_retries(3);
293        let router = ModelRouter::new(Arc::new(registry), policy);
294        let request = ChatRequest::new(ModelName::new("test"));
295
296        let result = router
297            .route_chat(&ProviderId::new("test"), request, None)
298            .await;
299
300        assert!(result.is_ok());
301    }
302
303    #[tokio::test]
304    async fn route_chat_should_fail_after_max_retries() {
305        let mut registry = ProviderRegistry::new();
306        registry.register_chat(MockChatProvider::new("test", 10));
307
308        let policy = RuntimePolicy::new().with_max_retries(2);
309        let router = ModelRouter::new(Arc::new(registry), policy);
310        let request = ChatRequest::new(ModelName::new("test"));
311
312        let result = router
313            .route_chat(&ProviderId::new("test"), request, None)
314            .await;
315
316        assert!(result.is_err());
317    }
318
319    #[tokio::test]
320    async fn route_chat_should_check_capabilities() {
321        let mut registry = ProviderRegistry::new();
322        registry.register_chat(MockChatProvider::with_capabilities(
323            "test",
324            ProviderCapabilities::chat(),
325        ));
326
327        let router = ModelRouter::new(Arc::new(registry), RuntimePolicy::new());
328        let request = ChatRequest::new(ModelName::new("test"));
329
330        let required = ProviderCapabilities {
331            chat_stream: true,
332            ..ProviderCapabilities::chat()
333        };
334
335        let result = router
336            .route_chat(&ProviderId::new("test"), request, Some(&required))
337            .await;
338
339        assert!(result.is_err());
340        assert!(matches!(
341            result.unwrap_err(),
342            RuntimeError::ProviderNotFound(_)
343        ));
344    }
345
346    #[tokio::test]
347    async fn route_chat_with_fallback_should_try_alternatives() {
348        let mut registry = ProviderRegistry::new();
349        registry.register_chat(MockChatProvider::new("primary", 10));
350        registry.register_chat(MockChatProvider::new("fallback", 0));
351
352        let policy = RuntimePolicy::new().with_max_retries(0);
353        let router = ModelRouter::new(Arc::new(registry), policy);
354        let request = ChatRequest::new(ModelName::new("test"));
355
356        let providers = vec![ProviderId::new("primary"), ProviderId::new("fallback")];
357        let result = router
358            .route_chat_with_fallback(&providers, request, None)
359            .await;
360
361        assert!(result.is_ok());
362    }
363
364    #[tokio::test]
365    async fn route_chat_should_return_error_for_unknown_provider() {
366        let registry = ProviderRegistry::new();
367        let router = ModelRouter::new(Arc::new(registry), RuntimePolicy::new());
368        let request = ChatRequest::new(ModelName::new("test"));
369
370        let result = router
371            .route_chat(&ProviderId::new("unknown"), request, None)
372            .await;
373
374        assert!(result.is_err());
375        assert!(matches!(
376            result.unwrap_err(),
377            RuntimeError::ProviderNotFound(_)
378        ));
379    }
380}