move snowflake (#10218)
This commit is contained in:
@@ -3,3 +3,5 @@ pub mod databricks;
|
||||
pub mod ollama;
|
||||
pub mod openai;
|
||||
pub mod openai_responses;
|
||||
|
||||
pub mod snowflake;
|
||||
|
||||
+5
-5
@@ -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";
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,3 +16,5 @@ pub mod openai;
|
||||
pub mod openai_compatible;
|
||||
|
||||
pub use declarative::declarative_providers::*;
|
||||
|
||||
pub mod snowflake;
|
||||
|
||||
@@ -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<crate::providers::api_client::TlsConfig>,
|
||||
pub fn new(
|
||||
mut host: String,
|
||||
token: String,
|
||||
tls_config: Option<TlsConfig>,
|
||||
request_builder: Option<RequestBuilderDecorator>,
|
||||
) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let mut host: Result<String, ConfigError> = 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<String, ConfigError> = 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<crate::config::ExtensionConfig>,
|
||||
tls_config: Option<crate::providers::api_client::TlsConfig>,
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
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,
|
||||
))
|
||||
@@ -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::*;
|
||||
}
|
||||
|
||||
@@ -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<ProviderRegistry> {
|
||||
);
|
||||
#[cfg(feature = "aws-providers")]
|
||||
registry.register::<SageMakerTgiProvider>(false);
|
||||
registry.register::<SnowflakeProvider>(false);
|
||||
registry.register::<SnowflakeProviderDef>(false);
|
||||
registry.register::<TetrateProvider>(true);
|
||||
registry.register::<XaiProvider>(false);
|
||||
registry.register_with_inventory::<XaiOAuthProvider>(
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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<ExtensionConfig>,
|
||||
tls_config: Option<TlsConfig>,
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
Box::pin(from_env(tls_config))
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn from_env(tls_config: Option<TlsConfig>) -> Result<SnowflakeProvider> {
|
||||
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<String> {
|
||||
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()
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user