Provider scenario tests (#3688)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -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>>>,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)),
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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") {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user