refactor: remove agent flavours, move provider to Agent (#2091)
This commit is contained in:
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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))),
|
||||
|
||||
@@ -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>>>,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user