From 056d58c5f2a1618a2f9e006ca6ae6887515defbf Mon Sep 17 00:00:00 2001 From: Yanuar Date: Wed, 23 Sep 2026 15:54:36 +0700 Subject: [PATCH] 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();