+15
-15
@@ -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
@@ -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();
|
||||
|
||||
@@ -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'")];
|
||||
|
||||
Reference in New Issue
Block a user