1use 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
19pub struct ModelRouter {
22 registry: Arc<ProviderRegistry>,
23 policy: RuntimePolicy,
24}
25
26impl ModelRouter {
27 #[must_use]
29 pub fn new(registry: Arc<ProviderRegistry>, policy: RuntimePolicy) -> Self {
30 Self { registry, policy }
31 }
32
33 #[must_use]
35 pub fn registry(&self) -> &ProviderRegistry {
36 &self.registry
37 }
38
39 #[must_use]
41 pub fn policy(&self) -> &RuntimePolicy {
42 &self.policy
43 }
44
45 #[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 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 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 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}