use std::sync::Arc;
use crate::config::GoldConfig;
use bamboo_domain::reasoning::ReasoningEffort;
use bamboo_llm::Config;
use bamboo_llm::{LLMError, ProviderModelRouter, ProviderRegistry, ResolvedModel};
pub const GOLD_CONFIG_METADATA_KEY: &str = "gold_config";
pub fn infer_provider(model_name: &str) -> Option<String> {
let name = model_name.trim().to_ascii_lowercase();
if name.is_empty() {
return None;
}
if name.starts_with("claude") {
return Some("anthropic".to_string());
}
if name.starts_with("gpt") || is_openai_o_series(&name) {
return Some("openai".to_string());
}
if name.starts_with("gemini") {
return Some("gemini".to_string());
}
None
}
fn is_openai_o_series(name: &str) -> bool {
let mut chars = name.chars();
matches!(chars.next(), Some('o')) && matches!(chars.next(), Some(c) if c.is_ascii_digit())
}
pub fn resolve_provider_type(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<String> {
let trimmed = provider_name.trim();
if trimmed.is_empty() {
return None;
}
config
.provider_instances
.get(trimmed)
.map(|instance| instance.provider_type.clone())
.or_else(|| {
provider_registry
.get_metadata(trimmed)
.map(|meta| meta.provider_type)
})
.or_else(|| Some(trimmed.to_string()))
}
pub fn resolve_provider_routing_key(
config: &Config,
requested_provider: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Result<String, LLMError> {
let requested = requested_provider.trim();
if requested.is_empty() {
return Err(LLMError::Auth("provider routing key is empty".to_string()));
}
let require_live = |routing_key: &str| {
provider_registry
.get(routing_key)
.map(|_| routing_key.to_string())
.ok_or_else(|| {
LLMError::Auth(format!(
"Provider '{routing_key}' is configured but unavailable"
))
})
};
if let Some(instance) = config.provider_instances.get(requested) {
if !instance.enabled {
return Err(LLMError::Auth(format!(
"Provider instance '{requested}' is disabled"
)));
}
return require_live(requested);
}
if provider_registry.get(requested).is_some() {
return Ok(requested.to_string());
}
if !bamboo_llm::AVAILABLE_PROVIDERS.contains(&requested) {
return Err(LLMError::Auth(format!(
"Unknown provider instance or type '{requested}'"
)));
}
if let Some(default_id) = config.default_provider_instance.as_deref() {
if let Some(instance) = config.provider_instances.get(default_id) {
if instance.provider_type == requested {
if !instance.enabled {
return Err(LLMError::Auth(format!(
"Default provider instance '{default_id}' is disabled"
)));
}
return require_live(default_id);
}
}
}
let mut matching_ids = config
.provider_instances
.iter()
.filter(|(_, instance)| instance.enabled && instance.provider_type == requested)
.map(|(id, _)| id.as_str())
.collect::<Vec<_>>();
matching_ids.sort_unstable();
if let Some(instance_id) = matching_ids.first() {
return require_live(instance_id);
}
Err(LLMError::Auth(format!(
"No enabled provider instance is available for type '{requested}'"
)))
}
pub fn parse_session_gold_config(session_gold_config_json: Option<&str>) -> Option<GoldConfig> {
let raw = session_gold_config_json?.trim();
if raw.is_empty() {
return None;
}
serde_json::from_str::<GoldConfig>(raw).ok()
}
pub fn normalize_gold_config_json(value: &serde_json::Value) -> Result<String, serde_json::Error> {
let parsed = serde_json::from_value::<GoldConfig>(value.clone())?;
serde_json::to_string(&parsed)
}
pub fn resolve_global_gold_config(config: &Config) -> Option<GoldConfig> {
config
.extra
.get("gold")
.cloned()
.and_then(|value| serde_json::from_value::<GoldConfig>(value).ok())
}
pub fn resolve_gold_config(
config: &Config,
session_gold_config_json: Option<&str>,
) -> Option<GoldConfig> {
if session_gold_config_json.is_some() {
return parse_session_gold_config(session_gold_config_json);
}
resolve_global_gold_config(config)
}
pub fn get_default_model_for_provider(
config: &Config,
provider_name: &str,
) -> Result<String, LLMError> {
let provider_name = provider_name.trim();
if let Some(instance) = config.provider_instances.get(provider_name) {
if !instance.enabled {
return Err(LLMError::Auth(format!(
"Provider instance '{provider_name}' is disabled"
)));
}
if let Some(model) = instance
.model
.as_deref()
.map(str::trim)
.filter(|model| !model.is_empty())
{
return Ok(model.to_string());
}
if instance.provider_type == "copilot" {
return Ok("gpt-4o".to_string());
}
return Err(LLMError::Auth(format!(
"Model must be specified for provider instance '{provider_name}'"
)));
}
match provider_name {
"copilot" => {
let provider_model = config
.providers()
.copilot
.as_ref()
.and_then(|c| c.model.clone());
Ok(provider_model.unwrap_or_else(|| "gpt-4o".to_string()))
}
"openai" => {
let openai_config = config
.providers()
.openai
.as_ref()
.ok_or_else(|| LLMError::Auth("OpenAI configuration required".to_string()))?;
openai_config.model.clone().ok_or_else(|| {
LLMError::Auth("OpenAI model must be specified in config".to_string())
})
}
"anthropic" => {
let anthropic_config =
config.providers().anthropic.as_ref().ok_or_else(|| {
LLMError::Auth("Anthropic configuration required".to_string())
})?;
anthropic_config.model.clone().ok_or_else(|| {
LLMError::Auth("Anthropic model must be specified in config".to_string())
})
}
"gemini" => {
let gemini_config = config
.providers()
.gemini
.as_ref()
.ok_or_else(|| LLMError::Auth("Gemini configuration required".to_string()))?;
gemini_config.model.clone().ok_or_else(|| {
LLMError::Auth("Gemini model must be specified in config".to_string())
})
}
other => Err(LLMError::Auth(format!("Unknown provider: {}", other))),
}
}
pub fn get_default_model_from_config(config: &Config) -> Result<String, LLMError> {
get_default_model_for_provider(config, config.effective_default_provider())
}
pub fn get_schedule_model_from_config(config: &Config) -> Result<String, LLMError> {
config
.get_fast_model()
.map(|model| model.trim().to_string())
.filter(|model| !model.is_empty())
.ok_or_else(|| {
LLMError::Auth(format!(
"No fast/default model configured for provider '{}'",
config.effective_default_provider()
))
})
}
pub fn get_fast_model_for_provider(config: &Config, provider_name: &str) -> Option<String> {
if let Some(instance) = config.provider_instances.get(provider_name.trim()) {
if !instance.enabled {
return None;
}
return instance
.fast_model
.as_deref()
.map(str::trim)
.filter(|model| !model.is_empty())
.map(ToString::to_string)
.or_else(|| get_default_model_for_provider(config, provider_name).ok());
}
let fast = match provider_name.trim() {
"openai" => config
.providers()
.openai
.as_ref()
.and_then(|c| c.fast_model.clone()),
"anthropic" => config
.providers()
.anthropic
.as_ref()
.and_then(|c| c.fast_model.clone()),
"gemini" => config
.providers()
.gemini
.as_ref()
.and_then(|c| c.fast_model.clone()),
"copilot" => config
.providers()
.copilot
.as_ref()
.and_then(|c| c.fast_model.clone()),
_ => None,
};
fast.or_else(|| get_default_model_for_provider(config, provider_name).ok())
}
pub fn get_memory_background_model_for_provider(
config: &Config,
provider_name: &str,
) -> Option<String> {
let configured = config
.memory()
.as_ref()
.and_then(|memory| memory.background_model.as_ref())
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.map(ToString::to_string);
configured.or_else(|| get_fast_model_for_provider(config, provider_name))
}
pub fn get_reasoning_effort_for_provider(
config: &Config,
provider_name: &str,
) -> Option<ReasoningEffort> {
config.reasoning_effort_for_key(provider_name)
}
pub fn get_task_summary_model_from_config(config: &Config) -> Result<String, LLMError> {
config.get_task_summary_model().ok_or_else(|| {
LLMError::Auth(format!(
"No task summary model configured for provider '{}'",
config.effective_default_provider()
))
})
}
pub fn get_memory_background_model_from_config(config: &Config) -> Result<String, LLMError> {
config.get_memory_background_model().ok_or_else(|| {
LLMError::Auth(format!(
"No background memory model configured for provider '{}'",
config.effective_default_provider()
))
})
}
pub fn get_vision_model_from_config(config: &Config) -> Result<String, LLMError> {
config.get_vision_model().ok_or_else(|| {
LLMError::Auth(format!(
"No model configured for provider '{}'",
config.effective_default_provider()
))
})
}
pub fn resolve_task_summary_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config
.defaults
.as_ref()
.and_then(|d| d.task_summary.as_ref())
{
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
resolve_background_model(config, provider_name, provider_registry)
.or_else(|| resolve_default_chat_model(config, provider_name, provider_registry))
}
pub fn resolve_background_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config
.defaults
.as_ref()
.and_then(|d| d.memory_background.as_ref())
.or_else(|| config.defaults.as_ref().and_then(|d| d.fast.as_ref()))
{
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
let model_name = get_memory_background_model_for_provider(config, provider_name)?;
let provider = provider_registry.get(provider_name)?;
Some(ResolvedModel {
provider,
model_name,
})
}
pub fn resolve_fast_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config.defaults.as_ref().and_then(|d| d.fast.as_ref()) {
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
let model_name = get_fast_model_for_provider(config, provider_name)?;
let provider = provider_registry.get(provider_name)?;
Some(ResolvedModel {
provider,
model_name,
})
}
pub fn resolve_vision_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config.defaults.as_ref().and_then(|d| d.vision.as_ref()) {
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
let model_name = if let Some(instance) = config.provider_instances.get(provider_name) {
if !instance.enabled {
return None;
}
instance
.vision_model
.as_deref()
.map(str::trim)
.filter(|model| !model.is_empty())
.map(ToString::to_string)
.or_else(|| get_default_model_for_provider(config, provider_name).ok())?
} else {
config.get_vision_model()?
};
let provider = provider_registry.get(provider_name)?;
Some(ResolvedModel {
provider,
model_name,
})
}
pub fn resolve_planning_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config.defaults.as_ref().and_then(|d| d.planning.as_ref()) {
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
resolve_default_chat_model(config, provider_name, provider_registry)
}
pub fn resolve_search_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config.defaults.as_ref().and_then(|d| d.search.as_ref()) {
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
resolve_fast_model(config, provider_name, provider_registry)
.or_else(|| resolve_default_chat_model(config, provider_name, provider_registry))
}
pub fn resolve_code_review_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config
.defaults
.as_ref()
.and_then(|d| d.code_review.as_ref())
{
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
resolve_default_chat_model(config, provider_name, provider_registry)
}
pub fn resolve_subagent_model_ref(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
subagent_type: &str,
) -> Option<bamboo_domain::ProviderModelRef> {
if config.features.provider_model_ref {
let router = ProviderModelRouter::new(provider_registry.clone());
if let Some(defaults) = config.defaults.as_ref() {
let candidate_refs = [
defaults.subagent_models.get(subagent_type),
defaults.sub_agent.as_ref(),
defaults.fast.as_ref(),
];
for model_ref in candidate_refs.into_iter().flatten() {
if router.route(model_ref).is_ok() {
return Some(model_ref.clone());
}
}
}
}
resolve_fast_model(config, provider_name, provider_registry)
.or_else(|| resolve_default_chat_model(config, provider_name, provider_registry))
.map(|resolved| bamboo_domain::ProviderModelRef::new(provider_name, resolved.model_name))
}
pub fn resolve_subagent_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
subagent_type: &str,
) -> Option<ResolvedModel> {
let model_ref =
resolve_subagent_model_ref(config, provider_name, provider_registry, subagent_type)?;
let provider = ProviderModelRouter::new(provider_registry.clone())
.route(&model_ref)
.or_else(|_| {
provider_registry.get(&model_ref.provider).ok_or_else(|| {
LLMError::Auth(format!("Provider '{}' not available", model_ref.provider))
})
})
.ok()?;
Some(ResolvedModel {
provider,
model_name: model_ref.model,
})
}
fn resolve_default_chat_model(
config: &Config,
provider_name: &str,
provider_registry: &Arc<ProviderRegistry>,
) -> Option<ResolvedModel> {
if config.features.provider_model_ref {
if let Some(model_ref) = config.defaults.as_ref().map(|d| &d.chat) {
if let Ok(provider) =
ProviderModelRouter::new(provider_registry.clone()).route(model_ref)
{
return Some(ResolvedModel {
provider,
model_name: model_ref.model.clone(),
});
}
}
}
let model_name = get_default_model_for_provider(config, provider_name).ok()?;
let provider = provider_registry.get(provider_name)?;
Some(ResolvedModel {
provider,
model_name,
})
}
pub fn resolve_image_fallback(
config_snapshot: &Config,
) -> Result<Option<crate::ImageFallbackConfig>, String> {
use crate::ImageFallbackMode;
if !config_snapshot.hooks.image_fallback.enabled {
return Ok(None);
}
let mode_str = config_snapshot
.hooks
.image_fallback
.mode
.trim()
.to_ascii_lowercase();
let mode = match mode_str.as_str() {
"placeholder" => ImageFallbackMode::Placeholder,
"error" => ImageFallbackMode::Error,
"ocr" => ImageFallbackMode::Ocr,
"vision" => ImageFallbackMode::Vision,
other => {
return Err(format!(
"Invalid config: hooks.image_fallback.mode must be 'placeholder', 'error', 'ocr', or 'vision' (got '{other}')"
));
}
};
let vision_model = if mode == ImageFallbackMode::Vision {
config_snapshot.get_vision_model()
} else {
None
};
Ok(Some(crate::ImageFallbackConfig { mode, vision_model }))
}
#[cfg(test)]
mod tests {
use super::*;
use bamboo_agent_core::tools::ToolSchema;
use bamboo_agent_core::Message;
use bamboo_config::CopilotConfig;
use bamboo_config::DefaultsConfig;
use bamboo_config::{OpenAIConfig, ProviderConfigs};
use bamboo_domain::ProviderModelRef;
use bamboo_llm::{LLMProvider, LLMStream};
use std::collections::HashMap;
macro_rules! test_config {
(@assign $config:ident, providers, $value:expr) => {
*$config.providers_mut() = $value;
};
(@assign $config:ident, memory, $value:expr) => {
*$config.memory_mut() = $value;
};
(@assign $config:ident, subagents, $value:expr) => {
*$config.subagents_mut() = $value;
};
(@assign $config:ident, $field:ident, $value:expr) => {
$config.$field = $value;
};
($($field:ident: $value:expr),* $(,)?) => {{
let mut config = Config::default();
$(test_config!(@assign config, $field, $value);)*
config
}};
}
struct NoopProvider;
#[async_trait::async_trait]
impl LLMProvider for NoopProvider {
async fn chat_stream(
&self,
_messages: &[Message],
_tools: &[ToolSchema],
_max_output_tokens: Option<u32>,
_model: &str,
) -> Result<LLMStream, LLMError> {
Err(LLMError::Api("noop".to_string()))
}
}
fn test_registry() -> Arc<ProviderRegistry> {
let mut providers: HashMap<String, Arc<dyn LLMProvider>> = HashMap::new();
providers.insert("openai".to_string(), Arc::new(NoopProvider));
Arc::new(ProviderRegistry::new(providers, "openai".to_string()))
}
fn registry_with_provider_ids(ids: &[&str], default: &str) -> Arc<ProviderRegistry> {
let providers = ids
.iter()
.map(|id| {
(
(*id).to_string(),
Arc::new(NoopProvider) as Arc<dyn LLMProvider>,
)
})
.collect();
Arc::new(ProviderRegistry::new(providers, default.to_string()))
}
#[test]
fn provider_type_alias_resolves_to_matching_default_instance() {
let mut config = Config::default();
config.provider_instances.insert(
"work-openai".to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "openai",
"enabled": true
}))
.unwrap(),
);
config.provider_instances.insert(
"main-anthropic".to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "anthropic",
"enabled": true
}))
.unwrap(),
);
config.default_provider_instance = Some("main-anthropic".to_string());
let registry =
registry_with_provider_ids(&["main-anthropic", "work-openai"], "main-anthropic");
assert_eq!(
resolve_provider_routing_key(&config, "openai", ®istry).unwrap(),
"work-openai"
);
assert_eq!(
resolve_provider_routing_key(&config, "anthropic", ®istry).unwrap(),
"main-anthropic"
);
}
#[test]
fn provider_type_alias_uses_lexical_instance_but_exact_unavailable_id_fails_closed() {
let mut config = Config::default();
for id in ["z-openai", "a-openai"] {
config.provider_instances.insert(
id.to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "openai",
"enabled": true
}))
.unwrap(),
);
}
config.provider_instances.insert(
"openai".to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "openai",
"enabled": false
}))
.unwrap(),
);
let registry = registry_with_provider_ids(&["a-openai", "z-openai"], "a-openai");
let error = resolve_provider_routing_key(&config, "openai", ®istry)
.unwrap_err()
.to_string();
assert!(error.contains("disabled"));
config.provider_instances.remove("openai");
assert_eq!(
resolve_provider_routing_key(&config, "openai", ®istry).unwrap(),
"a-openai"
);
assert!(resolve_provider_routing_key(&config, "typo", ®istry).is_err());
}
#[test]
fn test_get_model_from_openai_config() {
let config = test_config! {
provider: "openai".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: None,
vision_model: None,
reasoning_effort: None,
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_default_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-4o");
}
#[test]
fn instance_models_are_authoritative_over_stale_legacy_slot() {
let mut config = test_config! {
provider: "openai".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "legacy".to_string(),
model: Some("gpt-stale".to_string()),
fast_model: Some("gpt-stale-fast".to_string()),
vision_model: Some("gpt-stale-vision".to_string()),
..OpenAIConfig::default()
}),
..ProviderConfigs::default()
},
};
config.provider_instances.insert(
"work".to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "openai",
"model": "gpt-instance",
"fast_model": "gpt-instance-fast",
"vision_model": "gpt-instance-vision",
"enabled": true
}))
.unwrap(),
);
config.default_provider_instance = Some("work".to_string());
assert_eq!(
get_default_model_from_config(&config).unwrap(),
"gpt-instance"
);
assert_eq!(
get_fast_model_for_provider(&config, "work").as_deref(),
Some("gpt-instance-fast")
);
assert_eq!(
get_schedule_model_from_config(&config).unwrap(),
"gpt-instance-fast"
);
assert_eq!(
get_vision_model_from_config(&config).unwrap(),
"gpt-instance-vision"
);
}
#[test]
fn disabled_instance_model_does_not_fall_back_to_legacy_provider() {
let mut config = test_config! {
provider: "openai".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
model: Some("gpt-stale".to_string()),
..OpenAIConfig::default()
}),
..ProviderConfigs::default()
},
};
config.provider_instances.insert(
"work".to_string(),
serde_json::from_value(serde_json::json!({
"provider_type": "openai",
"model": "gpt-instance",
"enabled": false
}))
.unwrap(),
);
config.default_provider_instance = Some("work".to_string());
let error = get_default_model_from_config(&config)
.unwrap_err()
.to_string();
assert!(error.contains("disabled"));
assert!(get_fast_model_for_provider(&config, "work").is_none());
}
#[test]
fn test_error_when_model_not_configured() {
let config = test_config! {
provider: "openai".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: None, fast_model: None,
vision_model: None,
reasoning_effort: None,
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_default_model_from_config(&config);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("model must be specified"));
}
#[test]
fn test_get_model_from_copilot_provider_config() {
let config = test_config! {
provider: "copilot".to_string(),
providers: ProviderConfigs {
copilot: Some(CopilotConfig {
enabled: true,
headless_auth: false,
model: Some("gpt-4o-mini".to_string()),
fast_model: None,
vision_model: None,
reasoning_effort: None,
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_default_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-4o-mini");
}
#[test]
fn test_get_model_from_copilot_default_fallback() {
let config = test_config! {
provider: "copilot".to_string(),
providers: ProviderConfigs::default(),
};
let result = get_default_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-4o");
}
#[test]
fn test_get_default_model_for_specific_provider() {
let config = test_config! {
provider: "anthropic".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: Some("gpt-4o-mini".to_string()),
vision_model: None,
reasoning_effort: Some(ReasoningEffort::Medium),
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_default_model_for_provider(&config, "openai").expect("openai config");
assert_eq!(result, "gpt-4o");
}
#[test]
fn test_get_fast_model_for_specific_provider() {
let config = test_config! {
provider: "anthropic".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: Some("gpt-4o-mini".to_string()),
vision_model: None,
reasoning_effort: Some(ReasoningEffort::Medium),
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
assert_eq!(
get_fast_model_for_provider(&config, "openai").as_deref(),
Some("gpt-4o-mini")
);
}
#[test]
fn test_get_schedule_model_from_config_prefers_fast_model() {
let config = test_config! {
provider: "openai".to_string(),
defaults: None,
features: bamboo_config::FeatureFlags {
provider_model_ref: false,
..Default::default()
},
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: Some("gpt-4o-mini".to_string()),
vision_model: None,
reasoning_effort: None,
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_schedule_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-4o-mini");
}
#[test]
fn test_get_schedule_model_from_config_falls_back_to_default_model() {
let config = test_config! {
provider: "openai".to_string(),
defaults: None,
features: bamboo_config::FeatureFlags {
provider_model_ref: false,
..Default::default()
},
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: None,
vision_model: None,
reasoning_effort: None,
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
let result = get_schedule_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-4o");
}
#[test]
fn test_get_schedule_model_from_config_prefers_defaults_fast_over_chat() {
let config = test_config! {
provider: "openai".to_string(),
features: bamboo_config::FeatureFlags {
provider_model_ref: true,
..Default::default()
},
defaults: Some(DefaultsConfig {
chat: ProviderModelRef::new("openai", "gpt-chat"),
fast: Some(ProviderModelRef::new("openai", "gpt-fast")),
task_summary: None,
vision: None,
memory_background: None,
planning: None,
search: None,
code_review: None,
sub_agent: None,
subagent_models: HashMap::new(),
}),
};
let result = get_schedule_model_from_config(&config);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "gpt-fast");
}
#[test]
fn test_get_reasoning_effort_for_specific_provider() {
let config = test_config! {
provider: "anthropic".to_string(),
providers: ProviderConfigs {
openai: Some(OpenAIConfig {
api_key: "test".to_string(),
api_key_from_env: false,
api_key_encrypted: None,
credential_ref: None,
base_url: None,
model: Some("gpt-4o".to_string()),
fast_model: Some("gpt-4o-mini".to_string()),
vision_model: None,
reasoning_effort: Some(ReasoningEffort::Medium),
responses_only_models: vec![],
request_overrides: None,
extra: Default::default(),
}),
..ProviderConfigs::default()
},
};
assert_eq!(
get_reasoning_effort_for_provider(&config, "openai"),
Some(ReasoningEffort::Medium)
);
}
#[test]
fn resolve_subagent_model_ref_prefers_sub_agent_over_fast() {
let config = test_config! {
provider: "openai".to_string(),
features: bamboo_config::FeatureFlags {
provider_model_ref: true,
..Default::default()
},
defaults: Some(DefaultsConfig {
chat: ProviderModelRef::new("openai", "gpt-chat"),
fast: Some(ProviderModelRef::new("openai", "gpt-fast")),
task_summary: None,
vision: None,
memory_background: None,
planning: None,
search: None,
code_review: None,
sub_agent: Some(ProviderModelRef::new("openai", "gpt-sub-agent")),
subagent_models: HashMap::new(),
}),
};
let resolved = resolve_subagent_model_ref(&config, "openai", &test_registry(), "coder")
.expect("sub-agent model should resolve");
assert_eq!(resolved, ProviderModelRef::new("openai", "gpt-sub-agent"));
}
#[test]
fn resolve_subagent_model_ref_falls_back_to_fast_when_sub_agent_unset() {
let config = test_config! {
provider: "openai".to_string(),
features: bamboo_config::FeatureFlags {
provider_model_ref: true,
..Default::default()
},
defaults: Some(DefaultsConfig {
chat: ProviderModelRef::new("openai", "gpt-chat"),
fast: Some(ProviderModelRef::new("openai", "gpt-fast")),
task_summary: None,
vision: None,
memory_background: None,
planning: None,
search: None,
code_review: None,
sub_agent: None,
subagent_models: HashMap::new(),
}),
};
let resolved = resolve_subagent_model_ref(&config, "openai", &test_registry(), "coder")
.expect("fast model should resolve");
assert_eq!(resolved, ProviderModelRef::new("openai", "gpt-fast"));
}
#[test]
fn infer_provider_maps_known_families() {
assert_eq!(
infer_provider("claude-3-7-sonnet").as_deref(),
Some("anthropic")
);
assert_eq!(infer_provider("Claude-Opus").as_deref(), Some("anthropic"));
assert_eq!(infer_provider("gpt-4o").as_deref(), Some("openai"));
assert_eq!(infer_provider("GPT-4o-mini").as_deref(), Some("openai"));
assert_eq!(infer_provider("o1").as_deref(), Some("openai"));
assert_eq!(infer_provider("o3-mini").as_deref(), Some("openai"));
assert_eq!(infer_provider("o4-mini").as_deref(), Some("openai"));
assert_eq!(infer_provider("gemini-1.5-pro").as_deref(), Some("gemini"));
assert_eq!(infer_provider("llama-3"), None);
assert_eq!(infer_provider("opus"), None); assert_eq!(infer_provider(" "), None);
assert_eq!(infer_provider(""), None);
}
}