1use std::collections::BTreeMap;
13
14use crate::provider::{
15 CompletionRequest, CompletionResponse, ModelProvider, ModelSelection, ProviderError,
16};
17
18pub struct RoutingProvider {
20 routes: BTreeMap<String, Box<dyn ModelProvider>>,
21 fallback: Option<Box<dyn ModelProvider>>,
23 fallback_name: String,
24 last: String,
27}
28
29impl RoutingProvider {
30 pub fn new() -> RoutingProvider {
31 RoutingProvider {
32 routes: BTreeMap::new(),
33 fallback: None,
34 fallback_name: "none".to_string(),
35 last: "router".to_string(),
36 }
37 }
38
39 pub fn with(mut self, vendor: impl Into<String>, provider: Box<dyn ModelProvider>) -> Self {
41 self.routes.insert(vendor.into(), provider);
42 self
43 }
44
45 pub fn or_else(mut self, vendor: impl Into<String>, provider: Box<dyn ModelProvider>) -> Self {
51 self.fallback_name = vendor.into();
52 self.fallback = Some(provider);
53 self
54 }
55
56 pub fn vendors(&self) -> Vec<&str> {
62 let mut vendors: Vec<&str> = self.routes.keys().map(String::as_str).collect();
63 if self.fallback.is_some() && !vendors.contains(&self.fallback_name.as_str()) {
64 vendors.push(&self.fallback_name);
65 }
66 vendors.sort_unstable();
67 vendors
68 }
69
70 pub fn is_empty(&self) -> bool {
71 self.routes.is_empty() && self.fallback.is_none()
72 }
73
74 pub fn describe(&self) -> String {
76 match (self.routes.is_empty(), self.fallback.is_some()) {
77 (true, true) => format!("model calls go to {}", self.fallback_name),
78 (_, true) => format!(
79 "model calls route by vendor ({}), otherwise {}",
80 self.vendors().join(", "),
81 self.fallback_name
82 ),
83 (_, false) => format!(
84 "model calls route by vendor ({}); the artifact must pin one",
85 self.vendors().join(", ")
86 ),
87 }
88 }
89
90 fn vendor_of(selection: &ModelSelection) -> Option<&str> {
92 match selection {
93 ModelSelection::Exact(reference) => reference.split_once('/').map(|(vendor, _)| vendor),
94 _ => None,
95 }
96 }
97
98 fn pick(
99 &mut self,
100 selection: &ModelSelection,
101 ) -> Result<&mut dyn ModelProvider, ProviderError> {
102 if let Some(vendor) = RoutingProvider::vendor_of(selection) {
103 let vendor = vendor.to_string();
104 if self.routes.contains_key(&vendor) {
105 let provider = self.routes.get_mut(&vendor).expect("just checked");
106 return Ok(provider.as_mut());
107 }
108 if self.fallback.is_some() && self.fallback_name == vendor {
109 return Ok(self
110 .fallback
111 .as_mut()
112 .map(Box::as_mut)
113 .expect("checked immediately above"));
114 }
115 let reachable = match self.vendors() {
119 vendors if vendors.is_empty() => "none".to_string(),
120 vendors => vendors.join(", "),
121 };
122 return Err(ProviderError::Configuration(format!(
123 "the artifact pins `{vendor}/…`, and this run has no provider for `{vendor}`\n \
124 available: {reachable}\n \
125 export that vendor's key, or override the model with --model"
126 )));
127 }
128
129 if self.fallback.is_none() {
130 let reachable = if self.routes.is_empty() {
131 "nothing".to_string()
132 } else {
133 self.vendors().join(", ")
134 };
135 return Err(ProviderError::Configuration(format!(
136 "this artifact names no vendor and no default provider is configured\n \
137 pin one with `model exact \"<vendor>/<model>\"`; this run can reach: {reachable}"
138 )));
139 }
140 Ok(self
141 .fallback
142 .as_mut()
143 .map(Box::as_mut)
144 .expect("checked immediately above"))
145 }
146}
147
148impl Default for RoutingProvider {
149 fn default() -> Self {
150 RoutingProvider::new()
151 }
152}
153
154impl ModelProvider for RoutingProvider {
155 fn name(&self) -> &str {
156 &self.last
157 }
158
159 fn complete(
160 &mut self,
161 request: &CompletionRequest,
162 ) -> Result<CompletionResponse, ProviderError> {
163 let provider = self.pick(&request.model)?;
164 let name = provider.name().to_string();
165 let response = provider.complete(request);
166 self.last = name;
167 response
168 }
169
170 fn streams(&self) -> bool {
179 !self.is_empty()
180 && self.routes.values().all(|provider| provider.streams())
181 && self
182 .fallback
183 .as_ref()
184 .is_none_or(|provider| provider.streams())
185 }
186
187 fn complete_streaming(
188 &mut self,
189 request: &CompletionRequest,
190 on_delta: crate::provider::DeltaSink<'_>,
191 ) -> Result<CompletionResponse, ProviderError> {
192 let provider = self.pick(&request.model)?;
193 let name = provider.name().to_string();
194 let response = provider.complete_streaming(request, on_delta);
195 self.last = name;
196 response
197 }
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203 use crate::provider::Usage;
204 use serde_json::{json, Value};
205
206 struct Echo(&'static str);
208
209 impl ModelProvider for Echo {
210 fn name(&self) -> &str {
211 self.0
212 }
213 fn complete(
214 &mut self,
215 _request: &CompletionRequest,
216 ) -> Result<CompletionResponse, ProviderError> {
217 Ok(CompletionResponse {
218 value: Value::String(self.0.to_string()),
219 usage: Usage::default(),
220 model: self.0.to_string(),
221 })
222 }
223 }
224
225 fn request(model: ModelSelection) -> CompletionRequest {
226 CompletionRequest {
227 node: "n0".into(),
228 model,
229 system: None,
230 prompt: "hello".into(),
231 context: Vec::new(),
232 response_type: "markdown".into(),
233 shape: crate::schema::ResponseShape::Prose,
234 max_tokens: 100,
235 }
236 }
237
238 fn router() -> RoutingProvider {
239 RoutingProvider::new()
240 .with("openai", Box::new(Echo("openai")))
241 .with("anthropic", Box::new(Echo("anthropic")))
242 }
243
244 #[test]
245 fn a_pinned_vendor_decides_which_provider_answers() {
246 let mut router = router();
247 for (reference, expected) in [
248 ("openai/gpt-test", "openai"),
249 ("anthropic/claude-test", "anthropic"),
250 ] {
251 let response = router
252 .complete(&request(ModelSelection::Exact(reference.into())))
253 .unwrap();
254 assert_eq!(response.value, json!(expected));
255 assert_eq!(router.name(), expected, "the run log names who answered");
256 }
257 }
258
259 #[test]
260 fn a_vendor_nobody_serves_is_an_error_rather_than_a_silent_substitution() {
261 let error = router()
264 .complete(&request(ModelSelection::Exact("mistral/large".into())))
265 .unwrap_err();
266 let text = error.to_string();
267 assert!(text.contains("mistral"), "{text}");
268 assert!(text.contains("anthropic, openai"), "{text}");
269 }
270
271 #[test]
272 fn an_unpinned_call_goes_to_the_fallback() {
273 let mut router = router().or_else("anthropic", Box::new(Echo("anthropic")));
274 let response = router.complete(&request(ModelSelection::Default)).unwrap();
275 assert_eq!(response.value, json!("anthropic"));
276 }
277
278 #[test]
279 fn an_unpinned_call_with_no_fallback_says_how_to_pin_one() {
280 let error = router()
281 .complete(&request(ModelSelection::Default))
282 .unwrap_err();
283 let text = error.to_string();
284 assert!(text.contains("model exact"), "{text}");
285 assert!(text.contains("anthropic, openai"), "{text}");
286 }
287
288 #[test]
289 fn a_single_provider_serves_its_own_vendor_without_being_registered_twice() {
290 let mut router = RoutingProvider::new().or_else("anthropic", Box::new(Echo("anthropic")));
293 assert_eq!(router.vendors(), vec!["anthropic"]);
294
295 let pinned = router
296 .complete(&request(ModelSelection::Exact("anthropic/claude-x".into())))
297 .unwrap();
298 assert_eq!(pinned.value, json!("anthropic"));
299
300 let unpinned = router.complete(&request(ModelSelection::Default)).unwrap();
301 assert_eq!(unpinned.value, json!("anthropic"));
302
303 let error = router
305 .complete(&request(ModelSelection::Exact("openai/gpt-x".into())))
306 .unwrap_err();
307 assert!(
308 error.to_string().contains("no provider for `openai`"),
309 "{error}"
310 );
311 }
312
313 #[test]
314 fn several_providers_and_no_fallback_ask_the_artifact_to_pin() {
315 let described = router().describe();
316 assert!(described.contains("must pin one"), "{described}");
317 }
318
319 #[test]
320 fn a_bare_model_name_is_not_read_as_a_vendor() {
321 let mut router = router().or_else("anthropic", Box::new(Echo("anthropic")));
324 let response = router
325 .complete(&request(ModelSelection::Exact("claude-test".into())))
326 .unwrap();
327 assert_eq!(response.value, json!("anthropic"));
328 }
329
330 #[test]
331 fn the_description_says_what_this_run_can_reach() {
332 let described = router()
333 .or_else("anthropic", Box::new(Echo("anthropic")))
334 .describe();
335 assert!(described.contains("anthropic, openai"), "{described}");
336 }
337
338 #[test]
339 fn an_empty_router_is_empty() {
340 assert!(RoutingProvider::new().is_empty());
341 assert!(!router().is_empty());
342 }
343}