From 8247fd6c1d25d219374b1e18c583d1d4a3c50cde Mon Sep 17 00:00:00 2001 From: Yingjie He Date: Mon, 21 Apr 2025 13:27:38 -0700 Subject: [PATCH] refactor: use the verify_secret_key util for all api handlers (#2284) --- crates/goose-server/src/routes/agent.rs | 30 ++----------- .../src/routes/config_management.rs | 19 ++------ crates/goose-server/src/routes/configs.rs | 31 ++----------- crates/goose-server/src/routes/extension.rs | 22 ++-------- crates/goose-server/src/routes/reply.rs | 43 +++---------------- crates/goose-server/src/routes/session.rs | 21 ++------- crates/goose-server/src/routes/utils.rs | 16 +++++++ 7 files changed, 40 insertions(+), 142 deletions(-) diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index d543e429..03696fcc 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -1,3 +1,4 @@ +use super::utils::verify_secret_key; use crate::state::AppState; use axum::{ extract::{Query, State}, @@ -85,15 +86,7 @@ async fn extend_prompt( headers: HeaderMap, Json(payload): Json, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let mut agent = state.agent.write().await; if let Some(ref mut agent) = *agent { @@ -110,15 +103,7 @@ async fn create_agent( headers: HeaderMap, Json(payload): Json, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; // Set the environment variable for the model if provided if let Some(model) = &payload.model { @@ -187,14 +172,7 @@ async fn get_tools( headers: HeaderMap, Query(query): Query, ) -> Result>, StatusCode> { - 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); - } + verify_secret_key(&headers, &state)?; let config = Config::global(); let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()); diff --git a/crates/goose-server/src/routes/config_management.rs b/crates/goose-server/src/routes/config_management.rs index 7a78bc39..87c953e7 100644 --- a/crates/goose-server/src/routes/config_management.rs +++ b/crates/goose-server/src/routes/config_management.rs @@ -1,3 +1,4 @@ +use super::utils::verify_secret_key; use crate::routes::utils::check_provider_configured; use crate::state::AppState; use axum::{ @@ -5,6 +6,7 @@ use axum::{ routing::{delete, get, post}, Json, Router, }; +use etcetera::{choose_app_strategy, AppStrategy, AppStrategyArgs}; use goose::config::Config; use goose::config::{extensions::name_to_key, PermissionManager}; use goose::config::{ExtensionConfigManager, ExtensionEntry}; @@ -12,26 +14,13 @@ use goose::providers::base::ProviderMetadata; use goose::providers::providers as get_providers; use goose::{agents::ExtensionConfig, config::permission::PermissionLevel}; use http::{HeaderMap, StatusCode}; +use once_cell::sync::Lazy; use serde::{Deserialize, Serialize}; use serde_json::Value; use serde_yaml; use std::collections::HashMap; use utoipa::ToSchema; -fn verify_secret_key(headers: &HeaderMap, state: &AppState) -> Result { - // 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 { - Err(StatusCode::UNAUTHORIZED) - } else { - Ok(StatusCode::OK) - } -} - #[derive(Serialize, ToSchema)] pub struct ExtensionResponse { pub extensions: Vec, @@ -431,8 +420,6 @@ pub async fn upsert_permissions( Ok(Json("Permissions updated successfully".to_string())) } -use etcetera::{choose_app_strategy, AppStrategy, AppStrategyArgs}; -use once_cell::sync::Lazy; pub static APP_STRATEGY: Lazy = Lazy::new(|| AppStrategyArgs { top_level_domain: "Block".to_string(), author: "Block".to_string(), diff --git a/crates/goose-server/src/routes/configs.rs b/crates/goose-server/src/routes/configs.rs index 858bbf69..ee918009 100644 --- a/crates/goose-server/src/routes/configs.rs +++ b/crates/goose-server/src/routes/configs.rs @@ -1,3 +1,4 @@ +use super::utils::verify_secret_key; use crate::state::AppState; use axum::{ extract::{Query, State}, @@ -29,15 +30,7 @@ async fn store_config( headers: HeaderMap, Json(request): Json, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let config = Config::global(); let result = if request.is_secret { @@ -159,15 +152,7 @@ pub async fn get_config( headers: HeaderMap, Query(query): Query, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; // Fetch the configuration value. Right now we don't allow get a secret. let config = Config::global(); @@ -193,15 +178,7 @@ async fn delete_config( headers: HeaderMap, Json(request): Json, ) -> Result { - // 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); - } + verify_secret_key(&headers, &state)?; // Attempt to delete the key let config = Config::global(); diff --git a/crates/goose-server/src/routes/extension.rs b/crates/goose-server/src/routes/extension.rs index 71489e6e..cc12e4e9 100644 --- a/crates/goose-server/src/routes/extension.rs +++ b/crates/goose-server/src/routes/extension.rs @@ -2,6 +2,7 @@ use std::env; use std::path::Path; use std::sync::OnceLock; +use super::utils::verify_secret_key; use crate::state::AppState; use axum::{extract::State, routing::post, Json, Router}; use goose::agents::{extension::Envs, ExtensionConfig}; @@ -82,6 +83,8 @@ async fn add_extension( headers: HeaderMap, raw: axum::extract::Json, ) -> Result, StatusCode> { + verify_secret_key(&headers, &state)?; + // Log the raw request for debugging tracing::info!( "Received extension request: {}", @@ -100,15 +103,6 @@ async fn add_extension( return Err(StatusCode::UNPROCESSABLE_ENTITY); } }; - // Verify the presence and validity of the 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); - } // If this is a Stdio extension that uses npx, check for Node.js installation #[cfg(target_os = "windows")] @@ -264,15 +258,7 @@ async fn remove_extension( headers: HeaderMap, Json(name): Json, ) -> Result, StatusCode> { - // Verify the presence and validity of the 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); - } + verify_secret_key(&headers, &state)?; // Acquire a lock on the agent and attempt to remove the extension let mut agent = state.agent.write().await; diff --git a/crates/goose-server/src/routes/reply.rs b/crates/goose-server/src/routes/reply.rs index 66c3dc8b..7781a4fb 100644 --- a/crates/goose-server/src/routes/reply.rs +++ b/crates/goose-server/src/routes/reply.rs @@ -1,3 +1,4 @@ +use super::utils::verify_secret_key; use crate::state::AppState; use axum::{ extract::State, @@ -104,15 +105,7 @@ async fn handler( headers: HeaderMap, Json(request): Json, ) -> Result { - // 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); - } + verify_secret_key(&headers, &state)?; // Create channel for streaming let (tx, rx) = mpsc::channel(100); @@ -273,15 +266,7 @@ async fn ask_handler( headers: HeaderMap, Json(request): Json, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let session_working_dir = request.session_working_dir; @@ -393,15 +378,7 @@ pub async fn confirm_permission( headers: HeaderMap, Json(request): Json, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let agent = state.agent.clone(); let agent = agent.read().await; @@ -437,6 +414,8 @@ async fn submit_tool_result( headers: HeaderMap, raw: axum::extract::Json, ) -> Result, StatusCode> { + verify_secret_key(&headers, &state)?; + // Log the raw request for debugging tracing::info!( "Received tool result request: {}", @@ -456,16 +435,6 @@ async fn submit_tool_result( } }; - // 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.read().await; let agent = agent.as_ref().ok_or(StatusCode::NOT_FOUND)?; agent.handle_tool_result(payload.id, payload.result).await; diff --git a/crates/goose-server/src/routes/session.rs b/crates/goose-server/src/routes/session.rs index e81c1f2a..02b232da 100644 --- a/crates/goose-server/src/routes/session.rs +++ b/crates/goose-server/src/routes/session.rs @@ -1,3 +1,4 @@ +use super::utils::verify_secret_key; use crate::state::AppState; use axum::{ extract::{Path, State}, @@ -27,15 +28,7 @@ async fn list_sessions( State(state): State, headers: HeaderMap, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let sessions = get_session_info().map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; @@ -48,15 +41,7 @@ async fn get_session_history( headers: HeaderMap, Path(session_id): Path, ) -> Result, 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); - } + verify_secret_key(&headers, &state)?; let session_path = session::get_path(session::Identifier::Name(session_id.clone())); diff --git a/crates/goose-server/src/routes/utils.rs b/crates/goose-server/src/routes/utils.rs index 06a1bc3e..a6659e5e 100644 --- a/crates/goose-server/src/routes/utils.rs +++ b/crates/goose-server/src/routes/utils.rs @@ -1,5 +1,7 @@ +use crate::state::AppState; use goose::config::Config; use goose::providers::base::{ConfigKey, ProviderMetadata}; +use http::{HeaderMap, StatusCode}; use serde::{Deserialize, Serialize}; use std::env; use std::error::Error; @@ -21,6 +23,20 @@ pub struct KeyInfo { pub value: Option, // Only populated for non-secret keys that are set } +pub fn verify_secret_key(headers: &HeaderMap, state: &AppState) -> Result { + // 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 { + Err(StatusCode::UNAUTHORIZED) + } else { + Ok(StatusCode::OK) + } +} + /// Inspects a configuration key to determine if it's set, its location, and value (for non-secret keys) #[allow(dead_code)] pub fn inspect_key(key_name: &str, is_secret: bool) -> Result> {