feat: support goose mode in UI (#1434)
Co-authored-by: Lily Delalande <ldelalande@squareup.com>
This commit is contained in:
@@ -86,7 +86,7 @@ async fn extend_prompt(
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
let mut agent = state.agent.lock().await;
|
||||
let mut agent = state.agent.write().await;
|
||||
if let Some(ref mut agent) = *agent {
|
||||
agent.extend_system_prompt(payload.extension).await;
|
||||
Ok(Json(ExtendPromptResponse { success: true }))
|
||||
@@ -134,7 +134,7 @@ async fn create_agent(
|
||||
|
||||
let new_agent = AgentFactory::create(&version, provider).expect("Failed to create agent");
|
||||
|
||||
let mut agent = state.agent.lock().await;
|
||||
let mut agent = state.agent.write().await;
|
||||
*agent = Some(new_agent);
|
||||
|
||||
Ok(Json(CreateAgentResponse { version }))
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
use crate::state::AppState;
|
||||
use axum::{extract::State, routing::delete, routing::post, Json, Router};
|
||||
use axum::{
|
||||
extract::{Query, State},
|
||||
routing::{delete, get, post},
|
||||
Json, Router,
|
||||
};
|
||||
use goose::config::Config;
|
||||
use http::{HeaderMap, StatusCode};
|
||||
use once_cell::sync::Lazy;
|
||||
@@ -140,6 +144,45 @@ async fn check_provider_configs(
|
||||
Ok(Json(response))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct GetConfigQuery {
|
||||
key: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct GetConfigResponse {
|
||||
value: Option<String>,
|
||||
}
|
||||
|
||||
pub async fn get_config(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
Query(query): Query<GetConfigQuery>,
|
||||
) -> Result<Json<GetConfigResponse>, StatusCode> {
|
||||
// Verify secret key
|
||||
let secret_key = headers
|
||||
.get("X-Secret-Key")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
if secret_key != state.secret_key {
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
// Fetch the configuration value. Right now we don't allow get a secret.
|
||||
let config = Config::global();
|
||||
let value = if let Ok(config_value) = config.get::<String>(&query.key) {
|
||||
Some(config_value)
|
||||
} else if let Ok(env_value) = std::env::var(&query.key) {
|
||||
Some(env_value)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Return the value
|
||||
Ok(Json(GetConfigResponse { value }))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct DeleteConfigRequest {
|
||||
@@ -178,6 +221,7 @@ async fn delete_config(
|
||||
pub fn routes(state: AppState) -> Router {
|
||||
Router::new()
|
||||
.route("/configs/providers", post(check_provider_configs))
|
||||
.route("/configs/get", get(get_config))
|
||||
.route("/configs/store", post(store_config))
|
||||
.route("/configs/delete", delete(delete_config))
|
||||
.with_state(state)
|
||||
|
||||
@@ -161,7 +161,7 @@ async fn add_extension(
|
||||
};
|
||||
|
||||
// Acquire a lock on the agent and attempt to add the extension.
|
||||
let mut agent = state.agent.lock().await;
|
||||
let mut agent = state.agent.write().await;
|
||||
let agent = agent.as_mut().ok_or(StatusCode::PRECONDITION_REQUIRED)?;
|
||||
let response = agent.add_extension(extension_config).await;
|
||||
|
||||
@@ -201,7 +201,7 @@ async fn remove_extension(
|
||||
}
|
||||
|
||||
// Acquire a lock on the agent and attempt to remove the extension
|
||||
let mut agent = state.agent.lock().await;
|
||||
let mut agent = state.agent.write().await;
|
||||
let agent = agent.as_mut().ok_or(StatusCode::PRECONDITION_REQUIRED)?;
|
||||
agent.remove_extension(&name).await;
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ use goose::message::{Message, MessageContent};
|
||||
|
||||
use mcp_core::role::Role;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::{
|
||||
convert::Infallible,
|
||||
pin::Pin,
|
||||
@@ -113,7 +114,7 @@ async fn handler(
|
||||
|
||||
// Spawn task to handle streaming
|
||||
tokio::spawn(async move {
|
||||
let agent = agent.lock().await;
|
||||
let agent = agent.read().await;
|
||||
let agent = match agent.as_ref() {
|
||||
Some(agent) => agent,
|
||||
None => {
|
||||
@@ -237,7 +238,7 @@ async fn ask_handler(
|
||||
}
|
||||
|
||||
let agent = state.agent.clone();
|
||||
let agent = agent.lock().await;
|
||||
let agent = agent.write().await;
|
||||
let agent = agent.as_ref().ok_or(StatusCode::NOT_FOUND)?;
|
||||
|
||||
// Create a single message for the prompt
|
||||
@@ -277,11 +278,42 @@ async fn ask_handler(
|
||||
}))
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ToolConfirmationRequest {
|
||||
id: String,
|
||||
confirmed: bool,
|
||||
}
|
||||
|
||||
async fn confirm_handler(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
Json(request): Json<ToolConfirmationRequest>,
|
||||
) -> Result<Json<Value>, StatusCode> {
|
||||
// Verify secret key
|
||||
let secret_key = headers
|
||||
.get("X-Secret-Key")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
if secret_key != state.secret_key {
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
let agent = state.agent.clone();
|
||||
let agent = agent.read().await;
|
||||
let agent = agent.as_ref().ok_or(StatusCode::NOT_FOUND)?;
|
||||
agent
|
||||
.handle_confirmation(request.id.clone(), request.confirmed)
|
||||
.await;
|
||||
Ok(Json(Value::Object(serde_json::Map::new())))
|
||||
}
|
||||
|
||||
// Configure routes for this module
|
||||
pub fn routes(state: AppState) -> Router {
|
||||
Router::new()
|
||||
.route("/reply", post(handler))
|
||||
.route("/ask", post(ask_handler))
|
||||
.route("/confirm", post(confirm_handler))
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
@@ -332,7 +364,7 @@ mod tests {
|
||||
use axum::{body::Body, http::Request};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use tower::ServiceExt;
|
||||
|
||||
// This test requires tokio runtime
|
||||
@@ -346,7 +378,7 @@ mod tests {
|
||||
let agent = AgentFactory::create("reference", mock_provider).unwrap();
|
||||
let state = AppState {
|
||||
config: Arc::new(Mutex::new(HashMap::new())),
|
||||
agent: Arc::new(Mutex::new(Some(agent))),
|
||||
agent: Arc::new(RwLock::new(Some(agent))),
|
||||
secret_key: "test-secret".to_string(),
|
||||
};
|
||||
|
||||
|
||||
@@ -3,13 +3,13 @@ use goose::agents::Agent;
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
|
||||
/// Shared application state
|
||||
#[allow(dead_code)]
|
||||
#[derive(Clone)]
|
||||
pub struct AppState {
|
||||
pub agent: Arc<Mutex<Option<Box<dyn Agent>>>>,
|
||||
pub agent: Arc<RwLock<Option<Box<dyn Agent>>>>,
|
||||
pub secret_key: String,
|
||||
pub config: Arc<Mutex<HashMap<String, Value>>>,
|
||||
}
|
||||
@@ -17,7 +17,7 @@ pub struct AppState {
|
||||
impl AppState {
|
||||
pub async fn new(secret_key: String) -> Result<Self> {
|
||||
Ok(Self {
|
||||
agent: Arc::new(Mutex::new(None)),
|
||||
agent: Arc::new(RwLock::new(None)),
|
||||
secret_key,
|
||||
config: Arc::new(Mutex::new(HashMap::new())),
|
||||
})
|
||||
|
||||
@@ -274,7 +274,8 @@ impl Agent for TruncateAgent {
|
||||
|
||||
// Wait for confirmation response through the channel
|
||||
let mut rx = self.confirmation_rx.lock().await;
|
||||
if let Some((req_id, confirmed)) = rx.recv().await {
|
||||
// Loop the recv until we have a matched req_id due to potential duplicate messages.
|
||||
while let Some((req_id, confirmed)) = rx.recv().await {
|
||||
if req_id == request.id {
|
||||
if confirmed {
|
||||
// User approved - dispatch the tool call
|
||||
@@ -290,6 +291,7 @@ impl Agent for TruncateAgent {
|
||||
Ok(vec![Content::text("User declined to run this tool.")]),
|
||||
);
|
||||
}
|
||||
break; // Exit the loop once the matching `req_id` is found
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user