diff --git a/_docs/config/opencode-compatibility.mdx b/_docs/config/opencode-compatibility.mdx index bb1155b..2e85001 100644 --- a/_docs/config/opencode-compatibility.mdx +++ b/_docs/config/opencode-compatibility.mdx @@ -36,6 +36,7 @@ Blank cells mean that runtime behavior is not supported by that project today. ` | `agent..model` / `temperature` / `top_p` | ✅ | partial | Subagent `model` overrides are applied for Task/`@agent`; sampling settings are parsed but not yet applied. | | Markdown agent files | ✅ | ✅ | `.opencode/agents/*.md` frontmatter is parsed; body content becomes agent instructions. | | `provider..options.timeout` | ✅ | partial | Integer milliseconds or `false` to disable timeout. | +| `provider.` using `@ai-sdk/openai-compatible` | ✅ | ✅ | When `options.baseURL` is configured, crabcode requests its OpenAI-compatible `/v1/models` endpoint for `/models`. Discovered models are added alongside manually configured models; manually configured model metadata takes precedence; unsupported endpoints leave manual models unchanged. | | `theme` | ✅ | ✅ | In crabcode config files only. OpenCode config `theme` is ignored by crabcode. | | `notifications` | | ✅ | crabcode-specific sounds, desktop notifications, and terminal alert signals such as Zed tab dots. | | `images` | | ✅ | crabcode-specific image placeholder opener. | diff --git a/src/aisdk/providers/mod.rs b/src/aisdk/providers/mod.rs index 2eee0b6..b34be07 100644 --- a/src/aisdk/providers/mod.rs +++ b/src/aisdk/providers/mod.rs @@ -35,25 +35,13 @@ pub(crate) fn apply_extra_headers( /// version segment when one is already present (which produced /// `/v1/v1/responses`-style 404s). pub(crate) fn base_url_has_version_segment(base_url: &str) -> bool { - // Check if the URL path already contains a /vN segment (e.g., /v4, /v1) - if let Some(pos) = base_url.find("://") { - let after_scheme = &base_url[pos + 3..]; - if let Some(path_start) = after_scheme.find('/') { - let path = &after_scheme[path_start..]; - // Match /vN where N is one or more digits, followed by / or end of string - let bytes = path.as_bytes(); - for i in 0..bytes.len().saturating_sub(2) { - if bytes[i] == b'/' - && bytes[i + 1] == b'v' - && bytes[i + 2].is_ascii_digit() - && (i + 3 >= bytes.len() || bytes[i + 3] == b'/') - { - return true; - } - } - } - } - false + reqwest::Url::parse(base_url).ok().is_some_and(|url| { + url.path().split('/').any(|segment| { + segment.strip_prefix('v').is_some_and(|version| { + !version.is_empty() && version.bytes().all(|byte| byte.is_ascii_digit()) + }) + }) + }) } #[cfg(test)] @@ -85,6 +73,18 @@ mod tests { "https://opencode.ai/zen/go/v1/" )); assert!(base_url_has_version_segment("http://localhost:11434/v1")); + assert!(base_url_has_version_segment( + "https://gateway.example/api/v4" + )); + assert!(base_url_has_version_segment( + "https://gateway.example/v10/openai" + )); + assert!(!base_url_has_version_segment( + "https://gateway.example?v=/v1" + )); + assert!(!base_url_has_version_segment( + "https://gateway.example/v1beta" + )); assert!(!base_url_has_version_segment("https://api.openai.com")); assert!(!base_url_has_version_segment("https://api.anthropic.com")); } diff --git a/src/command/handlers.rs b/src/command/handlers.rs index eebb495..ec502de 100644 --- a/src/command/handlers.rs +++ b/src/command/handlers.rs @@ -68,7 +68,6 @@ pub fn handle_sessions<'a>( } else { session.title.clone() }; - crate::command::registry::DialogItem { id: session.id.clone(), name, @@ -399,6 +398,10 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult { }; if let Ok(discovery) = discovery.as_ref() { + crate::model::discovery::merge_dialog_models( + &mut models, + discovery.discover_custom_models_for_dialog().await, + ); discovery.apply_custom_models_to_dialog(&mut models); } @@ -880,7 +883,6 @@ pub async fn refresh_models() -> CommandResult { return CommandResult::Success(String::new()); } }; - let (providers_result, runtime_result) = tokio::join!( discovery.refresh_cache(), crate::model::extensions::ModelExtensions::refresh_runtime_models() diff --git a/src/model/catalog.rs b/src/model/catalog.rs index d1ad9cc..c8a7ba4 100644 --- a/src/model/catalog.rs +++ b/src/model/catalog.rs @@ -57,6 +57,10 @@ pub async fn selectable_models( } else { Vec::new() }; + merge_dialog_models( + &mut models, + discovery.discover_custom_models_for_dialog().await, + ); discovery.apply_custom_models_to_dialog(&mut models); let mut runtime_errors = Vec::new(); diff --git a/src/model/discovery.rs b/src/model/discovery.rs index 5abf36d..5dd2523 100644 --- a/src/model/discovery.rs +++ b/src/model/discovery.rs @@ -28,10 +28,30 @@ pub struct Provider { pub models: HashMap, } +#[derive(Deserialize)] +struct OpenAIModelsResponse { + #[serde(default)] + data: Vec, +} + +#[derive(Deserialize)] +struct OpenAIModel { + id: String, +} + static HTTP_CLIENT: OnceLock = OnceLock::new(); static MEMORY_CACHE: OnceLock>>> = OnceLock::new(); static MEMORY_MODEL_CACHE: OnceLock), CachedModels>>> = OnceLock::new(); +static MEMORY_CUSTOM_MODEL_CACHE: OnceLock< + Mutex), CachedCustomModels>>, +> = OnceLock::new(); + +#[derive(Clone)] +struct CachedCustomModels { + ids: HashMap>, + cached_at: std::time::Instant, +} #[derive(Clone)] struct CachedModels { @@ -60,6 +80,11 @@ fn memory_model_cache() -> &'static Mutex), Cached MEMORY_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) } +fn memory_custom_model_cache() -> &'static Mutex), CachedCustomModels>> +{ + MEMORY_CUSTOM_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Model { pub id: String, @@ -201,6 +226,83 @@ pub fn merge_dialog_models( } } +fn is_openai_compatible(provider: &crate::config::CustomProviderConfig) -> bool { + if provider + .base_url + .as_deref() + .is_some_and(|base_url| !base_url.trim().is_empty()) + { + return true; + } + matches!( + provider.npm.as_deref(), + Some("@ai-sdk/openai-compatible" | "@ai-sdk/gateway" | "@openrouter/ai-sdk-provider") + ) +} + +fn openai_models_endpoint(base_url: &str) -> Result { + let mut url = reqwest::Url::parse(base_url.trim()).context("invalid URL")?; + let path = url.path().trim_end_matches('/'); + let models_path = if crate::aisdk::providers::base_url_has_version_segment(url.as_str()) { + format!("{path}/models") + } else { + format!("{path}/v1/models") + }; + url.set_path(&models_path); + url.set_query(None); + url.set_fragment(None); + Ok(url) +} + +fn catalog_model_metadata<'a>( + catalog_providers: &'a HashMap, + provider_id: &str, + model_id: &str, +) -> Option<&'a Model> { + catalog_providers + .get(provider_id) + .and_then(|provider| provider.models.get(model_id)) + .or_else(|| { + catalog_providers + .values() + .find_map(|provider| provider.models.get(model_id)) + }) +} + +fn dialog_model( + provider: &Provider, + model_id: &str, + model: &Model, +) -> Option { + if matches!(model.status.as_deref(), Some("alpha" | "deprecated")) + || model.modalities.as_ref().is_some_and(|m| { + !m.output.iter().any(|output| output == "text") + || m.output.iter().any(|output| output == "image") + }) + { + return None; + } + Some(crate::model::types::Model { + id: model_id.to_string(), + name: model.name.clone(), + family: model.family.clone(), + provider_id: provider.id.clone(), + provider_name: provider.name.clone(), + attachment: model.attachment, + structured_output: model.structured_output, + free: crate::model::extensions::ModelExtensions::is_unauthenticated_free_provider( + &provider.id, + ) && model.cost.as_ref().is_some_and(|cost| cost.input == 0.0), + local: false, + reasoning_options: model.reasoning_options.clone(), + context_window: model + .limit + .as_ref() + .map(|limit| limit.context) + .filter(|context| *context > 0), + }) +} + impl Discovery { pub fn custom_provider_ids(&self) -> std::collections::HashSet { self.custom_providers @@ -246,6 +348,27 @@ impl Discovery { signature } + fn custom_provider_endpoint_signature(&self) -> Vec { + let Some(custom_providers) = &self.custom_providers else { + return Vec::new(); + }; + + let mut signature = custom_providers + .iter() + .map(|(provider_id, provider)| { + format!( + "{provider_id}:{}:{}", + provider.npm.as_deref().unwrap_or_default(), + provider.base_url.as_deref().unwrap_or_default() + ) + }) + .collect::>(); + signature.push(format!("disabled:{:?}", self.disabled_providers)); + signature.push(format!("enabled:{:?}", self.enabled_providers)); + signature.sort(); + signature + } + pub fn custom_provider_matches_filter(&self, filter: &str) -> bool { let filter = filter.trim().to_ascii_lowercase(); self.custom_providers.as_ref().is_some_and(|providers| { @@ -317,6 +440,314 @@ impl Discovery { .resolved_api_key() } + fn custom_provider_discovery_api_keys(&self) -> HashMap { + Self::discovery_api_keys_from_auth( + crate::persistence::AuthDAO::new() + .and_then(|auth| auth.load()) + .unwrap_or_default(), + ) + } + + fn discovery_api_keys_from_auth( + providers: HashMap, + ) -> HashMap { + providers + .into_iter() + .filter_map(|(provider_id, auth)| match auth { + crate::persistence::AuthConfig::Api { key } => Some((provider_id, key)), + crate::persistence::AuthConfig::OAuth { access, .. } => Some((provider_id, access)), + crate::persistence::AuthConfig::Local => None, + }) + .collect() + } + + /// Query configured OpenAI-compatible endpoints for their advertised model + /// IDs. Endpoint failures are deliberately isolated to preserve manual + /// configuration for providers that do not implement `GET /v1/models`. + pub async fn discover_custom_models_for_dialog(&self) -> Vec { + let providers = match self.fetch_providers().await { + Ok(providers) => providers, + Err(error) => { + crate::emit_log!("Skipped custom provider model discovery: {}", error); + return Vec::new(); + } + }; + let ids = self.cached_custom_model_ids(); + providers + .values() + .flat_map(|provider| { + ids.get(&provider.id) + .into_iter() + .flatten() + .filter_map(|id| { + provider + .models + .get(id) + .and_then(|model| dialog_model(provider, id, model)) + }) + }) + .collect() + } + + fn cached_custom_model_ids(&self) -> HashMap> { + memory_custom_model_cache() + .lock() + .ok() + .and_then(|cache| { + cache + .get(&( + self.get_cache_path().clone(), + self.custom_provider_endpoint_signature(), + )) + .cloned() + }) + .filter(|cached| cached.cached_at.elapsed().as_secs() <= CACHE_TTL_SECONDS) + .map(|cached| cached.ids) + .unwrap_or_default() + } + + async fn discover_custom_model_ids(&self) -> HashMap> { + let cache_key = ( + self.get_cache_path().clone(), + self.custom_provider_endpoint_signature(), + ); + if let Some(cached) = memory_custom_model_cache() + .lock() + .ok() + .and_then(|cache| cache.get(&cache_key).cloned()) + .filter(|cached| cached.cached_at.elapsed().as_secs() <= CACHE_TTL_SECONDS) + { + return cached.ids; + } + + let Some(custom_providers) = &self.custom_providers else { + return HashMap::new(); + }; + let stored_api_keys = self.custom_provider_discovery_api_keys(); + + let mut advertised = HashMap::new(); + for (provider_id, provider) in custom_providers { + if !self.provider_is_enabled(provider_id) || !is_openai_compatible(provider) { + continue; + } + + let Some(base_url) = provider.base_url.as_deref() else { + continue; + }; + let endpoint = match openai_models_endpoint(base_url) { + Ok(endpoint) => endpoint, + Err(error) => { + crate::emit_log!( + "Skipped {} model discovery: invalid base URL '{}': {}", + provider_id, + base_url, + error + ); + continue; + } + }; + + let mut request = self + .client + .get(endpoint) + .header("Accept", "application/json"); + if let Some(api_key) = stored_api_keys + .get(provider_id) + .cloned() + .or_else(|| provider.resolved_api_key()) + { + request = request.bearer_auth(api_key); + } + + let response = match request.send().await { + Ok(response) if response.status().is_success() => response, + Ok(response) => { + crate::emit_log!( + "Skipped {} model discovery: GET /v1/models returned {}", + provider_id, + response.status() + ); + continue; + } + Err(error) => { + crate::emit_log!( + "Skipped {} model discovery: GET /v1/models failed: {}", + provider_id, + error + ); + continue; + } + }; + + let response = match response.json::().await { + Ok(response) => response, + Err(error) => { + crate::emit_log!( + "Skipped {} model discovery: invalid GET /v1/models response: {}", + provider_id, + error + ); + continue; + } + }; + + let mut ids = response + .data + .into_iter() + .map(|model| model.id.trim().to_string()) + .filter(|id| !id.is_empty()) + .collect::>(); + ids.sort(); + ids.dedup(); + + advertised.insert(provider_id.clone(), ids); + } + + if let Ok(mut cache) = memory_custom_model_cache().lock() { + cache.insert( + cache_key, + CachedCustomModels { + ids: advertised.clone(), + cached_at: std::time::Instant::now(), + }, + ); + } + advertised + } + + // Resolve advertised IDs against the full catalog, not picker rows. Manual + // configuration is applied last so it remains authoritative. + fn apply_custom_catalog( + &self, + providers: &mut HashMap, + advertised: &HashMap>, + ) { + let mut additions = Vec::new(); + for (provider_id, ids) in advertised { + if !self.provider_is_enabled(provider_id) { + continue; + } + for model_id in ids { + let mut model = catalog_model_metadata(providers, provider_id, model_id) + .cloned() + .unwrap_or_else(|| Model { + id: model_id.clone(), + name: model_id.clone(), + family: String::new(), + attachment: false, + reasoning: false, + reasoning_options: Vec::new(), + tool_call: false, + structured_output: false, + temperature: false, + knowledge: String::new(), + release_date: String::new(), + last_updated: String::new(), + status: None, + modalities: Some(Modalities { + input: vec!["text".to_string()], + output: vec!["text".to_string()], + }), + open_weights: false, + cost: None, + limit: None, + provider: None, + }); + // Upstream routing belongs to the reference provider. Requests + // must use the configured gateway and its advertised model ID. + model.id = model_id.clone(); + model.provider = None; + additions.push((provider_id.clone(), model_id.clone(), model)); + } + } + for (provider_id, model_id, model) in additions { + let provider = providers + .entry(provider_id.clone()) + .or_insert_with(|| Provider { + id: provider_id.clone(), + name: provider_id, + api: String::new(), + doc: String::new(), + env: Vec::new(), + npm: String::new(), + models: HashMap::new(), + }); + provider.models.insert(model_id, model); + } + self.apply_custom_provider_overlays(providers); + } + + fn cached_model(&self, provider_id: &str, model_id: &str) -> Option<(Provider, Model)> { + let entry = self.load_cache_entry().ok().flatten(); + let provider = entry.as_ref().and_then(|entry| entry.data.get(provider_id)); + let advertised = self.cached_custom_model_ids(); + let is_discovered = advertised + .get(provider_id) + .is_some_and(|ids| ids.iter().any(|id| id == model_id)); + // Only this provider/model is materialized; pricing and limit lookups + // must not clone the entire catalog on each streamed response. + let mut providers = HashMap::new(); + if let Some(provider) = provider { + providers.insert( + provider_id.to_string(), + Provider { + id: provider.id.clone(), + name: provider.name.clone(), + api: provider.api.clone(), + doc: provider.doc.clone(), + env: provider.env.clone(), + npm: provider.npm.clone(), + models: provider + .models + .get(model_id) + .map(|model| HashMap::from([(model_id.to_string(), model.clone())])) + .unwrap_or_default(), + }, + ); + } + let mut ids = HashMap::new(); + if is_discovered { + if let Some(metadata) = entry + .as_ref() + .and_then(|entry| catalog_model_metadata(&entry.data, provider_id, model_id)) + { + let reference = + providers + .entry(provider_id.to_string()) + .or_insert_with(|| Provider { + id: provider_id.to_string(), + name: provider_id.to_string(), + api: String::new(), + doc: String::new(), + env: Vec::new(), + npm: String::new(), + models: HashMap::new(), + }); + reference + .models + .insert(model_id.to_string(), metadata.clone()); + } + ids.insert(provider_id.to_string(), vec![model_id.to_string()]); + } + self.apply_custom_catalog(&mut providers, &ids); + let mut provider = providers.remove(provider_id)?; + let model = provider.models.remove(model_id)?; + Some((provider, model)) + } + + pub fn clear_custom_model_discovery_cache(&self) { + let cache_key = ( + self.get_cache_path().clone(), + self.custom_provider_endpoint_signature(), + ); + if let Ok(mut cache) = memory_custom_model_cache().lock() { + cache.remove(&cache_key); + } + if let Ok(mut cache) = memory_model_cache().lock() { + cache.retain(|(path, _), _| path != self.get_cache_path()); + } + } + pub fn new() -> Result { let loaded = crate::config::ConfigLoader::load().ok(); let custom_providers = loaded @@ -565,17 +996,20 @@ impl Discovery { }; crate::model::extensions::ModelExtensions::augment_runtime_catalog(&mut providers); - self.apply_custom_provider_overlays(&mut providers); + let advertised = self.discover_custom_model_ids().await; + self.apply_custom_catalog(&mut providers, &advertised); Ok(providers) } pub async fn refresh_cache(&self) -> Result> { + self.clear_custom_model_discovery_cache(); let cached = self.load_from_cache().ok().flatten(); let mut providers = self.fetch_with_internal_providers(cached.as_ref()).await?; self.save_to_cache(&providers)?; crate::model::extensions::ModelExtensions::augment_runtime_catalog(&mut providers); - self.apply_custom_provider_overlays(&mut providers); + let advertised = self.discover_custom_model_ids().await; + self.apply_custom_catalog(&mut providers, &advertised); Ok(providers) } @@ -712,10 +1146,19 @@ impl Discovery { pub async fn fetch_models(&self) -> Result> { let mut models = crate::model::extensions::ModelExtensions::runtime_models_from_cache(); models.retain(|model| self.provider_is_enabled(&model.provider_id)); - let cache_key = ( - self.get_cache_path().clone(), - self.custom_provider_dialog_signature(), - ); + let cache_key = (self.get_cache_path().clone(), { + let mut signature = self.custom_provider_dialog_signature(); + signature.extend(self.custom_provider_endpoint_signature()); + if let Some(providers) = &self.custom_providers { + for (id, provider) in providers { + for (model_id, model) in &provider.models { + signature.push(format!("metadata:{id}:{model_id}:{model:?}")); + } + } + } + signature.sort(); + signature + }); if let Some(cached) = memory_model_cache() .lock() .ok() @@ -732,7 +1175,6 @@ impl Discovery { Err(_err) if !models.is_empty() => return Ok(models), Err(err) => return Err(err), }; - let mut persistent_models = Vec::new(); for (provider_id, provider) in providers { @@ -743,40 +1185,9 @@ impl Discovery { continue; } - let provider_name = provider.name.clone(); - for (model_id, model) in provider.models { - if matches!(model.status.as_deref(), Some("alpha" | "deprecated")) { - continue; - } - - let free = - crate::model::extensions::ModelExtensions::is_unauthenticated_free_provider( - &provider_id, - ) && model.cost.as_ref().is_some_and(|cost| cost.input == 0.0); - - let is_text_model = model.modalities.as_ref().map_or(true, |m| { - m.output.contains(&"text".to_string()) - && !m.output.contains(&"image".to_string()) - }); - - if is_text_model { - persistent_models.push(crate::model::types::Model { - id: model_id.clone(), - name: model.name.clone(), - family: model.family.clone(), - provider_id: provider_id.clone(), - provider_name: provider_name.clone(), - attachment: model.attachment, - structured_output: model.structured_output, - free, - local: false, - reasoning_options: model.reasoning_options.clone(), - context_window: model - .limit - .as_ref() - .map(|limit| limit.context) - .filter(|context| *context > 0), - }); + for (model_id, model) in &provider.models { + if let Some(model) = dialog_model(&provider, model_id, model) { + persistent_models.push(model); } } } @@ -796,9 +1207,7 @@ impl Discovery { } pub fn get_model_pricing(&self, provider_id: &str, model_id: &str) -> Option { - let entry = self.load_cache_entry().ok()??; - let provider = entry.data.get(provider_id)?; - let model = provider.models.get(model_id)?; + let (_, model) = self.cached_model(provider_id, model_id)?; model.cost.clone() } @@ -817,18 +1226,7 @@ impl Discovery { } pub fn get_model_limit(&self, provider_id: &str, model_id: &str) -> Option { - if let Some(limit) = self - .custom_providers - .as_ref() - .and_then(|providers| providers.get(&provider_id.trim().to_ascii_lowercase())) - .and_then(|provider| provider.models.get(model_id)) - .and_then(|model| model.context_window) - { - return Some(limit); - } - let entry = self.load_cache_entry().ok()??; - let provider = entry.data.get(provider_id)?; - let model = provider.models.get(model_id)?; + let (_, model) = self.cached_model(provider_id, model_id)?; model.limit.as_ref().map(|l| l.context) } @@ -842,9 +1240,7 @@ impl Discovery { { return Some(limit); } - let entry = self.load_cache_entry().ok()??; - let provider = entry.data.get(provider_id)?; - let model = provider.models.get(model_id)?; + let (_, model) = self.cached_model(provider_id, model_id)?; model.limit.as_ref().map(|limit| limit.output) } @@ -864,12 +1260,8 @@ impl Discovery { { return true; } - self.load_cache_entry() - .ok() - .flatten() - .and_then(|entry| entry.data.get(provider_id).cloned()) - .and_then(|provider| provider.models.get(model_id).cloned()) - .and_then(|model| model.modalities) + self.cached_model(provider_id, model_id) + .and_then(|(_, model)| model.modalities) .is_some_and(|modalities| modalities.input.iter().any(|input| input == modality)) } @@ -884,9 +1276,7 @@ impl Discovery { return Some(name); } - let entry = self.load_cache_entry().ok()??; - let provider = entry.data.get(provider_id)?; - let model = provider.models.get(model_id)?; + let (_, model) = self.cached_model(provider_id, model_id)?; Some(model.name.clone()) } @@ -912,16 +1302,7 @@ impl Discovery { provider_id: &str, model_id: &str, ) -> Option { - let mut providers = self - .load_cache_entry() - .ok() - .flatten() - .map(|entry| entry.data.clone()) - .unwrap_or_default(); - self.apply_custom_provider_overlays(&mut providers); - - let provider = providers.get(provider_id)?; - let model = provider.models.get(model_id)?; + let (provider, model) = self.cached_model(provider_id, model_id)?; let provider_npm = model .provider .as_ref() @@ -1009,6 +1390,9 @@ impl Discovery { if let Ok(mut cache) = memory_model_cache().lock() { cache.clear(); } + if let Ok(mut cache) = memory_custom_model_cache().lock() { + cache.clear(); + } Ok(()) } } @@ -1026,6 +1410,40 @@ mod tests { CustomModelConfig, CustomModelModalities, CustomProviderConfig, }; + #[test] + fn discovery_uses_stored_api_and_oauth_credentials() { + let oauth: crate::persistence::AuthConfig = serde_json::from_value(serde_json::json!({ + "type": "oauth", + "refresh": "refresh-token", + "access": "access-token", + "expires": 9223372036854775807_i64 + })) + .expect("OAuth auth config"); + let credentials = Discovery::discovery_api_keys_from_auth(HashMap::from([ + ( + "api-provider".to_string(), + crate::persistence::AuthConfig::Api { + key: "api-key".to_string(), + }, + ), + ("oauth-provider".to_string(), oauth), + ( + "local-provider".to_string(), + crate::persistence::AuthConfig::Local, + ), + ])); + + assert_eq!( + credentials.get("api-provider"), + Some(&"api-key".to_string()) + ); + assert_eq!( + credentials.get("oauth-provider"), + Some(&"access-token".to_string()) + ); + assert!(!credentials.contains_key("local-provider")); + } + fn unique_test_cache_path(name: &str) -> PathBuf { let nanos = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -1039,6 +1457,410 @@ mod tests { )) } + #[test] + fn openai_models_endpoint_appends_models_once() { + for (base, expected) in [ + ( + "https://gateway.example/api/v4/", + "https://gateway.example/api/v4/models", + ), + ( + "https://gateway.example/v10/openai", + "https://gateway.example/v10/openai/models", + ), + ( + "https://gateway.example/api?hint=/v1#ignored", + "https://gateway.example/api/v1/models", + ), + ] { + assert_eq!( + openai_models_endpoint(base).expect("endpoint").as_str(), + expected + ); + } + assert_eq!( + openai_models_endpoint("https://gateway.example/v1/") + .expect("endpoint") + .as_str(), + "https://gateway.example/v1/models" + ); + assert_eq!( + openai_models_endpoint("https://gateway.example/api") + .expect("endpoint") + .as_str(), + "https://gateway.example/api/v1/models" + ); + } + + #[test] + fn catalog_metadata_falls_back_to_matching_model_id() { + let model = Model { + id: "gpt-6-astra".to_string(), + name: "GPT-6 Astra".to_string(), + family: "gpt".to_string(), + attachment: true, + reasoning: true, + reasoning_options: Vec::new(), + tool_call: true, + structured_output: true, + temperature: true, + knowledge: String::new(), + release_date: String::new(), + last_updated: String::new(), + status: None, + modalities: None, + open_weights: false, + cost: None, + limit: None, + provider: None, + }; + let catalog = HashMap::from([( + "openai".to_string(), + Provider { + id: "openai".to_string(), + name: "OpenAI".to_string(), + api: String::new(), + doc: String::new(), + env: Vec::new(), + npm: String::new(), + models: HashMap::from([(model.id.clone(), model)]), + }, + )]); + + let metadata = + catalog_model_metadata(&catalog, "my-gateway", "gpt-6-astra").expect("metadata"); + assert_eq!(metadata.name, "GPT-6 Astra"); + assert!(metadata.attachment); + assert!(metadata.structured_output); + } + + #[tokio::test] + async fn custom_openai_compatible_discovery_uses_auth_and_catalog_metadata() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + let listener = TcpListener::bind("127.0.0.1:0").await.expect("listener"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("connection"); + let mut request = vec![0; 4096]; + let count = stream.read(&mut request).await.expect("request"); + let request = String::from_utf8_lossy(&request[..count]); + assert!(request.starts_with("GET /api/v1/models HTTP/1.1")); + assert!(request.contains("authorization: Bearer test-key")); + stream + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: 52\r\nconnection: close\r\n\r\n{\"data\":[{\"id\":\"gpt-6-astra\"},{\"id\":\"gpt-6-astra\"}]}" + ) + .await + .expect("response"); + }); + + let provider = CustomProviderConfig { + name: Some("Test Gateway".to_string()), + npm: Some("@ai-sdk/openai-compatible".to_string()), + base_url: Some(format!("http://{address}/api")), + api_key: Some("test-key".to_string()), + models: HashMap::new(), + }; + let mut discovery = + Discovery::new_with_custom(Some(HashMap::from([("gateway".to_string(), provider)]))) + .expect("discovery"); + let catalog = HashMap::from([( + "openai".to_string(), + Provider { + id: "openai".to_string(), + name: "OpenAI".to_string(), + api: String::new(), + doc: String::new(), + env: Vec::new(), + npm: String::new(), + models: HashMap::from([( + "gpt-6-astra".to_string(), + Model { + id: "gpt-6-astra".to_string(), + name: "GPT-6 Astra".to_string(), + family: "gpt".to_string(), + attachment: true, + reasoning: true, + reasoning_options: serde_json::from_value(serde_json::json!([ + {"id": "reasoning_effort", "type": "enum", "values": ["low", "high"]} + ])) + .expect("reasoning options"), + tool_call: true, + structured_output: true, + temperature: true, + knowledge: String::new(), + release_date: String::new(), + last_updated: String::new(), + status: None, + modalities: Some(Modalities { + input: vec!["text".into(), "image".into()], + output: vec!["text".into()], + }), + open_weights: false, + cost: Some(Cost { + input: 2.0, + output: 8.0, + cache_read: None, + cache_write: None, + }), + limit: Some(Limit { + context: 200_000, + output: 8192, + }), + provider: Some(ModelProvider { + npm: Some("@ai-sdk/anthropic".into()), + api: Some("https://reference.invalid".into()), + }), + }, + )]), + }, + )]); + + discovery.cache_path = unique_test_cache_path("custom_catalog"); + discovery.save_to_cache(&catalog).expect("cache"); + let models = discovery.discover_custom_models_for_dialog().await; + assert_eq!(models.len(), 1); + assert_eq!(models[0].provider_id, "gateway"); + assert_eq!(models[0].name, "GPT-6 Astra"); + assert!(models[0].attachment); + server.await.expect("server"); + + // A fresh request-time instance must see the same full model without + // another endpoint request (the mock server has already shut down). + let mut request_discovery = Discovery::new_with_custom(discovery.custom_providers.clone()) + .expect("request discovery"); + request_discovery.cache_path = discovery.cache_path.clone(); + let providers = request_discovery + .fetch_providers() + .await + .expect("request catalog"); + let provider = &providers["gateway"]; + let model = &provider.models["gpt-6-astra"]; + assert!(model.attachment && model.reasoning && model.tool_call && model.temperature); + assert_eq!(model.modalities.as_ref().unwrap().input, ["text", "image"]); + assert_eq!(model.limit.as_ref().unwrap().output, 8192); + assert!( + model.provider.is_none(), + "reference routing must not leak into gateway requests" + ); + assert_eq!(provider.api, format!("http://{address}/api")); + assert_eq!(provider.npm, "@ai-sdk/openai-compatible"); + assert_eq!( + request_discovery.get_model_limit("gateway", "gpt-6-astra"), + Some(200_000) + ); + assert_eq!( + request_discovery + .get_model_name("gateway", "gpt-6-astra") + .as_deref(), + Some("GPT-6 Astra") + ); + assert_eq!( + request_discovery + .get_model_pricing("gateway", "gpt-6-astra") + .unwrap() + .input, + 2.0 + ); + assert_eq!( + request_discovery.get_model_output_limit("gateway", "gpt-6-astra"), + Some(8192) + ); + assert!(request_discovery.model_supports_input_modality("gateway", "gpt-6-astra", "image")); + assert!(request_discovery + .get_model_reasoning_capability("gateway", "gpt-6-astra") + .is_some()); + let picker = request_discovery.fetch_models().await.expect("picker"); + assert!(picker + .iter() + .any(|model| model.provider_id == "gateway" && model.attachment)); + assert!( + !request_discovery + .load_from_cache() + .unwrap() + .unwrap() + .contains_key("gateway"), + "custom discovery must not pollute the persistent models.dev cache" + ); + + let custom = request_discovery + .custom_providers + .as_mut() + .unwrap() + .get_mut("gateway") + .unwrap(); + custom.models.insert( + "gpt-6-astra".into(), + CustomModelConfig { + name: Some("Manual Vision".into()), + context_window: Some(100_000), + max_tokens: Some(4096), + attachment: Some(false), + reasoning: Some(false), + reasoning_options: Some(Vec::new()), + temperature: Some(false), + tool_call: Some(false), + modalities: None, + launch: false, + }, + ); + let overridden = request_discovery + .fetch_providers() + .await + .expect("manual catalog"); + let model = &overridden["gateway"].models["gpt-6-astra"]; + assert_eq!(model.name, "Manual Vision"); + assert!(!model.attachment && !model.reasoning && !model.tool_call); + assert_eq!(model.limit.as_ref().unwrap().context, 100_000); + let picker = request_discovery + .fetch_models() + .await + .expect("manual picker"); + assert!(picker.iter().any(|model| model.provider_id == "gateway" + && model.name == "Manual Vision" + && !model.attachment)); + request_discovery.clear_custom_model_discovery_cache(); + assert!(request_discovery.cached_custom_model_ids().is_empty()); + assert!( + !request_discovery + .cached_model("gateway", "gpt-6-astra") + .unwrap() + .1 + .reasoning + ); + let _ = fs::remove_file(discovery.cache_path); + } + + #[tokio::test] + async fn custom_endpoint_discovery_is_additive_with_configured_models() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + let listener = TcpListener::bind("127.0.0.1:0").await.expect("listener"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("connection"); + let mut request = vec![0; 8192]; + let count = stream.read(&mut request).await.expect("request"); + let request = String::from_utf8_lossy(&request[..count]); + assert!(request.starts_with("GET /v1/models HTTP/1.1")); + stream + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: 52\r\nconnection: close\r\n\r\n{\"data\":[{\"id\":\"manual-model\"},{\"id\":\"live-model\"}]}" + ) + .await + .expect("response"); + }); + + let provider = CustomProviderConfig { + name: Some("Test Endpoint".to_string()), + npm: None, + base_url: Some(format!("http://{address}")), + api_key: None, + models: HashMap::from([( + "manual-model".to_string(), + CustomModelConfig { + name: Some("Manual Model".to_string()), + context_window: None, + max_tokens: None, + attachment: None, + reasoning: None, + reasoning_options: None, + temperature: None, + tool_call: None, + modalities: None, + launch: false, + }, + )]), + }; + let mut discovery = + Discovery::new_with_custom(Some(HashMap::from([("gateway".to_string(), provider)]))) + .expect("discovery"); + + discovery.cache_path = unique_test_cache_path("custom_additive"); + discovery.save_to_cache(&HashMap::new()).expect("cache"); + let discovered = discovery.discover_custom_models_for_dialog().await; + assert_eq!(discovered.len(), 2); + server.await.expect("server"); + + let mut models = discovered; + discovery.apply_custom_models_to_dialog(&mut models); + models.sort_by(|left, right| left.id.cmp(&right.id)); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "live-model"); + assert_eq!(models[1].id, "manual-model"); + assert_eq!(models[1].name, "Manual Model"); + } + + #[tokio::test] + async fn custom_endpoint_failure_preserves_manual_models_and_is_cached() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut request = [0; 4096]; + stream.read(&mut request).await.unwrap(); + stream + .write_all( + b"HTTP/1.1 404 Not Found\r\ncontent-length: 0\r\nconnection: close\r\n\r\n", + ) + .await + .unwrap(); + // Keep listening so an accidental second discovery request is + // observable, rather than just becoming another silent failure. + assert!( + tokio::time::timeout(Duration::from_millis(200), listener.accept()) + .await + .is_err() + ); + }); + let mut discovery = Discovery::new_with_custom(Some(HashMap::from([( + "manual-gateway".into(), + CustomProviderConfig { + name: None, + npm: Some("@ai-sdk/openai-compatible".into()), + base_url: Some(format!("http://{address}/v1")), + api_key: None, + models: HashMap::from([( + "manual".into(), + CustomModelConfig { + name: Some("Manual".into()), + context_window: Some(128_000), + max_tokens: Some(4096), + attachment: Some(true), + reasoning: None, + reasoning_options: None, + temperature: None, + tool_call: Some(true), + modalities: None, + launch: false, + }, + )]), + }, + )]))) + .unwrap(); + discovery.cache_path = unique_test_cache_path("unsupported_custom_endpoint"); + discovery.save_to_cache(&HashMap::new()).unwrap(); + let providers = discovery.fetch_providers().await.unwrap(); + assert!(providers["manual-gateway"].models["manual"].attachment); + assert_eq!( + discovery.get_model_limit("manual-gateway", "manual"), + Some(128_000) + ); + assert!(discovery + .discover_custom_models_for_dialog() + .await + .is_empty()); + let picker = discovery.fetch_models().await.unwrap(); + assert_eq!(picker.len(), 1); + assert_eq!(picker[0].name, "Manual"); + server.await.unwrap(); + let _ = fs::remove_file(discovery.cache_path); + } + #[test] fn estimate_tokens_uses_cache_rates_and_falls_back_to_input() { let priced = Cost { @@ -1621,7 +2443,8 @@ mod tests { #[tokio::test] async fn fetch_models_filters_deprecated_models() { - let mut discovery = Discovery::new().unwrap(); + let mut discovery = + Discovery::new_with_config(None, Default::default(), Default::default()).unwrap(); let cache_path = unique_test_cache_path("deprecated_model_filter"); if let Some(parent) = cache_path.parent() { fs::create_dir_all(parent).unwrap();