Provider scenario tests (#3688)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-07-28 20:04:11 +02:00
committed by GitHub
parent 4442ccfe4a
commit 73a274d311
33 changed files with 4424 additions and 220 deletions
+1 -1
View File
@@ -63,7 +63,7 @@ const DEFAULT_MAX_TURNS: u32 = 1000;
/// The main goose Agent
pub struct Agent {
pub(super) provider: Mutex<Option<Arc<dyn Provider>>>,
pub(super) extension_manager: Arc<RwLock<ExtensionManager>>,
pub extension_manager: Arc<RwLock<ExtensionManager>>,
pub(super) sub_recipe_manager: Mutex<SubRecipeManager>,
pub(super) tasks_manager: TasksManager,
pub(super) final_output_tool: Arc<Mutex<Option<FinalOutputTool>>>,
+7 -3
View File
@@ -338,12 +338,16 @@ impl ExtensionManager {
.insert(sanitized_name.clone());
}
self.clients
.insert(sanitized_name.clone(), Arc::new(Mutex::new(client)));
self.add_client(sanitized_name, client);
Ok(())
}
pub fn add_client(&mut self, client_name: String, client: Box<dyn McpClientTrait>) {
let sanitized_name = normalize(client_name);
self.clients
.insert(sanitized_name, Arc::new(Mutex::new(client)));
}
/// Get extensions info
pub async fn get_extensions_info(&self) -> Vec<ExtensionInfo> {
self.clients
+1 -3
View File
@@ -34,9 +34,7 @@ async fn toolshim_postprocess(
impl Agent {
/// Prepares tools and system prompt for a provider request
pub(crate) async fn prepare_tools_and_prompt(
&self,
) -> anyhow::Result<(Vec<Tool>, Vec<Tool>, String)> {
pub async fn prepare_tools_and_prompt(&self) -> anyhow::Result<(Vec<Tool>, Vec<Tool>, String)> {
// Get tool selection strategy from config
let config = Config::global();
let router_tool_selection_strategy = config
+6 -8
View File
@@ -71,8 +71,6 @@ pub fn create(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
return create_lead_worker_from_env(name, &model, &lead_model_name);
}
// Default: create regular provider
create_provider(name, model)
}
@@ -152,23 +150,23 @@ fn create_lead_worker_from_env(
fn create_provider(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
// We use Arc instead of Box to be able to clone for multiple async tasks
match name {
"openai" => Ok(Arc::new(OpenAiProvider::from_env(model)?)),
"anthropic" => Ok(Arc::new(AnthropicProvider::from_env(model)?)),
"azure_openai" => Ok(Arc::new(AzureProvider::from_env(model)?)),
"aws_bedrock" => Ok(Arc::new(BedrockProvider::from_env(model)?)),
"azure_openai" => Ok(Arc::new(AzureProvider::from_env(model)?)),
"claude-code" => Ok(Arc::new(ClaudeCodeProvider::from_env(model)?)),
"databricks" => Ok(Arc::new(DatabricksProvider::from_env(model)?)),
"gcp_vertex_ai" => Ok(Arc::new(GcpVertexAIProvider::from_env(model)?)),
"gemini-cli" => Ok(Arc::new(GeminiCliProvider::from_env(model)?)),
// "github_copilot" => Ok(Arc::new(GithubCopilotProvider::from_env(model)?)),
"google" => Ok(Arc::new(GoogleProvider::from_env(model)?)),
"groq" => Ok(Arc::new(GroqProvider::from_env(model)?)),
"litellm" => Ok(Arc::new(LiteLLMProvider::from_env(model)?)),
"ollama" => Ok(Arc::new(OllamaProvider::from_env(model)?)),
"openai" => Ok(Arc::new(OpenAiProvider::from_env(model)?)),
"openrouter" => Ok(Arc::new(OpenRouterProvider::from_env(model)?)),
"gcp_vertex_ai" => Ok(Arc::new(GcpVertexAIProvider::from_env(model)?)),
"google" => Ok(Arc::new(GoogleProvider::from_env(model)?)),
"sagemaker_tgi" => Ok(Arc::new(SageMakerTgiProvider::from_env(model)?)),
"venice" => Ok(Arc::new(VeniceProvider::from_env(model)?)),
"snowflake" => Ok(Arc::new(SnowflakeProvider::from_env(model)?)),
// "github_copilot" => Ok(Arc::new(GithubCopilotProvider::from_env(model)?)),
"venice" => Ok(Arc::new(VeniceProvider::from_env(model)?)),
"xai" => Ok(Arc::new(XaiProvider::from_env(model)?)),
_ => Err(anyhow::anyhow!("Unknown provider: {}", name)),
}
+10 -10
View File
@@ -70,28 +70,28 @@ impl GroqProvider {
.await?;
let status = response.status();
let payload: Option<Value> = response.json().await.ok();
let response_payload: Option<Value> = response.json().await.ok();
let formatted_payload = format!("{:?}", response_payload);
match status {
StatusCode::OK => payload.ok_or_else( || ProviderError::RequestFailed("Response body is not valid JSON".to_string()) ),
StatusCode::OK => response_payload.ok_or_else( || ProviderError::RequestFailed("Response body is not valid JSON".to_string()) ),
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
Err(ProviderError::Authentication(format!("Authentication failed. Please ensure your API keys are valid and have the required permissions. \
Status: {}. Response: {:?}", status, payload)))
Status: {}. Response: {:?}", status, response_payload)))
}
StatusCode::PAYLOAD_TOO_LARGE => {
Err(ProviderError::ContextLengthExceeded(format!("{:?}", payload)))
Err(ProviderError::ContextLengthExceeded(formatted_payload))
}
StatusCode::TOO_MANY_REQUESTS => {
Err(ProviderError::RateLimitExceeded(format!("{:?}", payload)))
Err(ProviderError::RateLimitExceeded(formatted_payload))
}
StatusCode::INTERNAL_SERVER_ERROR | StatusCode::SERVICE_UNAVAILABLE => {
Err(ProviderError::ServerError(format!("{:?}", payload)))
Err(ProviderError::ServerError(formatted_payload))
}
_ => {
tracing::debug!(
"{}", format!("Provider request failed with status: {}. Payload: {:?}", status, payload)
);
Err(ProviderError::RequestFailed(format!("Request failed with status: {}", status)))
let error_msg = format!("Provider request failed with status: {}. Payload: {:?}", status, response_payload);
tracing::debug!(error_msg);
Err(ProviderError::RequestFailed(error_msg))
}
}
}
+7
View File
@@ -93,6 +93,13 @@ impl OpenRouterProvider {
.await
.map_err(|e| ProviderError::RequestFailed(format!("Failed to parse response: {e}")))?;
let _debug = format!(
"OpenRouter request with payload: {} and response: {}",
serde_json::to_string_pretty(payload).unwrap_or_else(|_| "Invalid JSON".to_string()),
serde_json::to_string_pretty(&response_body)
.unwrap_or_else(|_| "Invalid JSON".to_string())
);
// OpenRouter can return errors in 200 OK responses, so we have to check for errors explicitly
// https://openrouter.ai/docs/api-reference/errors
if let Some(error_obj) = response_body.get("error") {
+9 -17
View File
@@ -58,6 +58,13 @@ impl TestProvider {
})
}
pub fn finish_recording(self) -> Result<()> {
if self.inner.is_some() {
self.save_records()?;
}
Ok(())
}
fn hash_input(messages: &[Message]) -> String {
let stable_messages: Vec<_> = messages
.iter()
@@ -68,6 +75,7 @@ impl TestProvider {
hasher.update(serialized.as_bytes());
format!("{:x}", hasher.finalize())
}
fn load_records(file_path: &str) -> Result<HashMap<String, TestRecord>> {
if !Path::new(file_path).exists() {
return Ok(HashMap::new());
@@ -113,7 +121,6 @@ impl Provider for TestProvider {
let hash = Self::hash_input(messages);
if let Some(inner) = &self.inner {
// Recording mode
let (message, usage) = inner.complete(system, messages, tools).await?;
let record = TestRecord {
@@ -135,7 +142,6 @@ impl Provider for TestProvider {
Ok((message, usage))
} else {
// Replay mode
let records = self.records.lock().unwrap();
if let Some(record) = records.get(&hash) {
Ok((record.output.message.clone(), record.output.usage.clone()))
@@ -153,16 +159,6 @@ impl Provider for TestProvider {
}
}
impl Drop for TestProvider {
fn drop(&mut self) {
if self.inner.is_some() {
if let Err(e) = self.save_records() {
eprintln!("Failed to save test records: {}", e);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -231,7 +227,6 @@ mod tests {
response: "Hello, world!".to_string(),
});
// Record phase
{
let test_provider = TestProvider::new_recording(mock, &temp_file);
@@ -244,11 +239,10 @@ mod tests {
assert_eq!(content.text, "Hello, world!");
}
test_provider.save_records().unwrap();
assert_eq!(test_provider.get_record_count(), 1);
test_provider.finish_recording().unwrap();
}
// Replay phase
{
let replay_provider = TestProvider::new_replaying(&temp_file).unwrap();
@@ -262,7 +256,6 @@ mod tests {
}
}
// Cleanup
let _ = fs::remove_file(temp_file);
}
@@ -286,7 +279,6 @@ mod tests {
.to_string()
.contains("No recorded response found"));
// Cleanup
let _ = fs::remove_file(temp_file);
}
}