From 056d58c5f2a1618a2f9e006ca6ae6887515defbf Mon Sep 17 00:00:00 2001 From: Yanuar Date: Wed, 23 Sep 2026 15:54:36 +0700 Subject: [PATCH 1/4] feat(models): discover OpenAI-compatible provider models --- _docs/config/opencode-compatibility.mdx | 1 + src/command/handlers.rs | 6 +- src/model/catalog.rs | 4 + src/model/discovery.rs | 407 +++++++++++++++++++++++- 4 files changed, 415 insertions(+), 3 deletions(-) diff --git a/_docs/config/opencode-compatibility.mdx b/_docs/config/opencode-compatibility.mdx index c5aeff5..3a19b56 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`. 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/command/handlers.rs b/src/command/handlers.rs index be61868..d0c7de0 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); } @@ -875,6 +878,7 @@ pub async fn refresh_models() -> CommandResult { return CommandResult::Success(String::new()); } }; + discovery.clear_custom_model_discovery_cache(); let (providers_result, runtime_result) = tokio::join!( discovery.refresh_cache(), diff --git a/src/model/catalog.rs b/src/model/catalog.rs index d5e7016..0a051e2 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 3d5c97a..142db6d 100644 --- a/src/model/discovery.rs +++ b/src/model/discovery.rs @@ -28,10 +28,23 @@ 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), CachedModels>>> = + OnceLock::new(); #[derive(Clone)] struct CachedModels { @@ -60,6 +73,10 @@ fn memory_model_cache() -> &'static Mutex), Cached MEMORY_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) } +fn memory_custom_model_cache() -> &'static Mutex), CachedModels>> { + MEMORY_CUSTOM_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Model { pub id: String, @@ -201,6 +218,72 @@ pub fn merge_dialog_models( } } +fn is_openai_compatible(provider: &crate::config::CustomProviderConfig) -> bool { + 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 path.ends_with("/v1") { + 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 discovery_model_from_dialog_model(model: &crate::model::types::Model) -> Model { + Model { + id: model.id.clone(), + name: model.name.clone(), + family: model.family.clone(), + attachment: model.attachment, + reasoning: !model.reasoning_options.is_empty(), + reasoning_options: model.reasoning_options.clone(), + tool_call: false, + structured_output: model.structured_output, + temperature: false, + knowledge: String::new(), + release_date: String::new(), + last_updated: String::new(), + status: None, + modalities: Some(Modalities { + input: if model.attachment { + vec!["text".to_string(), "image".to_string()] + } else { + vec!["text".to_string()] + }, + output: vec!["text".to_string()], + }), + open_weights: false, + cost: None, + limit: None, + provider: None, + } +} + impl Discovery { pub fn custom_provider_ids(&self) -> std::collections::HashSet { self.custom_providers @@ -246,6 +329,25 @@ 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.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| { @@ -316,6 +418,158 @@ impl Discovery { .resolved_api_key() } + /// 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(); + } + }; + self.discover_custom_models_from_catalog(&providers).await + } + + async fn discover_custom_models_from_catalog( + &self, + catalog_providers: &HashMap, + ) -> Vec { + 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.models; + } + + let Some(custom_providers) = &self.custom_providers else { + return Vec::new(); + }; + + let mut models = Vec::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) = 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 provider_name = provider.name.as_deref().unwrap_or(provider_id); + let mut ids = response + .data + .into_iter() + .map(|model| model.id.trim().to_string()) + .filter(|id| !id.is_empty()) + .collect::>(); + ids.sort(); + ids.dedup(); + + for model_id in ids { + let metadata = catalog_model_metadata(catalog_providers, provider_id, &model_id); + models.push(crate::model::types::Model { + id: model_id.clone(), + name: metadata + .map(|model| model.name.clone()) + .unwrap_or_else(|| model_id.clone()), + family: metadata + .map(|model| model.family.clone()) + .unwrap_or_default(), + provider_id: provider_id.clone(), + provider_name: provider_name.to_string(), + attachment: metadata.is_some_and(|model| model.attachment), + structured_output: metadata.is_some_and(|model| model.structured_output), + free: false, + local: false, + reasoning_options: metadata + .map(|model| model.reasoning_options.clone()) + .unwrap_or_default(), + }); + } + } + + if let Ok(mut cache) = memory_custom_model_cache().lock() { + cache.insert( + cache_key, + CachedModels { + models: models.clone(), + cached_at: std::time::Instant::now(), + }, + ); + } + models + } + + 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); + } + } + pub fn new() -> Result { let loaded = crate::config::ConfigLoader::load().ok(); let custom_providers = loaded @@ -726,11 +980,21 @@ impl Discovery { return Ok(models); } - let providers = match self.fetch_providers().await { + let mut providers = match self.fetch_providers().await { Ok(providers) => providers, Err(_err) if !models.is_empty() => return Ok(models), Err(err) => return Err(err), }; + let discovered_models = self.discover_custom_models_from_catalog(&providers).await; + for discovered_model in discovered_models { + let Some(provider) = providers.get_mut(&discovered_model.provider_id) else { + continue; + }; + provider + .models + .entry(discovered_model.id.clone()) + .or_insert_with(|| discovery_model_from_dialog_model(&discovered_model)); + } let mut persistent_models = Vec::new(); @@ -953,6 +1217,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(()) } } @@ -983,6 +1250,141 @@ mod tests { )) } + #[test] + fn openai_models_endpoint_appends_models_once() { + 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 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: 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 models = discovery + .discover_custom_models_from_catalog(&catalog) + .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"); + } + #[test] fn estimate_tokens_uses_cache_rates_and_falls_back_to_input() { let priced = Cost { @@ -1530,7 +1932,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(); From 766e58a0ffb7b81beb59813dbf6c29e55c6a44ad Mon Sep 17 00:00:00 2001 From: Yanuar Date: Thu, 1 Oct 2026 11:12:29 +0700 Subject: [PATCH 2/4] feat(models): include discovered models alongside configured models --- _docs/config/opencode-compatibility.mdx | 2 +- src/model/discovery.rs | 68 +++++++++++++++++++++++++ 2 files changed, 69 insertions(+), 1 deletion(-) diff --git a/_docs/config/opencode-compatibility.mdx b/_docs/config/opencode-compatibility.mdx index 3a19b56..951fd6a 100644 --- a/_docs/config/opencode-compatibility.mdx +++ b/_docs/config/opencode-compatibility.mdx @@ -36,7 +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`. Manually configured model metadata takes precedence; unsupported endpoints leave manual models unchanged. | +| `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/model/discovery.rs b/src/model/discovery.rs index 142db6d..cbeaa16 100644 --- a/src/model/discovery.rs +++ b/src/model/discovery.rs @@ -219,6 +219,13 @@ 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") @@ -1385,6 +1392,67 @@ mod tests { server.await.expect("server"); } + #[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 discovery = + Discovery::new_with_custom(Some(HashMap::from([("gateway".to_string(), provider)]))) + .expect("discovery"); + + let discovered = discovery + .discover_custom_models_from_catalog(&HashMap::new()) + .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"); + } + #[test] fn estimate_tokens_uses_cache_rates_and_falls_back_to_input() { let priced = Cost { From 4c23a444161b63118a1432720f3569ff878e4a99 Mon Sep 17 00:00:00 2001 From: Yanuar Date: Thu, 1 Oct 2026 11:39:53 +0700 Subject: [PATCH 3/4] fix(models): authenticate compatible model discovery --- src/model/discovery.rs | 62 +++++++++++++++++++++++++++++++++++++++++- 1 file changed, 61 insertions(+), 1 deletion(-) diff --git a/src/model/discovery.rs b/src/model/discovery.rs index cbeaa16..bb3c0bf 100644 --- a/src/model/discovery.rs +++ b/src/model/discovery.rs @@ -425,6 +425,27 @@ 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`. @@ -459,6 +480,7 @@ impl Discovery { let Some(custom_providers) = &self.custom_providers else { return Vec::new(); }; + let stored_api_keys = self.custom_provider_discovery_api_keys(); let mut models = Vec::new(); for (provider_id, provider) in custom_providers { @@ -486,7 +508,11 @@ impl Discovery { .client .get(endpoint) .header("Accept", "application/json"); - if let Some(api_key) = provider.resolved_api_key() { + if let Some(api_key) = stored_api_keys + .get(provider_id) + .cloned() + .or_else(|| provider.resolved_api_key()) + { request = request.bearer_auth(api_key); } @@ -1244,6 +1270,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) From e53e9c61c2c8138658233ff582405cc49ba40c9b Mon Sep 17 00:00:00 2001 From: Blankeos Date: Mon, 5 Oct 2026 02:01:55 +0800 Subject: [PATCH 4/4] fix(models): share compatible discovery metadata with runtime --- src/aisdk/providers/mod.rs | 38 +-- src/command/handlers.rs | 2 - src/model/discovery.rs | 582 ++++++++++++++++++++++++++++--------- 3 files changed, 461 insertions(+), 161 deletions(-) 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 d0c7de0..952e443 100644 --- a/src/command/handlers.rs +++ b/src/command/handlers.rs @@ -878,8 +878,6 @@ pub async fn refresh_models() -> CommandResult { return CommandResult::Success(String::new()); } }; - discovery.clear_custom_model_discovery_cache(); - let (providers_result, runtime_result) = tokio::join!( discovery.refresh_cache(), crate::model::extensions::ModelExtensions::refresh_runtime_models() diff --git a/src/model/discovery.rs b/src/model/discovery.rs index bb3c0bf..c84e42c 100644 --- a/src/model/discovery.rs +++ b/src/model/discovery.rs @@ -43,8 +43,15 @@ 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), 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 { @@ -73,7 +80,8 @@ fn memory_model_cache() -> &'static Mutex), Cached MEMORY_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) } -fn memory_custom_model_cache() -> &'static Mutex), CachedModels>> { +fn memory_custom_model_cache() -> &'static Mutex), CachedCustomModels>> +{ MEMORY_CUSTOM_MODEL_CACHE.get_or_init(|| Mutex::new(HashMap::new())) } @@ -235,7 +243,7 @@ fn is_openai_compatible(provider: &crate::config::CustomProviderConfig) -> bool 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 path.ends_with("/v1") { + let models_path = if crate::aisdk::providers::base_url_has_version_segment(url.as_str()) { format!("{path}/models") } else { format!("{path}/v1/models") @@ -261,34 +269,33 @@ fn catalog_model_metadata<'a>( }) } -fn discovery_model_from_dialog_model(model: &crate::model::types::Model) -> Model { - Model { - id: model.id.clone(), +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, - reasoning: !model.reasoning_options.is_empty(), - reasoning_options: model.reasoning_options.clone(), - tool_call: false, structured_output: model.structured_output, - temperature: false, - knowledge: String::new(), - release_date: String::new(), - last_updated: String::new(), - status: None, - modalities: Some(Modalities { - input: if model.attachment { - vec!["text".to_string(), "image".to_string()] - } else { - vec!["text".to_string()] - }, - output: vec!["text".to_string()], - }), - open_weights: false, - cost: None, - limit: None, - provider: None, - } + 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(), + }) } impl Discovery { @@ -351,6 +358,8 @@ impl Discovery { ) }) .collect::>(); + signature.push(format!("disabled:{:?}", self.disabled_providers)); + signature.push(format!("enabled:{:?}", self.enabled_providers)); signature.sort(); signature } @@ -457,13 +466,41 @@ impl Discovery { return Vec::new(); } }; - self.discover_custom_models_from_catalog(&providers).await + 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() } - async fn discover_custom_models_from_catalog( - &self, - catalog_providers: &HashMap, - ) -> Vec { + 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(), @@ -474,15 +511,15 @@ impl Discovery { .and_then(|cache| cache.get(&cache_key).cloned()) .filter(|cached| cached.cached_at.elapsed().as_secs() <= CACHE_TTL_SECONDS) { - return cached.models; + return cached.ids; } let Some(custom_providers) = &self.custom_providers else { - return Vec::new(); + return HashMap::new(); }; let stored_api_keys = self.custom_provider_discovery_api_keys(); - let mut models = Vec::new(); + let mut advertised = HashMap::new(); for (provider_id, provider) in custom_providers { if !self.provider_is_enabled(provider_id) || !is_openai_compatible(provider) { continue; @@ -548,7 +585,6 @@ impl Discovery { } }; - let provider_name = provider.name.as_deref().unwrap_or(provider_id); let mut ids = response .data .into_iter() @@ -558,39 +594,139 @@ impl Discovery { ids.sort(); ids.dedup(); - for model_id in ids { - let metadata = catalog_model_metadata(catalog_providers, provider_id, &model_id); - models.push(crate::model::types::Model { - id: model_id.clone(), - name: metadata - .map(|model| model.name.clone()) - .unwrap_or_else(|| model_id.clone()), - family: metadata - .map(|model| model.family.clone()) - .unwrap_or_default(), - provider_id: provider_id.clone(), - provider_name: provider_name.to_string(), - attachment: metadata.is_some_and(|model| model.attachment), - structured_output: metadata.is_some_and(|model| model.structured_output), - free: false, - local: false, - reasoning_options: metadata - .map(|model| model.reasoning_options.clone()) - .unwrap_or_default(), - }); - } + advertised.insert(provider_id.clone(), ids); } if let Ok(mut cache) = memory_custom_model_cache().lock() { cache.insert( cache_key, - CachedModels { - models: models.clone(), + CachedCustomModels { + ids: advertised.clone(), cached_at: std::time::Instant::now(), }, ); } - models + 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) { @@ -601,6 +737,9 @@ impl Discovery { 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 { @@ -851,17 +990,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) } @@ -998,10 +1140,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() @@ -1013,22 +1164,11 @@ impl Discovery { return Ok(models); } - let mut providers = match self.fetch_providers().await { + let providers = match self.fetch_providers().await { Ok(providers) => providers, Err(_err) if !models.is_empty() => return Ok(models), Err(err) => return Err(err), }; - let discovered_models = self.discover_custom_models_from_catalog(&providers).await; - for discovered_model in discovered_models { - let Some(provider) = providers.get_mut(&discovered_model.provider_id) else { - continue; - }; - provider - .models - .entry(discovered_model.id.clone()) - .or_insert_with(|| discovery_model_from_dialog_model(&discovered_model)); - } - let mut persistent_models = Vec::new(); for (provider_id, provider) in providers { @@ -1039,35 +1179,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(), - }); + for (model_id, model) in &provider.models { + if let Some(model) = dialog_model(&provider, model_id, model) { + persistent_models.push(model); } } } @@ -1087,9 +1201,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() } @@ -1108,9 +1220,7 @@ impl Discovery { } pub fn get_model_limit(&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.limit.as_ref().map(|l| l.context) } @@ -1125,9 +1235,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()) } @@ -1153,16 +1261,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() @@ -1319,6 +1418,25 @@ 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") @@ -1404,7 +1522,7 @@ mod tests { api_key: Some("test-key".to_string()), models: HashMap::new(), }; - let discovery = + let mut discovery = Discovery::new_with_custom(Some(HashMap::from([("gateway".to_string(), provider)]))) .expect("discovery"); let catalog = HashMap::from([( @@ -1424,7 +1542,10 @@ mod tests { family: "gpt".to_string(), attachment: true, reasoning: true, - reasoning_options: Vec::new(), + 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, @@ -1432,24 +1553,138 @@ mod tests { release_date: String::new(), last_updated: String::new(), status: None, - modalities: None, + modalities: Some(Modalities { + input: vec!["text".into(), "image".into()], + output: vec!["text".into()], + }), open_weights: false, - cost: None, - limit: None, - provider: None, + 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()), + }), }, )]), }, )]); - let models = discovery - .discover_custom_models_from_catalog(&catalog) - .await; + 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!(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] @@ -1494,13 +1729,13 @@ mod tests { }, )]), }; - let discovery = + let mut discovery = Discovery::new_with_custom(Some(HashMap::from([("gateway".to_string(), provider)]))) .expect("discovery"); - let discovered = discovery - .discover_custom_models_from_catalog(&HashMap::new()) - .await; + 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"); @@ -1513,6 +1748,73 @@ mod tests { 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 {