refactor: remove agent flavours, move provider to Agent (#2091)

This commit is contained in:
Salman Mohammed
2025-04-09 15:02:47 -04:00
committed by GitHub
parent a8cbd81c61
commit 513d5c8f5a
30 changed files with 1297 additions and 2272 deletions
+8 -13
View File
@@ -5,11 +5,10 @@ use axum::{
routing::{get, post},
Json, Router,
};
use goose::{agents::AgentFactory, config::PermissionManager, model::ModelConfig, providers};
use goose::{
agents::{capabilities::get_parameter_names, extension::ToolInfo},
config::Config,
};
use goose::agents::{extension::ToolInfo, extension_manager::get_parameter_names};
use goose::config::Config;
use goose::config::PermissionManager;
use goose::{agents::Agent, model::ModelConfig, providers};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::env;
@@ -32,7 +31,6 @@ struct ExtendPromptResponse {
#[derive(Deserialize)]
struct CreateAgentRequest {
version: Option<String>,
provider: String,
model: Option<String>,
}
@@ -70,8 +68,8 @@ pub struct GetToolsQuery {
}
async fn get_versions() -> Json<VersionsResponse> {
let versions = AgentFactory::available_versions();
let default_version = AgentFactory::default_version().to_string();
let versions = ["goose".to_string()];
let default_version = "goose".to_string();
Json(VersionsResponse {
available_versions: versions.iter().map(|v| v.to_string()).collect(),
@@ -136,11 +134,8 @@ async fn create_agent(
let provider =
providers::create(&payload.provider, model_config).expect("Failed to create provider");
let version = payload
.version
.unwrap_or_else(|| AgentFactory::default_version().to_string());
let new_agent = AgentFactory::create(&version, provider).expect("Failed to create agent");
let version = String::from("goose");
let new_agent = Agent::new(provider);
let mut agent = state.agent.write().await;
*agent = Some(new_agent);
@@ -8,7 +8,7 @@ use axum::{
use goose::agents::ExtensionConfig;
use goose::config::extensions::name_to_key;
use goose::config::Config;
use goose::config::{ExtensionEntry, ExtensionManager};
use goose::config::{ExtensionConfigManager, ExtensionEntry};
use goose::providers::base::ProviderMetadata;
use goose::providers::providers as get_providers;
use http::{HeaderMap, StatusCode};
@@ -184,7 +184,7 @@ pub async fn get_extensions(
) -> Result<Json<ExtensionResponse>, StatusCode> {
verify_secret_key(&headers, &state)?;
match ExtensionManager::get_all() {
match ExtensionConfigManager::get_all() {
Ok(extensions) => Ok(Json(ExtensionResponse { extensions })),
Err(err) => {
// Return UNPROCESSABLE_ENTITY only for DeserializeError, INTERNAL_SERVER_ERROR for everything else
@@ -219,13 +219,13 @@ pub async fn add_extension(
verify_secret_key(&headers, &state)?;
// Get existing extensions to check if this is an update
let extensions = ExtensionManager::get_all().map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
let extensions =
ExtensionConfigManager::get_all().map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
let key = name_to_key(&extension_query.name);
let is_update = extensions.iter().any(|e| e.config.key() == key);
// Use ExtensionManager to set the extension
match ExtensionManager::set(ExtensionEntry {
match ExtensionConfigManager::set(ExtensionEntry {
enabled: extension_query.enabled,
config: extension_query.config,
}) {
@@ -257,8 +257,7 @@ pub async fn remove_extension(
verify_secret_key(&headers, &state)?;
let key = name_to_key(&name);
// Use ExtensionManager to remove the extension
match ExtensionManager::remove(&key) {
match ExtensionConfigManager::remove(&key) {
Ok(_) => Ok(Json(format!("Removed extension {}", name))),
Err(_) => Err(StatusCode::NOT_FOUND),
}
+5 -5
View File
@@ -153,7 +153,7 @@ async fn handler(
};
// Get the provider first, before starting the reply stream
let provider = agent.provider().await;
let provider = agent.provider();
let mut stream = match agent
.reply(
@@ -294,7 +294,7 @@ async fn ask_handler(
let agent = agent.as_ref().ok_or(StatusCode::NOT_FOUND)?;
// Get the provider first, before starting the reply stream
let provider = agent.provider().await;
let provider = agent.provider();
// Create a single message for the prompt
let messages = vec![Message::user().with_text(request.prompt)];
@@ -467,7 +467,7 @@ pub fn routes(state: AppState) -> Router {
mod tests {
use super::*;
use goose::{
agents::AgentFactory,
agents::Agent,
model::ModelConfig,
providers::{
base::{Provider, ProviderUsage, Usage},
@@ -518,10 +518,10 @@ mod tests {
async fn test_ask_endpoint() {
// Create a mock app state with mock provider
let mock_model_config = ModelConfig::new("test-model".to_string());
let mock_provider = Box::new(MockProvider {
let mock_provider = Arc::new(MockProvider {
model_config: mock_model_config,
});
let agent = AgentFactory::create("reference", mock_provider).unwrap();
let agent = Agent::new(mock_provider);
let state = AppState {
config: Arc::new(Mutex::new(HashMap::new())),
agent: Arc::new(RwLock::new(Some(agent))),
+1 -1
View File
@@ -9,7 +9,7 @@ use tokio::sync::{Mutex, RwLock};
#[allow(dead_code)]
#[derive(Clone)]
pub struct AppState {
pub agent: Arc<RwLock<Option<Box<dyn Agent>>>>,
pub agent: Arc<RwLock<Option<Agent>>>,
pub secret_key: String,
pub config: Arc<Mutex<HashMap<String, Value>>>,
}