Make async (#5126)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-10-10 19:31:50 -04:00
committed by GitHub
parent 69a7b7fe5a
commit 1015064367
43 changed files with 343 additions and 480 deletions
+15 -15
View File
@@ -1,5 +1,3 @@
// src/lib.rs or tests/truncate_agent_tests.rs
use std::sync::Arc;
use anyhow::Result;
@@ -71,19 +69,21 @@ impl ProviderType {
}
}
fn create_provider(&self, model_config: ModelConfig) -> Result<Arc<dyn Provider>> {
async fn create_provider(&self, model_config: ModelConfig) -> Result<Arc<dyn Provider>> {
Ok(match self {
ProviderType::Azure => Arc::new(AzureProvider::from_env(model_config)?),
ProviderType::OpenAi => Arc::new(OpenAiProvider::from_env(model_config)?),
ProviderType::Anthropic => Arc::new(AnthropicProvider::from_env(model_config)?),
ProviderType::Bedrock => Arc::new(BedrockProvider::from_env(model_config)?),
ProviderType::Databricks => Arc::new(DatabricksProvider::from_env(model_config)?),
ProviderType::GcpVertexAI => Arc::new(GcpVertexAIProvider::from_env(model_config)?),
ProviderType::Google => Arc::new(GoogleProvider::from_env(model_config)?),
ProviderType::Groq => Arc::new(GroqProvider::from_env(model_config)?),
ProviderType::Ollama => Arc::new(OllamaProvider::from_env(model_config)?),
ProviderType::OpenRouter => Arc::new(OpenRouterProvider::from_env(model_config)?),
ProviderType::Xai => Arc::new(XaiProvider::from_env(model_config)?),
ProviderType::Azure => Arc::new(AzureProvider::from_env(model_config).await?),
ProviderType::OpenAi => Arc::new(OpenAiProvider::from_env(model_config).await?),
ProviderType::Anthropic => Arc::new(AnthropicProvider::from_env(model_config).await?),
ProviderType::Bedrock => Arc::new(BedrockProvider::from_env(model_config).await?),
ProviderType::Databricks => Arc::new(DatabricksProvider::from_env(model_config).await?),
ProviderType::GcpVertexAI => {
Arc::new(GcpVertexAIProvider::from_env(model_config).await?)
}
ProviderType::Google => Arc::new(GoogleProvider::from_env(model_config).await?),
ProviderType::Groq => Arc::new(GroqProvider::from_env(model_config).await?),
ProviderType::Ollama => Arc::new(OllamaProvider::from_env(model_config).await?),
ProviderType::OpenRouter => Arc::new(OpenRouterProvider::from_env(model_config).await?),
ProviderType::Xai => Arc::new(XaiProvider::from_env(model_config).await?),
})
}
}
@@ -114,7 +114,7 @@ async fn run_truncate_test(
.unwrap()
.with_context_limit(Some(context_window))
.with_temperature(Some(0.0));
let provider = provider_type.create_provider(model_config)?;
let provider = provider_type.create_provider(model_config).await?;
let agent = Agent::new();
agent.update_provider(provider).await?;
+85 -135
View File
@@ -1,12 +1,21 @@
use anyhow::Result;
use dotenvy::dotenv;
use goose::conversation::message::{Message, MessageContent};
use goose::providers::anthropic::ANTHROPIC_DEFAULT_MODEL;
use goose::providers::azure::AZURE_DEFAULT_MODEL;
use goose::providers::base::Provider;
use goose::providers::bedrock::BEDROCK_DEFAULT_MODEL;
use goose::providers::create_with_named_model;
use goose::providers::databricks::DATABRICKS_DEFAULT_MODEL;
use goose::providers::errors::ProviderError;
use goose::providers::{
anthropic, azure, bedrock, databricks, google, groq, litellm, ollama, openai, openrouter,
snowflake, xai,
};
use goose::providers::google::GOOGLE_DEFAULT_MODEL;
use goose::providers::groq::GROQ_DEFAULT_MODEL;
use goose::providers::litellm::LITELLM_DEFAULT_MODEL;
use goose::providers::ollama::OLLAMA_DEFAULT_MODEL;
use goose::providers::openai::OPEN_AI_DEFAULT_MODEL;
use goose::providers::sagemaker_tgi::SAGEMAKER_TGI_DEFAULT_MODEL;
use goose::providers::snowflake::SNOWFLAKE_DEFAULT_MODEL;
use goose::providers::xai::XAI_DEFAULT_MODEL;
use rmcp::model::{AnnotateAble, Content, RawImageContent};
use rmcp::model::{CallToolRequestParam, Tool};
use rmcp::object;
@@ -77,18 +86,14 @@ lazy_static::lazy_static! {
static ref ENV_LOCK: Mutex<()> = Mutex::new(());
}
/// Generic test harness for any Provider implementation
struct ProviderTester {
provider: Arc<dyn Provider>,
name: String,
}
impl ProviderTester {
fn new<T: Provider + Send + Sync + 'static>(provider: T, name: String) -> Self {
Self {
provider: Arc::new(provider),
name,
}
fn new(provider: Arc<dyn Provider>, name: String) -> Self {
Self { provider, name }
}
async fn test_basic_response(&self) -> Result<()> {
@@ -99,14 +104,12 @@ impl ProviderTester {
.complete("You are a helpful assistant.", &[message], &[])
.await?;
// For a basic response, we expect a single text response
assert_eq!(
response.content.len(),
1,
"Expected single content item in response"
);
// Verify we got a text response
assert!(
matches!(response.content[0], MessageContent::Text(_)),
"Expected text response"
@@ -146,7 +149,6 @@ impl ProviderTester {
dbg!(&response1);
println!("===================");
// Verify we got a tool request
assert!(
response1
.content
@@ -177,7 +179,6 @@ impl ProviderTester {
)]),
);
// Verify we construct a valid payload including the request/response pair for the next inference
let (response2, _) = self
.provider
.complete(
@@ -203,7 +204,6 @@ impl ProviderTester {
}
async fn test_context_length_exceeded_error(&self) -> Result<()> {
// Google Gemini has a really long context window
let large_message_content = if self.name.to_lowercase() == "google" {
"hello ".repeat(1_300_000)
} else {
@@ -215,7 +215,6 @@ impl ProviderTester {
Message::assistant().with_text("hey! I think it's 4."),
Message::user().with_text(&large_message_content),
Message::assistant().with_text("heyy!!"),
// Messages before this mark should be truncated
Message::user().with_text("what's the meaning of life?"),
Message::assistant().with_text("the meaning of life is 42"),
Message::user().with_text(
@@ -223,18 +222,15 @@ impl ProviderTester {
),
];
// Test that we get ProviderError::ContextLengthExceeded when the context window is exceeded
let result = self
.provider
.complete("You are a helpful assistant.", &messages, &[])
.await;
// Print some debug info
println!("=== {}::context_length_exceeded_error ===", self.name);
dbg!(&result);
println!("===================");
// Ollama truncates by default even when the context window is exceeded
if self.name.to_lowercase() == "ollama" {
assert!(
result.is_ok(),
@@ -260,7 +256,6 @@ impl ProviderTester {
use goose::conversation::message::Message;
use std::fs;
// Try to read the test image
let image_path = "crates/goose/examples/test_assets/test_image.png";
let image_data = match fs::read(image_path) {
Ok(data) => data,
@@ -281,7 +276,6 @@ impl ProviderTester {
}
.no_annotation();
// Test 1: Direct image message
let message_with_image =
Message::user().with_image(image_content.data.clone(), image_content.mime_type.clone());
@@ -297,7 +291,6 @@ impl ProviderTester {
println!("=== {}::image_content_support ===", self.name);
let (response, _) = result?;
println!("Image response: {:?}", response);
// Verify we got a text response
assert!(
response
.content
@@ -307,7 +300,6 @@ impl ProviderTester {
);
println!("===================");
// Test 2: Tool response with image (this should be handled gracefully)
let screenshot_tool = Tool::new(
"get_screenshot",
"Get a screenshot of the current screen",
@@ -350,7 +342,6 @@ impl ProviderTester {
Ok(())
}
/// Run all provider tests
async fn run_test_suite(&self) -> Result<()> {
self.test_basic_response().await?;
self.test_tool_usage().await?;
@@ -366,74 +357,75 @@ fn load_env() {
}
}
/// Helper function to run a provider test with proper error handling and reporting
async fn test_provider<F, T>(
async fn test_provider(
name: &str,
model_name: &str,
required_vars: &[&str],
env_modifications: Option<HashMap<&str, Option<String>>>,
provider_fn: F,
) -> Result<()>
where
F: FnOnce() -> T,
T: Provider + Send + Sync + 'static,
{
// We start off as failed, so that if the process panics it is seen as a failure
) -> Result<()> {
TEST_REPORT.record_fail(name);
// Take exclusive access to environment modifications
let lock = ENV_LOCK.lock().unwrap();
let original_env = {
let _lock = ENV_LOCK.lock().unwrap();
load_env();
load_env();
// Save current environment state for required vars and modified vars
let mut original_env = HashMap::new();
for &var in required_vars {
if let Ok(val) = std::env::var(var) {
original_env.insert(var, val);
}
}
if let Some(mods) = &env_modifications {
for &var in mods.keys() {
let mut original_env = HashMap::new();
for &var in required_vars {
if let Ok(val) = std::env::var(var) {
original_env.insert(var, val);
}
}
}
if let Some(mods) = &env_modifications {
for &var in mods.keys() {
if let Ok(val) = std::env::var(var) {
original_env.insert(var, val);
}
}
}
// Apply any environment modifications
if let Some(mods) = &env_modifications {
for (&var, value) in mods.iter() {
match value {
Some(val) => std::env::set_var(var, val),
None => std::env::remove_var(var),
if let Some(mods) = &env_modifications {
for (&var, value) in mods.iter() {
match value {
Some(val) => std::env::set_var(var, val),
None => std::env::remove_var(var),
}
}
}
let missing_vars = required_vars.iter().any(|var| std::env::var(var).is_err());
if missing_vars {
println!("Skipping {} tests - credentials not configured", name);
TEST_REPORT.record_skip(name);
return Ok(());
}
original_env
};
let provider = match create_with_named_model(&name.to_lowercase(), model_name).await {
Ok(p) => p,
Err(e) => {
println!("Skipping {} tests - failed to create provider: {}", name, e);
TEST_REPORT.record_skip(name);
return Ok(());
}
};
{
let _lock = ENV_LOCK.lock().unwrap();
for (&var, value) in original_env.iter() {
std::env::set_var(var, value);
}
if let Some(mods) = env_modifications {
for &var in mods.keys() {
if !original_env.contains_key(var) {
std::env::remove_var(var);
}
}
}
}
// Setup the provider
let missing_vars = required_vars.iter().any(|var| std::env::var(var).is_err());
if missing_vars {
println!("Skipping {} tests - credentials not configured", name);
TEST_REPORT.record_skip(name);
return Ok(());
}
let provider = provider_fn();
// Restore original environment
for (&var, value) in original_env.iter() {
std::env::set_var(var, value);
}
if let Some(mods) = env_modifications {
for &var in mods.keys() {
if !original_env.contains_key(var) {
std::env::remove_var(var);
}
}
}
std::mem::drop(lock);
let tester = ProviderTester::new(provider, name.to_string());
match tester.run_test_suite().await {
Ok(_) => {
@@ -450,26 +442,20 @@ where
#[tokio::test]
async fn test_openai_provider() -> Result<()> {
test_provider(
"OpenAI",
&["OPENAI_API_KEY"],
None,
openai::OpenAiProvider::default,
)
.await
test_provider("openai", OPEN_AI_DEFAULT_MODEL, &["OPENAI_API_KEY"], None).await
}
#[tokio::test]
async fn test_azure_provider() -> Result<()> {
test_provider(
"Azure",
AZURE_DEFAULT_MODEL,
&[
"AZURE_OPENAI_API_KEY",
"AZURE_OPENAI_ENDPOINT",
"AZURE_OPENAI_DEPLOYMENT_NAME",
],
None,
azure::AzureProvider::default,
)
.await
}
@@ -478,26 +464,23 @@ async fn test_azure_provider() -> Result<()> {
async fn test_bedrock_provider_long_term_credentials() -> Result<()> {
test_provider(
"Bedrock",
BEDROCK_DEFAULT_MODEL,
&["AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY"],
None,
bedrock::BedrockProvider::default,
)
.await
}
#[tokio::test]
async fn test_bedrock_provider_aws_profile_credentials() -> Result<()> {
let env_mods = HashMap::from_iter([
// Ensure to unset long-term credentials to use AWS Profile provider
("AWS_ACCESS_KEY_ID", None),
("AWS_SECRET_ACCESS_KEY", None),
]);
let env_mods =
HashMap::from_iter([("AWS_ACCESS_KEY_ID", None), ("AWS_SECRET_ACCESS_KEY", None)]);
test_provider(
"Bedrock AWS Profile Credentials",
"Bedrock",
BEDROCK_DEFAULT_MODEL,
&["AWS_PROFILE"],
Some(env_mods),
bedrock::BedrockProvider::default,
)
.await
}
@@ -506,50 +489,30 @@ async fn test_bedrock_provider_aws_profile_credentials() -> Result<()> {
async fn test_databricks_provider() -> Result<()> {
test_provider(
"Databricks",
DATABRICKS_DEFAULT_MODEL,
&["DATABRICKS_HOST", "DATABRICKS_TOKEN"],
None,
databricks::DatabricksProvider::default,
)
.await
}
#[tokio::test]
async fn test_databricks_provider_oauth() -> Result<()> {
let mut env_mods = HashMap::new();
env_mods.insert("DATABRICKS_TOKEN", None);
test_provider(
"Databricks OAuth",
&["DATABRICKS_HOST"],
Some(env_mods),
databricks::DatabricksProvider::default,
)
.await
}
#[tokio::test]
async fn test_ollama_provider() -> Result<()> {
test_provider(
"Ollama",
&["OLLAMA_HOST"],
None,
ollama::OllamaProvider::default,
)
.await
test_provider("Ollama", OLLAMA_DEFAULT_MODEL, &["OLLAMA_HOST"], None).await
}
#[tokio::test]
async fn test_groq_provider() -> Result<()> {
test_provider("Groq", &["GROQ_API_KEY"], None, groq::GroqProvider::default).await
test_provider("Groq", GROQ_DEFAULT_MODEL, &["GROQ_API_KEY"], None).await
}
#[tokio::test]
async fn test_anthropic_provider() -> Result<()> {
test_provider(
"Anthropic",
ANTHROPIC_DEFAULT_MODEL,
&["ANTHROPIC_API_KEY"],
None,
anthropic::AnthropicProvider::default,
)
.await
}
@@ -558,31 +521,25 @@ async fn test_anthropic_provider() -> Result<()> {
async fn test_openrouter_provider() -> Result<()> {
test_provider(
"OpenRouter",
OPEN_AI_DEFAULT_MODEL,
&["OPENROUTER_API_KEY"],
None,
openrouter::OpenRouterProvider::default,
)
.await
}
#[tokio::test]
async fn test_google_provider() -> Result<()> {
test_provider(
"Google",
&["GOOGLE_API_KEY"],
None,
google::GoogleProvider::default,
)
.await
test_provider("Google", GOOGLE_DEFAULT_MODEL, &["GOOGLE_API_KEY"], None).await
}
#[tokio::test]
async fn test_snowflake_provider() -> Result<()> {
test_provider(
"Snowflake",
SNOWFLAKE_DEFAULT_MODEL,
&["SNOWFLAKE_HOST", "SNOWFLAKE_TOKEN"],
None,
snowflake::SnowflakeProvider::default,
)
.await
}
@@ -591,9 +548,9 @@ async fn test_snowflake_provider() -> Result<()> {
async fn test_sagemaker_tgi_provider() -> Result<()> {
test_provider(
"SageMakerTgi",
SAGEMAKER_TGI_DEFAULT_MODEL,
&["SAGEMAKER_ENDPOINT_NAME"],
None,
goose::providers::sagemaker_tgi::SageMakerTgiProvider::default,
)
.await
}
@@ -611,21 +568,14 @@ async fn test_litellm_provider() -> Result<()> {
("LITELLM_API_KEY", Some("".to_string())),
]);
test_provider(
"LiteLLM",
&[], // No required environment variables
Some(env_mods),
litellm::LiteLLMProvider::default,
)
.await
test_provider("LiteLLM", LITELLM_DEFAULT_MODEL, &[], Some(env_mods)).await
}
#[tokio::test]
async fn test_xai_provider() -> Result<()> {
test_provider("Xai", &["XAI_API_KEY"], None, xai::XaiProvider::default).await
test_provider("Xai", XAI_DEFAULT_MODEL, &["XAI_API_KEY"], None).await
}
// Print the final test report
#[ctor::dtor]
fn print_test_report() {
TEST_REPORT.print_summary();
+8 -8
View File
@@ -13,17 +13,17 @@ use serial_test::serial;
mod tetrate_streaming_tests {
use super::*;
fn create_test_provider() -> Result<TetrateProvider> {
async fn create_test_provider() -> Result<TetrateProvider> {
// Create a test provider with the default model
let model_config = ModelConfig::new("claude-3-5-sonnet-latest")?;
TetrateProvider::from_env(model_config)
TetrateProvider::from_env(model_config).await
}
#[tokio::test]
#[serial]
#[ignore] // Ignore by default, run with --ignored flag when API key is available
async fn test_tetrate_streaming_basic() -> Result<()> {
let provider = create_test_provider()?;
let provider = create_test_provider().await?;
let messages = vec![Message::user().with_text("Count from 1 to 5, one number at a time.")];
@@ -78,7 +78,7 @@ mod tetrate_streaming_tests {
#[serial]
#[ignore]
async fn test_tetrate_streaming_with_tools() -> Result<()> {
let provider = create_test_provider()?;
let provider = create_test_provider().await?;
// Define a simple tool
let weather_tool = Tool::new(
@@ -140,7 +140,7 @@ mod tetrate_streaming_tests {
#[serial]
#[ignore]
async fn test_tetrate_streaming_empty_response() -> Result<()> {
let provider = create_test_provider()?;
let provider = create_test_provider().await?;
// This might result in a very short or empty response
let messages = vec![Message::user().with_text("")];
@@ -169,7 +169,7 @@ mod tetrate_streaming_tests {
#[serial]
#[ignore]
async fn test_tetrate_streaming_long_response() -> Result<()> {
let provider = create_test_provider()?;
let provider = create_test_provider().await?;
let messages = vec![Message::user().with_text(
"Write a detailed 3-paragraph essay about the importance of streaming in modern APIs.",
@@ -230,7 +230,7 @@ mod tetrate_streaming_tests {
std::env::set_var("TETRATE_API_KEY", "invalid-key-for-testing");
let model_config = ModelConfig::new("claude-3-5-sonnet-latest")?;
let provider = TetrateProvider::from_env(model_config)?;
let provider = TetrateProvider::from_env(model_config).await?;
let messages = vec![Message::user().with_text("Hello")];
@@ -251,7 +251,7 @@ mod tetrate_streaming_tests {
#[serial]
#[ignore]
async fn test_tetrate_streaming_concurrent_streams() -> Result<()> {
let provider = create_test_provider()?;
let provider = create_test_provider().await?;
// Create multiple concurrent streams
let messages1 = vec![Message::user().with_text("Say 'Stream 1'")];