feat: support goose mode in UI (#1434)

Co-authored-by: Lily Delalande <ldelalande@squareup.com>
This commit is contained in:
Yingjie He
2025-02-28 17:00:41 -08:00
committed by GitHub
parent 8f5fba97b8
commit f7f2540287
16 changed files with 320 additions and 19 deletions
+2 -2
View File
@@ -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 }))
+45 -1
View File
@@ -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)
+2 -2
View File
@@ -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;
+36 -4
View File
@@ -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 -3
View File
@@ -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())),
})
+3 -1
View File
@@ -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
}
}
}