diff --git a/_docs/config/opencode-compatibility.mdx b/_docs/config/opencode-compatibility.mdx index c5aeff5..951fd6a 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/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..bb3c0bf 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,79 @@ 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 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 +336,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 +425,184 @@ 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(); + } + }; + 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 stored_api_keys = self.custom_provider_discovery_api_keys(); + + 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) = 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 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 +1013,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 +1250,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(()) } } @@ -970,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) @@ -983,6 +1317,202 @@ 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"); + } + + #[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 { @@ -1530,7 +2060,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();