diff --git a/crates/goose-provider-types/src/formats.rs b/crates/goose-provider-types/src/formats.rs index 208835687..3db82a6ea 100644 --- a/crates/goose-provider-types/src/formats.rs +++ b/crates/goose-provider-types/src/formats.rs @@ -3,3 +3,5 @@ pub mod databricks; pub mod ollama; pub mod openai; pub mod openai_responses; + +pub mod snowflake; diff --git a/crates/goose/src/providers/formats/snowflake.rs b/crates/goose-provider-types/src/formats/snowflake.rs similarity index 99% rename from crates/goose/src/providers/formats/snowflake.rs rename to crates/goose-provider-types/src/formats/snowflake.rs index bd72fa6e5..1c2ac06ea 100644 --- a/crates/goose/src/providers/formats/snowflake.rs +++ b/crates/goose-provider-types/src/formats/snowflake.rs @@ -1,9 +1,9 @@ use crate::conversation::message::{Message, MessageContent}; +use crate::conversation::token_usage::Usage; +use crate::errors::ProviderError; use crate::mcp_utils::extract_text_from_resource; +use crate::model::ModelConfig; use anyhow::{anyhow, Result}; -use goose_providers::conversation::token_usage::Usage; -use goose_providers::errors::ProviderError; -use goose_providers::model::ModelConfig; use rmcp::model::{object, CallToolRequestParams, Role, Tool}; use rmcp::object; use serde_json::{json, Value}; @@ -560,7 +560,7 @@ data: {"id":"a9537c2c-2017-4906-9817-2456168d89fa","model":"claude-sonnet-4-2025 #[test] fn test_create_request_format() -> Result<()> { use crate::conversation::message::Message; - use goose_providers::model::ModelConfig; + use crate::model::ModelConfig; let model_config = ModelConfig::new("claude-4-sonnet").with_canonical_limits("snowflake"); @@ -669,7 +669,7 @@ data: {"id":"a9537c2c-2017-4906-9817-2456168d89fa","model":"claude-sonnet-4-2025 #[test] fn test_create_request_excludes_tools_for_description() -> Result<()> { use crate::conversation::message::Message; - use goose_providers::model::ModelConfig; + use crate::model::ModelConfig; let model_config = ModelConfig::new("claude-4-sonnet").with_canonical_limits("snowflake"); let system = "Reply with only a description in four words or less"; diff --git a/crates/goose-provider-types/src/utils.rs b/crates/goose-provider-types/src/utils.rs index 3768ad6a3..1b7777ff2 100644 --- a/crates/goose-provider-types/src/utils.rs +++ b/crates/goose-provider-types/src/utils.rs @@ -12,3 +12,16 @@ pub fn sanitize_unicode_tags(text: &str) -> String { .filter(|&c| !is_in_unicode_tag_range(c)) .collect() } + +/// Extract the model name from a JSON object. Common with most providers to have this top level attribute. +pub fn get_model(data: &serde_json::Value) -> String { + if let Some(model) = data.get("model") { + if let Some(model_str) = model.as_str() { + model_str.to_string() + } else { + "Unknown".to_string() + } + } else { + "Unknown".to_string() + } +} diff --git a/crates/goose-providers/src/lib.rs b/crates/goose-providers/src/lib.rs index 4eebda49a..2c1e3a8e7 100644 --- a/crates/goose-providers/src/lib.rs +++ b/crates/goose-providers/src/lib.rs @@ -16,3 +16,5 @@ pub mod openai; pub mod openai_compatible; pub use declarative::declarative_providers::*; + +pub mod snowflake; diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose-providers/src/snowflake.rs similarity index 82% rename from crates/goose/src/providers/snowflake.rs rename to crates/goose-providers/src/snowflake.rs index 41db6df2b..1ae0537d2 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose-providers/src/snowflake.rs @@ -1,23 +1,21 @@ +use crate::conversation::token_usage::ProviderUsage; +use crate::images::ImageFormat; use anyhow::Result; use async_trait::async_trait; -use goose_providers::conversation::token_usage::ProviderUsage; -use goose_providers::images::ImageFormat; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; -use super::api_client::{ApiClient, AuthMethod}; -use super::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; -use super::formats::snowflake::{create_request, get_usage, response_to_message}; -use super::openai_compatible::{map_http_error_to_provider_error, sanitize_url}; -use super::retry::ProviderRetry; -use super::utils::get_model; -use crate::config::ConfigError; +use crate::api_client::{ApiClient, AuthMethod, RequestBuilderDecorator, TlsConfig}; +use crate::base::{ConfigKey, MessageStream, Provider, ProviderMetadata}; use crate::conversation::message::Message; -use goose_providers::errors::ProviderError; +use crate::errors::ProviderError; +use crate::formats::snowflake::{create_request, get_usage, response_to_message}; +use crate::openai_compatible::{map_http_error_to_provider_error, sanitize_url}; +use crate::retry::ProviderRetry; +use crate::utils::get_model; -use futures::future::BoxFuture; -use goose_providers::model::ModelConfig; -use goose_providers::request_log::{start_log, LoggerHandleExt}; +use crate::model::ModelConfig; +use crate::request_log::{start_log, LoggerHandleExt}; use rmcp::model::Tool; const SNOWFLAKE_PROVIDER_NAME: &str = "snowflake"; @@ -58,23 +56,12 @@ pub struct SnowflakeProvider { } impl SnowflakeProvider { - pub async fn from_env( - tls_config: Option, + pub fn new( + mut host: String, + token: String, + tls_config: Option, + request_builder: Option, ) -> Result { - let config = crate::config::Config::global(); - let mut host: Result = config.get_param("SNOWFLAKE_HOST"); - if host.is_err() { - host = config.get_secret("SNOWFLAKE_HOST") - } - if host.is_err() { - return Err(ConfigError::NotFound( - "Did not find SNOWFLAKE_HOST in either config file or keyring".to_string(), - ) - .into()); - } - - let mut host = host?; - // Convert host to lowercase host = host.to_lowercase(); @@ -83,19 +70,6 @@ impl SnowflakeProvider { host = format!("{}.snowflakecomputing.com", host); } - let mut token: Result = config.get_param("SNOWFLAKE_TOKEN"); - - if token.is_err() { - token = config.get_secret("SNOWFLAKE_TOKEN") - } - - if token.is_err() { - return Err(ConfigError::NotFound( - "Did not find SNOWFLAKE_TOKEN in either config file or keyring".to_string(), - ) - .into()); - } - // Ensure host has https:// prefix let base_url = if !host.starts_with("https://") && !host.starts_with("http://") { format!("https://{}", host) @@ -103,10 +77,12 @@ impl SnowflakeProvider { host }; - let auth = AuthMethod::BearerToken(token?); - let api_client = ApiClient::new_with_tls(base_url, auth, tls_config)? - .with_request_builder(crate::session_context::session_id_request_builder()) - .with_header("User-Agent", "goose")?; + let auth = AuthMethod::BearerToken(token); + let mut api_client = ApiClient::new_with_tls(base_url, auth, tls_config)?; + if let Some(request_builder) = request_builder { + api_client = api_client.with_request_builder(request_builder); + } + let api_client = api_client.with_header("User-Agent", "goose")?; Ok(Self { api_client, @@ -297,7 +273,7 @@ impl SnowflakeProvider { } } -impl goose_providers::base::ProviderDescriptor for SnowflakeProvider { +impl crate::base::ProviderDescriptor for SnowflakeProvider { fn metadata() -> ProviderMetadata { ProviderMetadata::new( SNOWFLAKE_PROVIDER_NAME, @@ -314,17 +290,6 @@ impl goose_providers::base::ProviderDescriptor for SnowflakeProvider { } } -impl ProviderDef for SnowflakeProvider { - type Provider = Self; - - fn from_env( - _extensions: Vec, - tls_config: Option, - ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(tls_config)) - } -} - #[async_trait] impl Provider for SnowflakeProvider { fn get_name(&self) -> &str { @@ -363,7 +328,7 @@ impl Provider for SnowflakeProvider { log.write(&response, Some(&usage))?; let provider_usage = ProviderUsage::new(response_model, usage); - Ok(super::base::stream_from_single_message( + Ok(crate::base::stream_from_single_message( message, provider_usage, )) diff --git a/crates/goose/src/providers/formats/mod.rs b/crates/goose/src/providers/formats/mod.rs index bec979c25..9f5531670 100644 --- a/crates/goose/src/providers/formats/mod.rs +++ b/crates/goose/src/providers/formats/mod.rs @@ -9,4 +9,6 @@ pub mod databricks { pub mod gcpvertexai; pub mod google; pub mod openrouter; -pub mod snowflake; +pub mod snowflake { + pub use goose_providers::formats::snowflake::*; +} diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index 41079ba83..c0d80d364 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -31,7 +31,7 @@ use super::{ openrouter::OpenRouterProvider, pi_acp::PiAcpProvider, provider_registry::ProviderRegistry, - snowflake::SnowflakeProvider, + snowflake_def::SnowflakeProviderDef, tetrate::TetrateProvider, xai::XaiProvider, xai_oauth::XaiOAuthProvider, @@ -130,7 +130,7 @@ async fn init_registry() -> RwLock { ); #[cfg(feature = "aws-providers")] registry.register::(false); - registry.register::(false); + registry.register::(false); registry.register::(true); registry.register::(false); registry.register_with_inventory::( diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 20b883c8b..c6ecf4561 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -73,7 +73,10 @@ mod retry { pub mod openai_def; #[cfg(feature = "aws-providers")] pub mod sagemaker_tgi; -pub mod snowflake; +pub mod snowflake { + pub use goose_providers::snowflake::*; +} +pub mod snowflake_def; pub mod testprovider; pub mod tetrate; pub mod toolshim; diff --git a/crates/goose/src/providers/snowflake_def.rs b/crates/goose/src/providers/snowflake_def.rs new file mode 100644 index 000000000..fb9894905 --- /dev/null +++ b/crates/goose/src/providers/snowflake_def.rs @@ -0,0 +1,53 @@ +use anyhow::Result; +use futures::future::BoxFuture; +use goose_providers::base::ProviderDescriptor; +use goose_providers::snowflake::SnowflakeProvider; + +use crate::config::{Config, ConfigError, ExtensionConfig}; +use crate::providers::api_client::TlsConfig; +use crate::providers::base::{ProviderDef, ProviderMetadata}; + +pub struct SnowflakeProviderDef; + +impl ProviderDescriptor for SnowflakeProviderDef { + fn metadata() -> ProviderMetadata { + SnowflakeProvider::metadata() + } +} + +impl ProviderDef for SnowflakeProviderDef { + type Provider = SnowflakeProvider; + + fn from_env( + _extensions: Vec, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Box::pin(from_env(tls_config)) + } +} + +pub async fn from_env(tls_config: Option) -> Result { + let config = Config::global(); + let host = get_config_or_secret(config, "SNOWFLAKE_HOST")?; + let token = get_config_or_secret(config, "SNOWFLAKE_TOKEN")?; + + SnowflakeProvider::new( + host, + token, + tls_config, + Some(crate::session_context::session_id_request_builder()), + ) +} + +fn get_config_or_secret(config: &Config, key: &str) -> Result { + config + .get_param(key) + .or_else(|_| config.get_secret(key)) + .map_err(|_| { + ConfigError::NotFound(format!( + "Did not find {} in either config file or keyring", + key + )) + .into() + }) +}