|
4 | 4 | //! A model is named `<provider>/<model>` (`vercel/deepseek/deepseek-v4-flash-0731`, |
5 | 5 | //! `anthropic/claude-haiku-4-5`). [`EndpointResolver`] turns that name into |
6 | 6 | //! an [`Endpoint`]. [`BuiltinProviders`] knows the providers BenchFlow ships |
7 | | -//! with until the provider registry (`providers/`) replaces it behind the |
8 | | -//! same trait. |
| 7 | +//! with, and falls back to the provider registry (`providers/`) for the |
| 8 | +//! rest. |
9 | 9 |
|
10 | 10 | use std::fmt; |
11 | 11 |
|
@@ -97,14 +97,12 @@ impl EndpointError { |
97 | 97 | match self { |
98 | 98 | EndpointError::NoProvider { .. } | EndpointError::UnknownProvider { .. } => format!( |
99 | 99 | "name the model as <provider>/<model> with one of: {}", |
100 | | - BUILTIN |
101 | | - .iter() |
102 | | - .map(|p| p.name) |
103 | | - .collect::<Vec<_>>() |
104 | | - .join(", ") |
| 100 | + known_providers().join(", ") |
105 | 101 | ), |
| 102 | + // Not `--env`: those values reach the agent's sandbox, and a |
| 103 | + // native loop's key must stay on this machine. |
106 | 104 | EndpointError::NoKey { wanted, .. } => { |
107 | | - format!("export {wanted}=<key> (or pass --env {wanted}=<key>), then run again") |
| 105 | + format!("export {wanted}=<key> in this shell, then run again") |
108 | 106 | } |
109 | 107 | } |
110 | 108 | } |
@@ -186,10 +184,7 @@ impl EndpointResolver for BuiltinProviders { |
186 | 184 | }); |
187 | 185 | }; |
188 | 186 | let Some(spec) = BUILTIN.iter().find(|p| p.name == provider) else { |
189 | | - return Err(EndpointError::UnknownProvider { |
190 | | - provider: provider.to_string(), |
191 | | - model: model.to_string(), |
192 | | - }); |
| 187 | + return resolve_from_registry(provider, rest, model, host); |
193 | 188 | }; |
194 | 189 | let found = spec.keys.iter().find_map(|name| { |
195 | 190 | host.get(name) |
@@ -238,6 +233,65 @@ impl EndpointResolver for BuiltinProviders { |
238 | 233 | } |
239 | 234 | } |
240 | 235 |
|
| 236 | +/// A provider the native loops have no entry for, from the provider |
| 237 | +/// registry (`zai/`, `deepseek/`, `mimo/`): its OpenAI-compatible endpoint |
| 238 | +/// when it takes a bearer key, else its Anthropic-compatible one when it takes |
| 239 | +/// `x-api-key`, the way the native clients send keys. |
| 240 | +fn resolve_from_registry( |
| 241 | + provider: &str, |
| 242 | + rest: &str, |
| 243 | + model: &str, |
| 244 | + host: &HostEnv, |
| 245 | +) -> Result<Endpoint, EndpointError> { |
| 246 | + use crate::providers::{Auth, ProviderRegistry}; |
| 247 | + let unknown = || EndpointError::UnknownProvider { |
| 248 | + provider: provider.to_string(), |
| 249 | + model: model.to_string(), |
| 250 | + }; |
| 251 | + let registry = ProviderRegistry::builtin(); |
| 252 | + let p = registry.find(provider).ok_or_else(unknown)?; |
| 253 | + let (api, base_url) = if let Some(e) = p.openai.as_ref().filter(|e| e.auth == Auth::Bearer) { |
| 254 | + (Api::OpenAiChat, e.base_url.clone()) |
| 255 | + } else if let Some(e) = p.anthropic.as_ref().filter(|e| e.auth == Auth::XApiKey) { |
| 256 | + // The registry's Anthropic base URL stops before the version. |
| 257 | + (Api::AnthropicMessages, format!("{}/v1", e.base_url)) |
| 258 | + } else { |
| 259 | + return Err(unknown()); |
| 260 | + }; |
| 261 | + let key = crate::providers::find_provider_key(p, host) |
| 262 | + .ok() |
| 263 | + .filter(|k| !k.value().trim().is_empty()) |
| 264 | + .ok_or_else(|| EndpointError::NoKey { |
| 265 | + provider: provider.to_string(), |
| 266 | + wanted: p.key.env.join(" or "), |
| 267 | + })?; |
| 268 | + let base_url = host |
| 269 | + .get(BASE_URL_OVERRIDE_ENV) |
| 270 | + .map(|(v, _)| v.trim().to_string()) |
| 271 | + .filter(|v| !v.is_empty()) |
| 272 | + .unwrap_or(base_url); |
| 273 | + Ok(Endpoint { |
| 274 | + api, |
| 275 | + base_url: base_url.trim_end_matches('/').to_string(), |
| 276 | + key: ApiKey::new(key.value()), |
| 277 | + model: rest.to_string(), |
| 278 | + provider: provider.to_string(), |
| 279 | + key_label: key.source.clone(), |
| 280 | + upstreams: None, |
| 281 | + }) |
| 282 | +} |
| 283 | + |
| 284 | +/// Every provider a native loop can name: its own, then the registry's. |
| 285 | +fn known_providers() -> Vec<String> { |
| 286 | + let mut names: Vec<String> = BUILTIN.iter().map(|p| p.name.to_string()).collect(); |
| 287 | + for p in crate::providers::ProviderRegistry::builtin().providers() { |
| 288 | + if !names.contains(&p.name) { |
| 289 | + names.push(p.name.clone()); |
| 290 | + } |
| 291 | + } |
| 292 | + names |
| 293 | +} |
| 294 | + |
241 | 295 | #[cfg(test)] |
242 | 296 | mod tests { |
243 | 297 | use std::collections::BTreeMap; |
@@ -332,6 +386,44 @@ mod tests { |
332 | 386 | assert!(matches!(err, EndpointError::NoProvider { .. })); |
333 | 387 | } |
334 | 388 |
|
| 389 | + #[test] |
| 390 | + fn registry_providers_work_for_native_loops_too() { |
| 391 | + // MiMo's token plan: OpenAI-compatible with a bearer key. |
| 392 | + let e = BuiltinProviders |
| 393 | + .resolve( |
| 394 | + "mimo/mimo-v2.6-flash", |
| 395 | + &host(&[("MIMO_API_KEY", "k-123456789")]), |
| 396 | + ) |
| 397 | + .unwrap(); |
| 398 | + assert_eq!(e.api, Api::OpenAiChat); |
| 399 | + assert_eq!( |
| 400 | + e.url(), |
| 401 | + "https://token-plan-sgp.xiaomimimo.com/v1/chat/completions" |
| 402 | + ); |
| 403 | + assert_eq!(e.model, "mimo-v2.6-flash"); |
| 404 | + assert_eq!(e.key_label, "MIMO_API_KEY (environment)"); |
| 405 | + // Z.ai and DeepSeek, which were refused before. |
| 406 | + for (model, var) in [ |
| 407 | + ("zai/glm-5.3", "ZAI_API_KEY"), |
| 408 | + ("deepseek/deepseek-chat", "DEEPSEEK_API_KEY"), |
| 409 | + ] { |
| 410 | + let e = BuiltinProviders |
| 411 | + .resolve(model, &host(&[(var, "k-123456789")])) |
| 412 | + .unwrap(); |
| 413 | + assert_eq!(e.api, Api::OpenAiChat, "{model}"); |
| 414 | + } |
| 415 | + let err = BuiltinProviders |
| 416 | + .resolve("mimo/mimo-v2.6-flash", &host(&[])) |
| 417 | + .unwrap_err(); |
| 418 | + assert!( |
| 419 | + err.next_step().contains("MIMO_API_KEY"), |
| 420 | + "{}", |
| 421 | + err.next_step() |
| 422 | + ); |
| 423 | + let err = BuiltinProviders.resolve("nope/x", &host(&[])).unwrap_err(); |
| 424 | + assert!(err.next_step().contains("mimo"), "{}", err.next_step()); |
| 425 | + } |
| 426 | + |
335 | 427 | #[test] |
336 | 428 | fn the_override_points_any_provider_at_another_server() { |
337 | 429 | let e = BuiltinProviders |
|
0 commit comments