diff --git a/Cargo.lock b/Cargo.lock index fcfec7c8da02..0cfdf6d94c71 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5137,6 +5137,7 @@ dependencies = [ "url", "utoipa 4.2.3", "uuid", + "wiremock", ] [[package]] diff --git a/crates/goose-providers/Cargo.toml b/crates/goose-providers/Cargo.toml index 7fd0180c6fe0..9228b045d498 100644 --- a/crates/goose-providers/Cargo.toml +++ b/crates/goose-providers/Cargo.toml @@ -61,6 +61,7 @@ tempfile = { workspace = true } tokio = { workspace = true, features = ["rt-multi-thread"] } tokio-stream = { workspace = true } env-lock = { workspace = true } +wiremock.workspace = true [[example]] name = "streaming" diff --git a/crates/goose-providers/examples/declarative.rs b/crates/goose-providers/examples/declarative.rs new file mode 100644 index 000000000000..7d63ff0e1c04 --- /dev/null +++ b/crates/goose-providers/examples/declarative.rs @@ -0,0 +1,33 @@ +use anyhow::Result; +use futures::StreamExt; +use goose_providers::{ + base::Provider, conversation::message::Message, declarative::EnvKeyResolver, model::ModelConfig, +}; + +async fn complete(provider: &dyn Provider, model: ModelConfig) -> Result<()> { + let system = "You are a knowledgable geography expert"; + let messages = [Message::user().with_text("what is the capital of France?")]; + let mut stream = provider.stream(&model, system, &messages, &[]).await?; + + while let Some((Some(msg), _)) = stream.next().await.transpose()? { + print!("{}", msg.as_concat_text()); + } + println!(); + + Ok(()) +} + +#[tokio::main] +async fn main() -> Result<()> { + let deepseek = include_str!("deepseek.json"); + let deepseek_model = ModelConfig::new("deepseek-v4-flash"); + let zai = include_str!("zai.json"); + let zai_model = ModelConfig::new("glm-4.5-flash"); + + for (json, model) in [(deepseek, deepseek_model), (zai, zai_model)] { + let provider = goose_providers::declarative::from_json(json, None, EnvKeyResolver {})?; + println!("{}:", provider.get_name()); + complete(provider.as_ref(), model).await?; + } + Ok(()) +} diff --git a/crates/goose-providers/examples/deepseek.json b/crates/goose-providers/examples/deepseek.json new file mode 100644 index 000000000000..1d220744375d --- /dev/null +++ b/crates/goose-providers/examples/deepseek.json @@ -0,0 +1,30 @@ +{ + "name": "deepseek", + "engine": "openai", + "display_name": "DeepSeek", + "description": "Custom DeepSeek provider", + "api_key_env": "DEEPSEEK_API_KEY", + "base_url": "https://api.deepseek.com", + "models": [ + { + "name": "deepseek-chat", + "context_limit": 128000, + "input_token_cost": null, + "output_token_cost": null, + "currency": null, + "supports_cache_control": null + }, + { + "name": "deepseek-reasoner", + "context_limit": 128000, + "input_token_cost": null, + "output_token_cost": null, + "currency": null, + "supports_cache_control": null + } + ], + "headers": null, + "timeout_seconds": null, + "preserves_thinking": true, + "supports_streaming": true +} diff --git a/crates/goose-providers/examples/zai.json b/crates/goose-providers/examples/zai.json new file mode 100644 index 000000000000..afcb16166216 --- /dev/null +++ b/crates/goose-providers/examples/zai.json @@ -0,0 +1,25 @@ +{ + "name": "zai", + "engine": "anthropic", + "display_name": "Z.AI", + "description": "Z.AI GLM models via Anthropic-compatible API.", + "api_key_env": "ZHIPU_API_KEY", + "base_url": "https://api.z.ai/api/anthropic", + "catalog_provider_id": "zai", + "model_doc_link": "https://docs.z.ai/devpack/tool/goose", + "fast_model": "glm-4.5-air", + "preserves_thinking": true, + "models": [ + { "name": "glm-5.1", "context_limit": 200000 }, + { "name": "glm-5", "context_limit": 204800 }, + { "name": "glm-5-turbo", "context_limit": 200000 }, + { "name": "glm-4.7", "context_limit": 204800 }, + { "name": "glm-4.7-flash", "context_limit": 200000 }, + { "name": "glm-4.7-flashx", "context_limit": 200000 }, + { "name": "glm-4.6", "context_limit": 204800 }, + { "name": "glm-4.5", "context_limit": 131072 }, + { "name": "glm-4.5-air", "context_limit": 131072 }, + { "name": "glm-4.5-flash", "context_limit": 131072 } + ], + "supports_streaming": true +} diff --git a/crates/goose-providers/src/anthropic.rs b/crates/goose-providers/src/anthropic.rs index a0c901721c76..33755a6862e8 100644 --- a/crates/goose-providers/src/anthropic.rs +++ b/crates/goose-providers/src/anthropic.rs @@ -1,4 +1,6 @@ +use crate::api_client::{AuthMethod, TlsConfig}; use crate::base::ProviderDescriptor; +use crate::declarative::{DeclarativeProviderConfig, KeyResolver}; use crate::errors::ProviderError; use crate::request_log::{start_log, LoggerHandleExt}; use anyhow::Result; @@ -90,6 +92,24 @@ impl AnthropicProviderBuilder { } } + pub fn api_client(mut self, api_client: ApiClient) -> Self { + self.api_client = api_client; + self + } + + pub fn map_api_client(mut self, f: impl FnOnce(ApiClient) -> ApiClient) -> Self { + self.api_client = f(self.api_client); + self + } + + pub fn try_map_api_client( + mut self, + f: impl FnOnce(ApiClient) -> Result, + ) -> Result { + self.api_client = f(self.api_client)?; + Ok(self) + } + pub fn supports_streaming(mut self, supports_streaming: bool) -> Self { self.supports_streaming = supports_streaming; self @@ -287,3 +307,96 @@ impl Provider for AnthropicProvider { })) } } + +fn format_options_for_provider(preserves_thinking: bool) -> AnthropicFormatOptions { + AnthropicFormatOptions { + preserve_unsigned_thinking: preserves_thinking, + preserve_thinking_context: preserves_thinking, + thinking_disabled: false, + } +} + +pub fn from_declarative_config( + config: DeclarativeProviderConfig, + tls_config: Option, + key_resolver: impl KeyResolver, +) -> Result { + let custom_models = if !config.models.is_empty() { + Some( + config + .models + .iter() + .map(|m| m.name.clone()) + .collect::>(), + ) + } else { + None + }; + + if config.dynamic_models == Some(false) && custom_models.is_none() { + return Err(anyhow::anyhow!( + "Provider '{}' has dynamic_models: false but no static models listed; \ + at least one entry in `models` is required.", + config.name + )); + } + + let api_key = if config.api_key_env.is_empty() { + None + } else { + match key_resolver.resolve_key(config.api_key_env.as_str()) { + Ok(key) => Some(key), + Err(err) => { + if config.requires_auth { + anyhow::bail!("missing required key {}: {}", config.api_key_env, err); + } + None + } + } + }; + + let auth = match api_key { + Some(key) if !key.is_empty() => AuthMethod::ApiKey { + header_name: "x-api-key".to_string(), + key, + }, + _ => AuthMethod::NoAuth, + }; + + let format_options = format_options_for_provider(config.preserves_thinking); + + let mut api_client = ApiClient::new_with_tls(config.base_url, auth, tls_config)?; + + if let Some(headers) = &config.headers { + let mut header_map = reqwest::header::HeaderMap::new(); + header_map.insert( + reqwest::header::HeaderName::from_static("anthropic-version"), + reqwest::header::HeaderValue::from_static(ANTHROPIC_API_VERSION), + ); + for (key, value) in headers { + let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; + let header_value = reqwest::header::HeaderValue::from_str(value)?; + header_map.insert(header_name, header_value); + } + api_client = api_client.with_headers(header_map)?; + } else { + api_client = api_client.with_header("anthropic-version", ANTHROPIC_API_VERSION)?; + } + + let supports_streaming = config.supports_streaming.unwrap_or(true); + + if !supports_streaming { + return Err(anyhow::anyhow!( + "Anthropic provider does not support non-streaming mode. All Claude models support streaming. \ + Please remove 'supports_streaming: false' from your provider configuration." + )); + } + + Ok(AnthropicProviderBuilder::new(api_client) + .supports_streaming(supports_streaming) + .name(config.name.clone()) + .custom_models(custom_models) + .dynamic_models(config.dynamic_models) + .skip_canonical_filtering(config.skip_canonical_filtering) + .format_options(format_options)) +} diff --git a/crates/goose-providers/src/declarative.rs b/crates/goose-providers/src/declarative.rs new file mode 100644 index 000000000000..d3c39929c06b --- /dev/null +++ b/crates/goose-providers/src/declarative.rs @@ -0,0 +1,388 @@ +use std::{collections::HashMap, str::FromStr}; + +use anyhow::Result; +use serde::{Deserialize, Deserializer, Serialize}; +use utoipa::ToSchema; + +use crate::{ + anthropic, + api_client::TlsConfig, + base::{ModelInfo, Provider}, + ollama, openai, +}; + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] +pub struct EnvVarConfig { + pub name: String, + #[serde(default)] + pub required: bool, + #[serde(default)] + pub secret: bool, + /// Defaults to the value of `required` if not specified. + /// UIs may use this to feature this config value more prominently. + pub primary: Option, + pub description: Option, + pub default: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, ToSchema)] +#[serde(rename_all = "lowercase")] +pub enum ProviderEngine { + #[serde(alias = "openai_compatible")] + OpenAI, + #[serde(alias = "ollama_compatible")] + Ollama, + #[serde(alias = "anthropic_compatible")] + Anthropic, +} + +impl FromStr for ProviderEngine { + type Err = anyhow::Error; + + fn from_str(engine: &str) -> Result { + match engine.trim().to_lowercase().as_str() { + "openai" | "openai_compatible" => Ok(Self::OpenAI), + "anthropic" | "anthropic_compatible" => Ok(Self::Anthropic), + "ollama" | "ollama_compatible" => Ok(Self::Ollama), + _ => Err(anyhow::anyhow!("Invalid provider type: {}", engine)), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] +pub struct DeclarativeProviderConfig { + pub name: String, + pub engine: ProviderEngine, + pub display_name: String, + pub description: Option, + #[serde(default)] + pub api_key_env: String, + pub base_url: String, + pub models: Vec, + pub headers: Option>, + pub timeout_seconds: Option, + pub supports_streaming: Option, + #[serde(default = "default_requires_auth")] + pub requires_auth: bool, + #[serde(default)] + pub catalog_provider_id: Option, + #[serde(default)] + pub base_path: Option, + #[serde(default)] + pub env_vars: Option>, + /// Controls whether `fetch_supported_models` calls the provider's `/v1/models` + /// endpoint or returns the static `models` list directly. + /// + /// - `Some(false)` + non-empty `models`: return the static list; no API call. + /// Construction fails if `models` is empty. + /// - `Some(true)` or `None`: try the API; fall back to `models` on 404. + #[serde(default)] + pub dynamic_models: Option, + #[serde(default)] + pub skip_canonical_filtering: bool, + #[serde(default, deserialize_with = "deserialize_non_empty_string")] + pub model_doc_link: Option, + #[serde(default)] + pub setup_steps: Vec, + #[serde(default, deserialize_with = "deserialize_non_empty_string")] + pub fast_model: Option, + #[serde(default)] + pub preserves_thinking: bool, +} + +fn default_requires_auth() -> bool { + true +} + +fn should_preserve_thinking_by_default(engine: &ProviderEngine) -> bool { + matches!(engine, ProviderEngine::OpenAI) +} + +/// Deserialize an optional string, treating empty/whitespace-only values as None. +fn deserialize_non_empty_string<'de, D>(deserializer: D) -> Result, D::Error> +where + D: Deserializer<'de>, +{ + let opt: Option = Option::deserialize(deserializer)?; + Ok(opt.filter(|s| !s.trim().is_empty())) +} + +impl DeclarativeProviderConfig { + pub fn id(&self) -> &str { + &self.name + } + + pub fn display_name(&self) -> &str { + &self.display_name + } + + pub fn models(&self) -> &[ModelInfo] { + &self.models + } +} + +pub trait KeyResolver { + type Error: std::error::Error + Send + Sync + 'static; + + fn resolve_key(&self, key: &str) -> std::result::Result; +} + +pub struct EnvKeyResolver; + +impl EnvKeyResolver { + pub fn new() -> Self { + EnvKeyResolver {} + } +} + +impl Default for EnvKeyResolver { + fn default() -> Self { + Self::new() + } +} + +impl KeyResolver for EnvKeyResolver { + type Error = std::env::VarError; + + fn resolve_key(&self, key: &str) -> std::result::Result { + std::env::var(key) + } +} + +fn expand_env_vars(template: &str, env_vars: &[EnvVarConfig]) -> Result { + let mut result = template.to_string(); + + for var in env_vars { + let placeholder = format!("${{{}}}", var.name); + if !result.contains(&placeholder) { + continue; + } + + let value = match std::env::var(&var.name) { + Ok(value) => value, + Err(_) => match &var.default { + Some(default) => default.clone(), + None if var.required => { + anyhow::bail!("Required environment variable {} is not set", var.name) + } + None => continue, + }, + }; + + result = result.replace(&placeholder, &value); + } + + Ok(result) +} + +fn resolve_config(config: &mut DeclarativeProviderConfig) -> Result<()> { + if let Some(env_vars) = &config.env_vars { + config.base_url = expand_env_vars(&config.base_url, env_vars)?; + + for var in env_vars { + if var.name.ends_with("_STREAMING") { + let value = std::env::var(&var.name) + .ok() + .or_else(|| var.default.clone()) + .map(|value| value.eq_ignore_ascii_case("true")); + if let Some(value) = value { + config.supports_streaming = Some(value); + } + } + } + } + + Ok(()) +} + +fn config_from_json(json: &str) -> Result { + let raw: serde_json::Value = serde_json::from_str(json)?; + let preserves_thinking_was_set = raw.get("preserves_thinking").is_some(); + let mut config: DeclarativeProviderConfig = serde_json::from_value(raw)?; + + if !preserves_thinking_was_set { + config.preserves_thinking = should_preserve_thinking_by_default(&config.engine); + } + + resolve_config(&mut config)?; + Ok(config) +} + +pub fn from_json( + json: &str, + tls_config: Option, + key_resolver: impl KeyResolver, +) -> Result> { + let config = config_from_json(json)?; + + match config.engine { + ProviderEngine::OpenAI => openai::from_declarative_config(config, tls_config, key_resolver) + .map(|provider| Box::new(provider.build()) as Box), + ProviderEngine::Ollama => ollama::from_declarative_config(config, tls_config, key_resolver) + .map(|provider| Box::new(provider.build()) as Box), + ProviderEngine::Anthropic => { + anthropic::from_declarative_config(config, tls_config, key_resolver) + .map(|provider| Box::new(provider.build()) as Box) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn model_json() -> serde_json::Value { + json!({ + "name": "test-model", + "context_limit": 4096, + "input_token_cost": null, + "output_token_cost": null, + "currency": null, + "supports_cache_control": null, + "reasoning": false + }) + } + + #[test] + fn provider_engine_deserializes_compatible_aliases() { + let openai: DeclarativeProviderConfig = serde_json::from_value(json!({ + "name": "test-openai", + "engine": "openai_compatible", + "display_name": "Test OpenAI", + "base_url": "http://localhost:1234", + "models": [model_json()] + })) + .unwrap(); + assert_eq!(openai.engine, ProviderEngine::OpenAI); + + let anthropic: DeclarativeProviderConfig = serde_json::from_value(json!({ + "name": "test-anthropic", + "engine": "anthropic_compatible", + "display_name": "Test Anthropic", + "base_url": "http://localhost:1234", + "models": [model_json()] + })) + .unwrap(); + assert_eq!(anthropic.engine, ProviderEngine::Anthropic); + + let ollama: DeclarativeProviderConfig = serde_json::from_value(json!({ + "name": "test-ollama", + "engine": "ollama_compatible", + "display_name": "Test Ollama", + "base_url": "http://localhost:11434", + "models": [model_json()] + })) + .unwrap(); + assert_eq!(ollama.engine, ProviderEngine::Ollama); + } + + #[test] + fn from_json_defaults_openai_preserves_thinking_to_true() { + let json = json!({ + "name": "test-provider", + "engine": "openai", + "display_name": "Test Provider", + "base_url": "http://localhost:1234/v1/chat/completions", + "models": [model_json()], + "requires_auth": false, + "dynamic_models": false + }) + .to_string(); + + let config = config_from_json(&json).unwrap(); + + assert!(config.preserves_thinking); + } + + #[test] + fn from_json_preserves_explicit_openai_preserves_thinking_false() { + let json = json!({ + "name": "test-provider", + "engine": "openai", + "display_name": "Test Provider", + "base_url": "http://localhost:1234/v1/chat/completions", + "models": [model_json()], + "requires_auth": false, + "dynamic_models": false, + "preserves_thinking": false + }) + .to_string(); + + let config = config_from_json(&json).unwrap(); + + assert!(!config.preserves_thinking); + } + + #[test] + fn from_json_expands_base_url_from_env_var_default() { + let _guard = env_lock::lock_env([("TEST_PROVIDER_HOST", None::<&str>)]); + let json = json!({ + "name": "test-provider", + "engine": "openai", + "display_name": "Test Provider", + "base_url": "${TEST_PROVIDER_HOST}/v1/chat/completions", + "models": [model_json()], + "requires_auth": false, + "dynamic_models": false, + "env_vars": [{ + "name": "TEST_PROVIDER_HOST", + "default": "http://localhost:1234" + }] + }) + .to_string(); + + let provider = from_json(&json, None, EnvKeyResolver).unwrap(); + + assert_eq!(provider.get_name(), "test-provider"); + } + + #[tokio::test] + async fn from_json_ollama_returns_static_models_when_dynamic_models_false() { + let json = json!({ + "name": "test-ollama", + "engine": "ollama", + "display_name": "Test Ollama", + "base_url": "http://localhost:11434", + "models": [model_json()], + "requires_auth": false, + "dynamic_models": false + }) + .to_string(); + + let provider = from_json(&json, None, EnvKeyResolver).unwrap(); + + assert_eq!( + provider.fetch_supported_models().await.unwrap(), + vec!["test-model".to_string()] + ); + } + + #[test] + fn from_json_errors_when_required_env_var_is_missing() { + let _guard = env_lock::lock_env([("TEST_PROVIDER_REQUIRED_HOST", None::<&str>)]); + let json = json!({ + "name": "test-provider", + "engine": "openai", + "display_name": "Test Provider", + "base_url": "${TEST_PROVIDER_REQUIRED_HOST}/v1/chat/completions", + "models": [model_json()], + "requires_auth": false, + "dynamic_models": false, + "env_vars": [{ + "name": "TEST_PROVIDER_REQUIRED_HOST", + "required": true + }] + }) + .to_string(); + + let err = match from_json(&json, None, EnvKeyResolver) { + Ok(_) => panic!("expected missing required env var error"), + Err(err) => err, + }; + + assert!(err + .to_string() + .contains("Required environment variable TEST_PROVIDER_REQUIRED_HOST is not set")); + } +} diff --git a/crates/goose-providers/src/lib.rs b/crates/goose-providers/src/lib.rs index 99944f92d906..f84486bdc0f8 100644 --- a/crates/goose-providers/src/lib.rs +++ b/crates/goose-providers/src/lib.rs @@ -3,6 +3,7 @@ pub mod api_client; pub mod base; pub mod canonical; pub mod conversation; +pub mod declarative; pub mod errors; pub mod formats; pub mod goose_mode; diff --git a/crates/goose-providers/src/ollama.rs b/crates/goose-providers/src/ollama.rs index 7db1fb374200..9d7fb35cf1d9 100644 --- a/crates/goose-providers/src/ollama.rs +++ b/crates/goose-providers/src/ollama.rs @@ -2,8 +2,10 @@ use super::api_client::ApiClient; use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata}; use super::openai_compatible::handle_status; use super::retry::{ProviderRetry, RetryConfig}; +use crate::api_client::{AuthMethod, TlsConfig}; use crate::base::ProviderDescriptor; use crate::conversation::message::Message; +use crate::declarative::{DeclarativeProviderConfig, KeyResolver}; use crate::errors::ProviderError; use crate::formats::ollama::{create_request, response_to_streaming_message_ollama}; use crate::images::ImageFormat; @@ -13,7 +15,7 @@ use anyhow::{Error, Result}; use async_stream::try_stream; use async_trait::async_trait; use futures::TryStreamExt; -use reqwest::Response; +use reqwest::{Response, StatusCode}; use rmcp::model::Tool; use serde_json::{json, Value}; use std::time::Duration; @@ -21,6 +23,7 @@ use tokio::pin; use tokio_stream::StreamExt; use tokio_util::codec::{FramedRead, LinesCodec}; use tokio_util::io::StreamReader; +use url::Url; pub const OLLAMA_PROVIDER_NAME: &str = "ollama"; pub const OLLAMA_HOST: &str = "localhost"; @@ -79,10 +82,88 @@ pub struct OllamaProvider { #[serde(skip)] api_client: ApiClient, name: String, + custom_models: Option>, + dynamic_models: Option, skip_canonical_filtering: bool, options: OllamaOptions, } +pub struct OllamaProviderBuilder { + api_client: ApiClient, + name: String, + custom_models: Option>, + dynamic_models: Option, + skip_canonical_filtering: bool, + options: OllamaOptions, +} + +impl OllamaProviderBuilder { + pub fn new(api_client: ApiClient) -> Self { + Self { + api_client, + name: OLLAMA_PROVIDER_NAME.to_string(), + custom_models: None, + dynamic_models: None, + skip_canonical_filtering: false, + options: OllamaOptions::default(), + } + } + + pub fn api_client(mut self, api_client: ApiClient) -> Self { + self.api_client = api_client; + self + } + + pub fn map_api_client(mut self, f: impl FnOnce(ApiClient) -> ApiClient) -> Self { + self.api_client = f(self.api_client); + self + } + + pub fn try_map_api_client( + mut self, + f: impl FnOnce(ApiClient) -> Result, + ) -> Result { + self.api_client = f(self.api_client)?; + Ok(self) + } + + pub fn name(mut self, name: impl Into) -> Self { + self.name = name.into(); + self + } + + pub fn custom_models(mut self, custom_models: Option>) -> Self { + self.custom_models = custom_models; + self + } + + pub fn dynamic_models(mut self, dynamic_models: Option) -> Self { + self.dynamic_models = dynamic_models; + self + } + + pub fn skip_canonical_filtering(mut self, skip_canonical_filtering: bool) -> Self { + self.skip_canonical_filtering = skip_canonical_filtering; + self + } + + pub fn options(mut self, options: OllamaOptions) -> Self { + self.options = options; + self + } + + pub fn build(self) -> OllamaProvider { + OllamaProvider { + api_client: self.api_client, + name: self.name, + custom_models: self.custom_models, + dynamic_models: self.dynamic_models, + skip_canonical_filtering: self.skip_canonical_filtering, + options: self.options, + } + } +} + impl OllamaProvider { pub fn new( api_client: ApiClient, @@ -90,12 +171,58 @@ impl OllamaProvider { skip_canonical_filtering: bool, options: OllamaOptions, ) -> Self { - Self { - api_client, - name, - skip_canonical_filtering, - options, + OllamaProviderBuilder::new(api_client) + .name(name) + .skip_canonical_filtering(skip_canonical_filtering) + .options(options) + .build() + } + + pub fn with_options(mut self, options: OllamaOptions) -> Self { + self.options = options; + self + } + + async fn fetch_models_from_api(&self) -> Result, ProviderError> { + let response = self + .api_client + .request("api/tags") + .response_get() + .await + .map_err(|e| ProviderError::RequestFailed(format!("Failed to fetch models: {}", e)))?; + + if response.status() == StatusCode::NOT_FOUND { + return Err(ProviderError::EndpointNotFound( + "Ollama models endpoint not found".to_string(), + )); + } + + if !response.status().is_success() { + return Err(ProviderError::RequestFailed(format!( + "Failed to fetch models: HTTP {}", + response.status() + ))); } + + let json_response = response.json::().await.map_err(|e| { + ProviderError::RequestFailed(format!("Failed to parse response: {}", e)) + })?; + + let models = json_response + .get("models") + .and_then(|m| m.as_array()) + .ok_or_else(|| { + ProviderError::RequestFailed("No models array in response".to_string()) + })?; + + let mut model_names: Vec = models + .iter() + .filter_map(|model| model.get("name").and_then(|n| n.as_str()).map(String::from)) + .collect(); + + model_names.sort(); + + Ok(model_names) } } @@ -136,6 +263,100 @@ fn apply_ollama_options(payload: &mut Value, options: &OllamaOptions, model_conf } } +pub fn from_declarative_config( + config: DeclarativeProviderConfig, + tls_config: Option, + key_resolver: impl KeyResolver, +) -> Result { + let custom_models = if !config.models.is_empty() { + Some( + config + .models + .iter() + .map(|m| m.name.clone()) + .collect::>(), + ) + } else { + None + }; + + if config.dynamic_models == Some(false) && custom_models.is_none() { + return Err(anyhow::anyhow!( + "Provider '{}' has dynamic_models: false but no static models listed; \ + at least one entry in `models` is required.", + config.name + )); + } + + let timeout = Duration::from_secs(config.timeout_seconds.unwrap_or(OLLAMA_TIMEOUT)); + + let base_has_scheme = + config.base_url.starts_with("http://") || config.base_url.starts_with("https://"); + let base = if base_has_scheme { + config.base_url.clone() + } else { + format!("http://{}", config.base_url) + }; + + let mut base_url = Url::parse(&base) + .map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?; + + let is_localhost = matches!(base_url.host_str(), Some("localhost" | "127.0.0.1" | "::1")); + + if base_url.port().is_none() && !base_has_scheme && is_localhost { + base_url + .set_port(Some(OLLAMA_DEFAULT_PORT)) + .map_err(|_| anyhow::anyhow!("Failed to set default port"))?; + } + + let api_key = if config.api_key_env.is_empty() { + None + } else { + match key_resolver.resolve_key(config.api_key_env.as_str()) { + Ok(key) => Some(key), + Err(err) => { + if config.requires_auth { + anyhow::bail!("missing required key {}: {}", config.api_key_env, err); + } + None + } + } + }; + + let auth = match api_key { + Some(key) if !key.is_empty() => AuthMethod::BearerToken(key), + _ => AuthMethod::NoAuth, + }; + + let mut api_client = + ApiClient::with_timeout_and_tls(base_url.to_string(), auth, timeout, tls_config)?; + + if let Some(headers) = &config.headers { + let mut header_map = reqwest::header::HeaderMap::new(); + for (key, value) in headers { + let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; + let header_value = reqwest::header::HeaderValue::from_str(value)?; + header_map.insert(header_name, header_value); + } + api_client = api_client.with_headers(header_map)?; + } + + let supports_streaming = config.supports_streaming.unwrap_or(true); + + if !supports_streaming { + return Err(anyhow::anyhow!( + "Ollama provider does not support non-streaming mode. All Ollama models support streaming. \ + Please remove 'supports_streaming: false' from your provider configuration." + )); + } + + Ok(OllamaProviderBuilder::new(api_client) + .name(config.name.clone()) + .custom_models(custom_models) + .dynamic_models(config.dynamic_models) + .skip_canonical_filtering(config.skip_canonical_filtering)) +} + impl ProviderDescriptor for OllamaProvider { fn metadata() -> ProviderMetadata { ProviderMetadata::new( @@ -213,39 +434,26 @@ impl Provider for OllamaProvider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { - let response = self - .api_client - .request("api/tags") - .response_get() - .await - .map_err(|e| ProviderError::RequestFailed(format!("Failed to fetch models: {}", e)))?; + if let Some(custom_models) = &self.custom_models { + if self.dynamic_models == Some(false) { + return Ok(custom_models.clone()); + } - if !response.status().is_success() { - return Err(ProviderError::RequestFailed(format!( - "Failed to fetch models: HTTP {}", - response.status() - ))); + match self.fetch_models_from_api().await { + Ok(models) => return Ok(models), + Err(e) if e.is_endpoint_not_found() => { + tracing::debug!( + "Models endpoint not implemented for provider '{}' ({}), using predefined list", + self.name, + e + ); + return Ok(custom_models.clone()); + } + Err(e) => return Err(e), + } } - let json_response = response.json::().await.map_err(|e| { - ProviderError::RequestFailed(format!("Failed to parse response: {}", e)) - })?; - - let models = json_response - .get("models") - .and_then(|m| m.as_array()) - .ok_or_else(|| { - ProviderError::RequestFailed("No models array in response".to_string()) - })?; - - let mut model_names: Vec = models - .iter() - .filter_map(|model| model.get("name").and_then(|n| n.as_str()).map(String::from)) - .collect(); - - model_names.sort(); - - Ok(model_names) + self.fetch_models_from_api().await } } @@ -323,6 +531,105 @@ fn stream_ollama( #[cfg(test)] mod tests { use super::*; + use crate::base::ModelInfo; + + fn ollama_config( + dynamic_models: Option, + models: Vec, + ) -> DeclarativeProviderConfig { + ollama_config_with_base_url(dynamic_models, models, "http://localhost:11434") + } + + fn ollama_config_with_base_url( + dynamic_models: Option, + models: Vec, + base_url: &str, + ) -> DeclarativeProviderConfig { + DeclarativeProviderConfig { + name: "test-ollama".to_string(), + engine: crate::declarative::ProviderEngine::Ollama, + display_name: "Test Ollama".to_string(), + description: None, + api_key_env: String::new(), + base_url: base_url.to_string(), + models, + headers: None, + timeout_seconds: None, + supports_streaming: None, + requires_auth: false, + catalog_provider_id: None, + base_path: None, + env_vars: None, + dynamic_models, + skip_canonical_filtering: false, + model_doc_link: None, + setup_steps: vec![], + fast_model: None, + preserves_thinking: false, + } + } + + #[tokio::test] + async fn fetch_supported_models_uses_static_models_when_dynamic_models_false() { + let provider = from_declarative_config( + ollama_config(Some(false), vec![ModelInfo::new("static-model", 4096)]), + None, + crate::declarative::EnvKeyResolver, + ) + .unwrap() + .build(); + + assert_eq!( + provider.fetch_supported_models().await.unwrap(), + vec!["static-model".to_string()] + ); + } + + #[test] + fn from_custom_config_requires_static_models_when_dynamic_models_false() { + let err = from_declarative_config( + ollama_config(Some(false), vec![]), + None, + crate::declarative::EnvKeyResolver, + ) + .err() + .expect("expected static models validation error"); + + assert!(err + .to_string() + .contains("dynamic_models: false but no static models listed")); + } + + #[tokio::test] + async fn fetch_supported_models_falls_back_to_static_models_on_404() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/tags")) + .respond_with(ResponseTemplate::new(404)) + .expect(1) + .mount(&server) + .await; + + let provider = from_declarative_config( + ollama_config_with_base_url( + None, + vec![ModelInfo::new("static-model", 4096)], + &server.uri(), + ), + None, + crate::declarative::EnvKeyResolver, + ) + .unwrap() + .build(); + + assert_eq!( + provider.fetch_supported_models().await.unwrap(), + vec!["static-model".to_string()] + ); + } #[test] fn test_apply_ollama_options_uses_input_limit() { diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index e44acbb60a04..a5e6e8630fd3 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -1,8 +1,10 @@ use super::api_client::ApiClient; use super::base::{ConfigKey, ModelInfo, Provider, ProviderMetadata}; use super::retry::ProviderRetry; +use crate::api_client::{AuthMethod, TlsConfig}; use crate::conversation::message::Message; use crate::conversation::token_usage::ProviderUsage; +use crate::declarative::{DeclarativeProviderConfig, KeyResolver}; use crate::errors::ProviderError; use crate::formats::openai::is_openai_responses_model; use crate::formats::openai::{ @@ -61,6 +63,7 @@ pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[ ]; pub const OPEN_AI_DOC_URL: &str = "https://platform.openai.com/docs/models"; +const DEFAULT_TIMEOUT_SECONDS: u64 = 600; type OpenAiBaseUrlParts = (String, Vec<(String, String)>, bool); @@ -178,6 +181,19 @@ impl OpenAiProviderBuilder { self } + pub fn map_api_client(mut self, f: impl FnOnce(ApiClient) -> ApiClient) -> Self { + self.api_client = f(self.api_client); + self + } + + pub fn try_map_api_client( + mut self, + f: impl FnOnce(ApiClient) -> Result, + ) -> Result { + self.api_client = f(self.api_client)?; + Ok(self) + } + pub fn base_path(mut self, base_path: impl Into) -> Self { self.base_path = base_path.into(); self @@ -687,6 +703,97 @@ impl Provider for OpenAiProvider { } } +pub fn from_declarative_config( + config: DeclarativeProviderConfig, + tls_config: Option, + key_resolver: impl KeyResolver, +) -> Result { + let custom_models = if !config.models.is_empty() { + Some( + config + .models + .iter() + .map(|m| m.name.clone()) + .collect::>(), + ) + } else { + None + }; + + if config.dynamic_models == Some(false) && custom_models.is_none() { + return Err(anyhow::anyhow!( + "Provider '{}' has dynamic_models: false but no static models listed; \ + at least one entry in `models` is required.", + config.name + )); + } + + let api_key = if config.api_key_env.is_empty() { + None + } else { + match key_resolver.resolve_key(config.api_key_env.as_str()) { + Ok(key) => Some(key), + Err(err) => { + if config.requires_auth { + anyhow::bail!("missing required key {}: {}", config.api_key_env, err); + } + None + } + } + }; + + let normalized_base_url = ensure_url_scheme(&config.base_url); + let url = url::Url::parse(&normalized_base_url) + .map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?; + + let host = url[..url::Position::BeforePath].to_string(); + let base_path = if let Some(ref explicit_path) = config.base_path { + explicit_path.trim_start_matches('/').to_string() + } else { + derive_base_path(url.path()) + }; + + let timeout_secs = config.timeout_seconds.unwrap_or(DEFAULT_TIMEOUT_SECONDS); + + let auth = match api_key { + Some(key) if !key.is_empty() => AuthMethod::BearerToken(key), + _ => AuthMethod::NoAuth, + }; + let mut api_client = ApiClient::with_timeout_and_tls( + host, + auth, + std::time::Duration::from_secs(timeout_secs), + tls_config, + )?; + + if let Some(query) = url.query() { + let query_params = url::form_urlencoded::parse(query.as_bytes()) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + api_client = api_client.with_query(query_params); + } + + if let Some(headers) = &config.headers { + let mut header_map = reqwest::header::HeaderMap::new(); + for (key, value) in headers { + let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; + let header_value = reqwest::header::HeaderValue::from_str(value)?; + header_map.insert(header_name, header_value); + } + api_client = api_client.with_headers(header_map)?; + } + + Ok(OpenAiProviderBuilder::new(api_client) + .base_path(base_path) + .custom_headers(config.headers) + .supports_streaming(config.supports_streaming.unwrap_or(true)) + .name(config.name.clone()) + .custom_models(custom_models) + .dynamic_models(config.dynamic_models) + .skip_canonical_filtering(config.skip_canonical_filtering) + .preserve_thinking_context(config.preserves_thinking)) +} + pub fn parse_custom_headers(s: String) -> HashMap { s.split(',') .filter_map(|header| { @@ -698,6 +805,26 @@ pub fn parse_custom_headers(s: String) -> HashMap { .collect() } +pub fn derive_base_path(url_path: &str) -> String { + let stripped = url_path.trim_start_matches('/'); + let normalized = stripped.trim_end_matches('/'); + if normalized.is_empty() { + "v1/chat/completions".to_string() + } else if normalized.ends_with("chat/completions") { + stripped.to_string() + } else if ends_with_version_segment(normalized) { + format!("{}/chat/completions", normalized) + } else { + format!("{}/v1/chat/completions", normalized) + } +} + +fn ends_with_version_segment(path: &str) -> bool { + let last = path.rsplit('/').next().unwrap_or(path); + last.strip_prefix('v') + .is_some_and(|rest| !rest.is_empty() && rest.bytes().all(|b| b.is_ascii_digit())) +} + #[cfg(test)] mod tests { use super::*; @@ -950,6 +1077,60 @@ mod tests { ); } + fn custom_config(base_url: &str) -> DeclarativeProviderConfig { + DeclarativeProviderConfig { + name: "test-openai".to_string(), + engine: crate::declarative::ProviderEngine::OpenAI, + display_name: "Test OpenAI".to_string(), + description: None, + api_key_env: String::new(), + base_url: base_url.to_string(), + models: vec![crate::base::ModelInfo::new("test-model", 4096)], + headers: None, + timeout_seconds: None, + supports_streaming: None, + requires_auth: false, + catalog_provider_id: None, + base_path: None, + env_vars: None, + dynamic_models: Some(false), + skip_canonical_filtering: false, + model_doc_link: None, + setup_steps: vec![], + fast_model: None, + preserves_thinking: false, + } + } + + #[test] + fn from_custom_config_preserves_ipv6_authority() { + let provider = from_declarative_config( + custom_config("http://[::1]:1234/v1"), + None, + crate::declarative::EnvKeyResolver, + ) + .unwrap() + .build(); + + assert_eq!(provider.api_client.host(), "http://[::1]:1234"); + } + + #[test] + fn from_custom_config_preserves_userinfo_authority() { + let provider = from_declarative_config( + custom_config("https://user:pass@gateway.example/v1"), + None, + crate::declarative::EnvKeyResolver, + ) + .unwrap() + .build(); + + assert_eq!( + provider.api_client.host(), + "https://user:pass@gateway.example" + ); + } + #[test] fn parse_n_ctx_falls_back_to_sole_entry_when_id_differs() { let body = json!({ @@ -970,4 +1151,37 @@ mod tests { }); assert_eq!(parse_n_ctx_from_models(&body, "model-c"), None); } + + #[test] + fn derive_base_path_not_removing_api_path() { + let r = derive_base_path("https://opencode.ai/zen/go"); + assert_eq!(r, "https://opencode.ai/zen/go/v1/chat/completions"); + } + + #[test] + fn derive_base_path_should_support_v1() { + let r = derive_base_path("https://opencode.ai/zen/go/v1"); + assert_eq!(r, "https://opencode.ai/zen/go/v1/chat/completions"); + } + + #[test] + fn derive_base_path_should_support_no_base_path() { + let r = derive_base_path("https://opencode.ai/"); + assert_eq!(r, "https://opencode.ai/v1/chat/completions"); + } + + #[test] + fn derive_base_path_preserves_non_v1_version_prefix() { + // Zhipu's default base_url is https://open.bigmodel.cn/api/paas/v4 and + // from_custom_config passes url.path() ("/api/paas/v4") here. The + // existing /api/paas/v4 version must not gain an extra /v1 segment. + let r = derive_base_path("/api/paas/v4"); + assert_eq!(r, "api/paas/v4/chat/completions"); + } + + #[test] + fn derive_base_path_does_not_treat_v_word_as_version() { + let r = derive_base_path("/api/voice"); + assert_eq!(r, "api/voice/v1/chat/completions"); + } } diff --git a/crates/goose/src/config/declarative_providers.rs b/crates/goose/src/config/declarative_providers.rs index 9b3d1ee182aa..f241bf8a18c7 100644 --- a/crates/goose/src/config/declarative_providers.rs +++ b/crates/goose/src/config/declarative_providers.rs @@ -10,126 +10,26 @@ use crate::providers::openai_def::OpenAiProviderDef; use anyhow::Result; use include_dir::{include_dir, Dir}; use once_cell::sync::Lazy; -use serde::{Deserialize, Deserializer, Serialize}; +use serde::{Deserialize, Serialize}; use std::str::FromStr; -/// Deserialize an optional string, treating empty/whitespace-only values as None. -fn deserialize_non_empty_string<'de, D>(deserializer: D) -> Result, D::Error> -where - D: Deserializer<'de>, -{ - let opt: Option = Option::deserialize(deserializer)?; - Ok(opt.filter(|s| !s.trim().is_empty())) -} use std::collections::HashMap; use std::path::{Path, PathBuf}; use std::sync::Mutex; use utoipa::ToSchema; +pub use goose_providers::declarative::*; + static FIXED_PROVIDERS: Dir = include_dir!("$CARGO_MANIFEST_DIR/src/providers/declarative"); pub fn custom_providers_dir() -> std::path::PathBuf { Paths::config_dir().join("custom_providers") } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, ToSchema)] -#[serde(rename_all = "lowercase")] -pub enum ProviderEngine { - OpenAI, - Ollama, - Anthropic, -} - -impl FromStr for ProviderEngine { - type Err = anyhow::Error; - - fn from_str(engine: &str) -> Result { - match engine.trim().to_lowercase().as_str() { - "openai" | "openai_compatible" => Ok(Self::OpenAI), - "anthropic" | "anthropic_compatible" => Ok(Self::Anthropic), - "ollama" | "ollama_compatible" => Ok(Self::Ollama), - _ => Err(anyhow::anyhow!("Invalid provider type: {}", engine)), - } - } -} - -#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] -pub struct EnvVarConfig { - pub name: String, - #[serde(default)] - pub required: bool, - #[serde(default)] - pub secret: bool, - /// When true, the field is shown prominently in the UI (not collapsed). - /// Defaults to the value of `required` if not specified. - pub primary: Option, - pub description: Option, - pub default: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] -pub struct DeclarativeProviderConfig { - pub name: String, - pub engine: ProviderEngine, - pub display_name: String, - pub description: Option, - #[serde(default)] - pub api_key_env: String, - pub base_url: String, - pub models: Vec, - pub headers: Option>, - pub timeout_seconds: Option, - pub supports_streaming: Option, - #[serde(default = "default_requires_auth")] - pub requires_auth: bool, - #[serde(default)] - pub catalog_provider_id: Option, - #[serde(default)] - pub base_path: Option, - #[serde(default)] - pub env_vars: Option>, - /// Controls whether `fetch_supported_models` calls the provider's `/v1/models` - /// endpoint or returns the static `models` list directly. - /// - /// - `Some(false)` + non-empty `models`: return the static list; no API call. - /// Construction fails if `models` is empty. - /// - `Some(true)` or `None`: try the API; fall back to `models` on 404. - #[serde(default)] - pub dynamic_models: Option, - #[serde(default)] - pub skip_canonical_filtering: bool, - #[serde(default, deserialize_with = "deserialize_non_empty_string")] - pub model_doc_link: Option, - #[serde(default)] - pub setup_steps: Vec, - #[serde(default, deserialize_with = "deserialize_non_empty_string")] - pub fast_model: Option, - #[serde(default)] - pub preserves_thinking: bool, -} - -fn default_requires_auth() -> bool { - true -} - fn should_preserve_thinking_by_default(engine: &ProviderEngine) -> bool { matches!(engine, ProviderEngine::OpenAI) } -impl DeclarativeProviderConfig { - pub fn id(&self) -> &str { - &self.name - } - - pub fn display_name(&self) -> &str { - &self.display_name - } - - pub fn models(&self) -> &[ModelInfo] { - &self.models - } -} - /// Expand `${VAR_NAME}` placeholders in a template string using the given env var configs. /// Resolves values via Config (secret if `secret`, param otherwise), falls back to `default`. /// Returns an error if a `required` var is missing. @@ -464,6 +364,7 @@ pub fn load_provider(id: &str) -> Result { Err(anyhow::anyhow!("Provider not found: {}", id)) } + pub fn load_custom_providers(dir: &Path) -> Result> { if !dir.exists() { return Ok(Vec::new()); diff --git a/crates/goose/src/providers/anthropic_def.rs b/crates/goose/src/providers/anthropic_def.rs index b0e695365cfe..cadbb9a85f6f 100644 --- a/crates/goose/src/providers/anthropic_def.rs +++ b/crates/goose/src/providers/anthropic_def.rs @@ -1,12 +1,14 @@ use anyhow::Result; use futures::future::BoxFuture; -use crate::{config::DeclarativeProviderConfig, providers::base::ProviderDef}; +use crate::{ + config::{Config, DeclarativeProviderConfig}, + providers::{base::ProviderDef, custom_provider_config::ConfigKeyResolver}, +}; use goose_providers::{ - anthropic::{AnthropicProvider, AnthropicProviderBuilder, ANTHROPIC_API_VERSION}, - api_client::{ApiClient, AuthMethod}, + anthropic::{self, AnthropicProvider, AnthropicProviderBuilder, ANTHROPIC_API_VERSION}, + api_client::{ApiClient, AuthMethod, TlsConfig}, base::ProviderDescriptor, - formats::anthropic::AnthropicFormatOptions, }; pub struct AnthropicProviderDef; @@ -51,79 +53,17 @@ async fn from_env( pub fn from_custom_config( config: DeclarativeProviderConfig, - tls_config: Option, + tls_config: Option, ) -> Result { - let custom_models = if !config.models.is_empty() { - Some( - config - .models - .iter() - .map(|m| m.name.clone()) - .collect::>(), - ) - } else { - None - }; - - if config.dynamic_models == Some(false) && custom_models.is_none() { - return Err(anyhow::anyhow!( - "Provider '{}' has dynamic_models: false but no static models listed; \ - at least one entry in `models` is required.", - config.name - )); - } - - let global_config = crate::config::Config::global(); - let api_key: String = global_config - .get_secret(&config.api_key_env) - .map_err(|_| anyhow::anyhow!("Missing API key: {}", config.api_key_env))?; - - let auth = AuthMethod::ApiKey { - header_name: "x-api-key".to_string(), - key: api_key, - }; - - let format_options = format_options_for_provider(config.preserves_thinking); - - let mut api_client = ApiClient::new_with_tls(config.base_url, auth, tls_config)? - .with_request_builder(crate::session_context::session_id_request_builder()) - .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; - - if let Some(headers) = &config.headers { - let mut header_map = reqwest::header::HeaderMap::new(); - for (key, value) in headers { - let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; - let header_value = reqwest::header::HeaderValue::from_str(value)?; - header_map.insert(header_name, header_value); - } - api_client = api_client.with_headers(header_map)?; - } - - let supports_streaming = config.supports_streaming.unwrap_or(true); - - if !supports_streaming { - return Err(anyhow::anyhow!( - "Anthropic provider does not support non-streaming mode. All Claude models support streaming. \ - Please remove 'supports_streaming: false' from your provider configuration." - )); - } - - Ok(AnthropicProviderBuilder::new(api_client) - .supports_streaming(supports_streaming) - .name(config.name.clone()) - .custom_models(custom_models) - .dynamic_models(config.dynamic_models) - .skip_canonical_filtering(config.skip_canonical_filtering) - .format_options(format_options) - .build()) -} - -fn format_options_for_provider(preserves_thinking: bool) -> AnthropicFormatOptions { - AnthropicFormatOptions { - preserve_unsigned_thinking: preserves_thinking, - preserve_thinking_context: preserves_thinking, - thinking_disabled: false, - } + anthropic::from_declarative_config(config, tls_config, ConfigKeyResolver::new(Config::global())) + .map(|builder| { + builder + .map_api_client(|api_client| { + api_client + .with_request_builder(crate::session_context::session_id_request_builder()) + }) + .build() + }) } #[cfg(test)] diff --git a/crates/goose/src/providers/custom_provider_config.rs b/crates/goose/src/providers/custom_provider_config.rs new file mode 100644 index 000000000000..aaf1d56f3312 --- /dev/null +++ b/crates/goose/src/providers/custom_provider_config.rs @@ -0,0 +1,21 @@ +use goose_providers::declarative::KeyResolver; + +use crate::config::{Config, ConfigError}; + +pub struct ConfigKeyResolver<'a> { + config: &'a Config, +} + +impl<'a> ConfigKeyResolver<'a> { + pub fn new(config: &'a Config) -> Self { + Self { config } + } +} + +impl<'a> KeyResolver for ConfigKeyResolver<'a> { + type Error = ConfigError; + + fn resolve_key(&self, key: &str) -> std::result::Result { + self.config.get_secret(key) + } +} diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 3b9cb867f3fa..dd91b1d94570 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -28,6 +28,7 @@ pub mod codex; pub mod codex_acp; pub mod copilot_acp; pub mod cursor_agent; +pub mod custom_provider_config; pub mod databricks; pub mod databricks_auth; pub mod databricks_v2; diff --git a/crates/goose/src/providers/ollama_def.rs b/crates/goose/src/providers/ollama_def.rs index 24a15e055f33..9022c81a5f65 100644 --- a/crates/goose/src/providers/ollama_def.rs +++ b/crates/goose/src/providers/ollama_def.rs @@ -5,14 +5,16 @@ use futures::future::BoxFuture; use url::Url; use crate::{ - config::declarative_providers::DeclarativeProviderConfig, providers::base::ProviderDef, + config::{declarative_providers::DeclarativeProviderConfig, Config}, + providers::{base::ProviderDef, custom_provider_config::ConfigKeyResolver}, }; use goose_providers::{ api_client::{ApiClient, AuthMethod}, base::ProviderDescriptor, ollama::{ - OllamaOptions, OllamaProvider, OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS, OLLAMA_DEFAULT_PORT, - OLLAMA_HOST, OLLAMA_PROVIDER_NAME, OLLAMA_TIMEOUT, + self, OllamaOptions, OllamaProvider, OllamaProviderBuilder, + OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS, OLLAMA_DEFAULT_PORT, OLLAMA_HOST, OLLAMA_PROVIDER_NAME, + OLLAMA_TIMEOUT, }, }; @@ -71,73 +73,26 @@ pub async fn from_env( )? .with_request_builder(crate::session_context::session_id_request_builder()); - Ok(OllamaProvider::new( - api_client, - OLLAMA_PROVIDER_NAME.to_string(), - false, - options_from_config(), - )) + Ok(OllamaProviderBuilder::new(api_client) + .name(OLLAMA_PROVIDER_NAME) + .options(options_from_config()) + .build()) } pub fn from_custom_config( config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { - let timeout = Duration::from_secs(config.timeout_seconds.unwrap_or(OLLAMA_TIMEOUT)); - - let base = if config.base_url.starts_with("http://") || config.base_url.starts_with("https://") - { - config.base_url.clone() - } else { - format!("http://{}", config.base_url) - }; - - let mut base_url = Url::parse(&base) - .map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?; - - let explicit_default_port = - config.base_url.ends_with(":80") || config.base_url.ends_with(":443"); - let is_https = base_url.scheme() == "https"; - - if base_url.port().is_none() && !explicit_default_port && !is_https { - base_url - .set_port(Some(OLLAMA_DEFAULT_PORT)) - .map_err(|_| anyhow::anyhow!("Failed to set default port"))?; - } - - let mut api_client = ApiClient::with_timeout_and_tls( - base_url.to_string(), - AuthMethod::NoAuth, - timeout, - tls_config, - )? - .with_request_builder(crate::session_context::session_id_request_builder()); - - if let Some(headers) = &config.headers { - let mut header_map = reqwest::header::HeaderMap::new(); - for (key, value) in headers { - let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; - let header_value = reqwest::header::HeaderValue::from_str(value)?; - header_map.insert(header_name, header_value); - } - api_client = api_client.with_headers(header_map)?; - } - - let supports_streaming = config.supports_streaming.unwrap_or(true); - - if !supports_streaming { - return Err(anyhow::anyhow!( - "Ollama provider does not support non-streaming mode. All Ollama models support streaming. \ - Please remove 'supports_streaming: false' from your provider configuration." - )); - } - - Ok(OllamaProvider::new( - api_client, - config.name.clone(), - config.skip_canonical_filtering, - options_from_config(), - )) + ollama::from_declarative_config(config, tls_config, ConfigKeyResolver::new(Config::global())) + .map(|builder| { + builder + .map_api_client(|api_client| { + api_client + .with_request_builder(crate::session_context::session_id_request_builder()) + }) + .options(options_from_config()) + .build() + }) } pub fn options_from_config() -> OllamaOptions { diff --git a/crates/goose/src/providers/openai_def.rs b/crates/goose/src/providers/openai_def.rs index 208d2cc7570c..679a32cbb7d8 100644 --- a/crates/goose/src/providers/openai_def.rs +++ b/crates/goose/src/providers/openai_def.rs @@ -4,12 +4,13 @@ use goose_providers::base::ProviderDescriptor; use std::collections::HashMap; use crate::config::declarative_providers::DeclarativeProviderConfig; +use crate::config::Config; use crate::providers::base::{ProviderDef, DEFAULT_PROVIDER_TIMEOUT_SECS}; +use crate::providers::custom_provider_config::ConfigKeyResolver; use goose_providers::api_client::{ApiClient, AuthMethod}; use goose_providers::openai::{ - ensure_url_scheme, parse_custom_headers, parse_openai_base_url, OpenAiProvider, - OpenAiProviderBuilder, OPEN_AI_DEFAULT_BASE_PATH, OPEN_AI_DEFAULT_FAST_MODEL, - OPEN_AI_VERSIONLESS_BASE_PATH, + parse_custom_headers, parse_openai_base_url, OpenAiProvider, OpenAiProviderBuilder, + OPEN_AI_DEFAULT_BASE_PATH, OPEN_AI_DEFAULT_FAST_MODEL, OPEN_AI_VERSIONLESS_BASE_PATH, }; pub struct OpenAiProviderDef; @@ -203,85 +204,19 @@ pub fn from_custom_config( config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { - let custom_models = if !config.models.is_empty() { - Some( - config - .models - .iter() - .map(|m| m.name.clone()) - .collect::>(), - ) - } else { - None - }; - - if config.dynamic_models == Some(false) && custom_models.is_none() { - return Err(anyhow::anyhow!( - "Provider '{}' has dynamic_models: false but no static models listed; \ - at least one entry in `models` is required.", - config.name - )); - } - - let global_config = crate::config::Config::global(); - let api_key = resolve_api_key(&config, &|key| global_config.get_secret(key))?; - - let normalized_base_url = ensure_url_scheme(&config.base_url); - let url = url::Url::parse(&normalized_base_url) - .map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?; - - let host = if let Some(port) = url.port() { - format!( - "{}://{}:{}", - url.scheme(), - url.host_str().unwrap_or(""), - port - ) - } else { - format!("{}://{}", url.scheme(), url.host_str().unwrap_or("")) - }; - let base_path = if let Some(ref explicit_path) = config.base_path { - explicit_path.trim_start_matches('/').to_string() - } else { - derive_base_path(url.path()) - }; - - let timeout_secs = config - .timeout_seconds - .unwrap_or(DEFAULT_PROVIDER_TIMEOUT_SECS); - - let auth = match api_key { - Some(key) if !key.is_empty() => AuthMethod::BearerToken(key), - _ => AuthMethod::NoAuth, - }; - let mut api_client = ApiClient::with_timeout_and_tls( - host, - auth, - std::time::Duration::from_secs(timeout_secs), + goose_providers::openai::from_declarative_config( + config, tls_config, - )? - .with_request_builder(crate::session_context::session_id_request_builder()); - - if let Some(headers) = &config.headers { - let mut header_map = reqwest::header::HeaderMap::new(); - for (key, value) in headers { - let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?; - let header_value = reqwest::header::HeaderValue::from_str(value)?; - header_map.insert(header_name, header_value); - } - api_client = api_client.with_headers(header_map)?; - } - - Ok(OpenAiProviderBuilder::new(api_client) - .base_path(base_path) - .custom_headers(config.headers) - .supports_streaming(config.supports_streaming.unwrap_or(true)) - .name(config.name.clone()) - .custom_models(custom_models) - .dynamic_models(config.dynamic_models) - .skip_canonical_filtering(config.skip_canonical_filtering) - .preserve_thinking_context(config.preserves_thinking) - .build()) + ConfigKeyResolver::new(Config::global()), + ) + .map(|builder| { + builder + .map_api_client(|api_client| { + api_client + .with_request_builder(crate::session_context::session_id_request_builder()) + }) + .build() + }) } /// Components extracted from an `OPENAI_BASE_URL` value. @@ -359,26 +294,6 @@ fn is_direct_openai_host(host: &str) -> bool { .unwrap_or(false) } -fn derive_base_path(url_path: &str) -> String { - let stripped = url_path.trim_start_matches('/'); - let normalized = stripped.trim_end_matches('/'); - if normalized.is_empty() { - "v1/chat/completions".to_string() - } else if normalized.ends_with("chat/completions") { - stripped.to_string() - } else if ends_with_version_segment(normalized) { - format!("{}/chat/completions", normalized) - } else { - format!("{}/v1/chat/completions", normalized) - } -} - -fn ends_with_version_segment(path: &str) -> bool { - let last = path.rsplit('/').next().unwrap_or(path); - last.strip_prefix('v') - .is_some_and(|rest| !rest.is_empty() && rest.bytes().all(|b| b.is_ascii_digit())) -} - #[cfg(test)] mod tests { use super::*; @@ -426,12 +341,6 @@ mod tests { assert!(!r.has_v1); } - #[test] - fn derive_base_path_not_removing_api_path() { - let r = derive_base_path("https://opencode.ai/zen/go"); - assert_eq!(r, "https://opencode.ai/zen/go/v1/chat/completions"); - } - #[test] fn is_direct_openai_host_matches_only_openai() { assert!(is_direct_openai_host("https://api.openai.com")); @@ -442,33 +351,6 @@ mod tests { assert!(!is_direct_openai_host("https://router.huggingface.co/v1")); } - #[test] - fn derive_base_path_should_support_v1() { - let r = derive_base_path("https://opencode.ai/zen/go/v1"); - assert_eq!(r, "https://opencode.ai/zen/go/v1/chat/completions"); - } - - #[test] - fn derive_base_path_should_support_no_base_path() { - let r = derive_base_path("https://opencode.ai/"); - assert_eq!(r, "https://opencode.ai/v1/chat/completions"); - } - - #[test] - fn derive_base_path_preserves_non_v1_version_prefix() { - // Zhipu's default base_url is https://open.bigmodel.cn/api/paas/v4 and - // from_custom_config passes url.path() ("/api/paas/v4") here. The - // existing /api/paas/v4 version must not gain an extra /v1 segment. - let r = derive_base_path("/api/paas/v4"); - assert_eq!(r, "api/paas/v4/chat/completions"); - } - - #[test] - fn derive_base_path_does_not_treat_v_word_as_version() { - let r = derive_base_path("/api/voice"); - assert_eq!(r, "api/voice/v1/chat/completions"); - } - #[test] fn parse_base_url_preserves_query_params() { let r = parse_base_url("https://gw.example.com/v1?api-version=2024-02-01").unwrap();