feat: ActionRequired (#5897)
This commit is contained in:
@@ -19,9 +19,9 @@ use goose::config::declarative_providers::{
|
||||
DeclarativeProviderConfig, LoadedProvider, ProviderEngine,
|
||||
};
|
||||
use goose::conversation::message::{
|
||||
FrontendToolRequest, Message, MessageContent, MessageMetadata, RedactedThinkingContent,
|
||||
SystemNotificationContent, SystemNotificationType, ThinkingContent, TokenState,
|
||||
ToolConfirmationRequest, ToolRequest, ToolResponse,
|
||||
ActionRequired, ActionRequiredData, FrontendToolRequest, Message, MessageContent,
|
||||
MessageMetadata, RedactedThinkingContent, SystemNotificationContent, SystemNotificationType,
|
||||
ThinkingContent, TokenState, ToolConfirmationRequest, ToolRequest, ToolResponse,
|
||||
};
|
||||
|
||||
use crate::routes::recipe_utils::RecipeManifest;
|
||||
@@ -358,7 +358,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::agent::agent_remove_extension,
|
||||
super::routes::agent::update_agent_provider,
|
||||
super::routes::agent::update_router_tool_selector,
|
||||
super::routes::reply::confirm_permission,
|
||||
super::routes::action_required::confirm_tool_action,
|
||||
super::routes::reply::reply,
|
||||
super::routes::session::list_sessions,
|
||||
super::routes::session::get_session,
|
||||
@@ -411,7 +411,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::config_management::UpdateCustomProviderRequest,
|
||||
super::routes::config_management::CheckProviderRequest,
|
||||
super::routes::config_management::SetProviderRequest,
|
||||
super::routes::reply::PermissionConfirmationRequest,
|
||||
super::routes::action_required::ConfirmToolActionRequest,
|
||||
super::routes::reply::ChatRequest,
|
||||
super::routes::session::ImportSessionRequest,
|
||||
super::routes::session::SessionListResponse,
|
||||
@@ -438,6 +438,8 @@ derive_utoipa!(Icon as IconSchema);
|
||||
ToolResponse,
|
||||
ToolRequest,
|
||||
ToolConfirmationRequest,
|
||||
ActionRequired,
|
||||
ActionRequiredData,
|
||||
ThinkingContent,
|
||||
RedactedThinkingContent,
|
||||
FrontendToolRequest,
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
use crate::state::AppState;
|
||||
use axum::{extract::State, http::StatusCode, routing::post, Json, Router};
|
||||
use goose::permission::permission_confirmation::PrincipalType;
|
||||
use goose::permission::{Permission, PermissionConfirmation};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
use utoipa::ToSchema;
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ConfirmToolActionRequest {
|
||||
id: String,
|
||||
#[serde(default = "default_principal_type")]
|
||||
principal_type: PrincipalType,
|
||||
action: String,
|
||||
session_id: String,
|
||||
}
|
||||
|
||||
fn default_principal_type() -> PrincipalType {
|
||||
PrincipalType::Tool
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/action-required/tool-confirmation",
|
||||
request_body = ConfirmToolActionRequest,
|
||||
responses(
|
||||
(status = 200, description = "Tool confirmation action is confirmed", body = Value),
|
||||
(status = 401, description = "Unauthorized - invalid secret key"),
|
||||
(status = 500, description = "Internal server error")
|
||||
)
|
||||
)]
|
||||
pub async fn confirm_tool_action(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(request): Json<ConfirmToolActionRequest>,
|
||||
) -> Result<Json<Value>, StatusCode> {
|
||||
let agent = state.get_agent_for_route(request.session_id).await?;
|
||||
let permission = match request.action.as_str() {
|
||||
"always_allow" => Permission::AlwaysAllow,
|
||||
"allow_once" => Permission::AllowOnce,
|
||||
"deny" => Permission::DenyOnce,
|
||||
_ => Permission::DenyOnce,
|
||||
};
|
||||
|
||||
agent
|
||||
.handle_confirmation(
|
||||
request.id.clone(),
|
||||
PermissionConfirmation {
|
||||
principal_type: request.principal_type,
|
||||
permission,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Json(Value::Object(serde_json::Map::new())))
|
||||
}
|
||||
|
||||
pub fn routes(state: Arc<AppState>) -> Router {
|
||||
Router::new()
|
||||
.route(
|
||||
"/action-required/tool-confirmation",
|
||||
post(confirm_tool_action),
|
||||
)
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
mod integration_tests {
|
||||
use super::*;
|
||||
use axum::{body::Body, http::Request};
|
||||
use tower::ServiceExt;
|
||||
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_tool_confirmation_endpoint() {
|
||||
let state = AppState::new().await.unwrap();
|
||||
|
||||
let app = routes(state);
|
||||
|
||||
let request = Request::builder()
|
||||
.uri("/action-required/tool-confirmation")
|
||||
.method("POST")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-secret-key", "test-secret")
|
||||
.body(Body::from(
|
||||
serde_json::to_string(&ConfirmToolActionRequest {
|
||||
id: "test-id".to_string(),
|
||||
principal_type: PrincipalType::Tool,
|
||||
action: "allow_once".to_string(),
|
||||
session_id: "test-session".to_string(),
|
||||
})
|
||||
.unwrap(),
|
||||
))
|
||||
.unwrap();
|
||||
|
||||
let response = app.oneshot(request).await.unwrap();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
pub mod action_required;
|
||||
pub mod agent;
|
||||
pub mod audio;
|
||||
pub mod config_management;
|
||||
@@ -22,6 +23,7 @@ pub fn configure(state: Arc<crate::state::AppState>, secret_key: String) -> Rout
|
||||
Router::new()
|
||||
.merge(status::routes())
|
||||
.merge(reply::routes(state.clone()))
|
||||
.merge(action_required::routes(state.clone()))
|
||||
.merge(agent::routes(state.clone()))
|
||||
.merge(audio::routes(state.clone()))
|
||||
.merge(config_management::routes(state.clone()))
|
||||
|
||||
@@ -8,17 +8,12 @@ use axum::{
|
||||
};
|
||||
use bytes::Bytes;
|
||||
use futures::{stream::StreamExt, Stream};
|
||||
use goose::agents::{AgentEvent, SessionConfig};
|
||||
use goose::conversation::message::{Message, MessageContent, TokenState};
|
||||
use goose::conversation::Conversation;
|
||||
use goose::permission::{Permission, PermissionConfirmation};
|
||||
use goose::session::SessionManager;
|
||||
use goose::{
|
||||
agents::{AgentEvent, SessionConfig},
|
||||
permission::permission_confirmation::PrincipalType,
|
||||
};
|
||||
use rmcp::model::ServerNotification;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::{
|
||||
convert::Infallible,
|
||||
pin::Pin,
|
||||
@@ -30,7 +25,6 @@ use tokio::sync::mpsc;
|
||||
use tokio::time::timeout;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use utoipa::ToSchema;
|
||||
|
||||
fn track_tool_telemetry(content: &MessageContent, all_messages: &[Message]) {
|
||||
match content {
|
||||
@@ -452,60 +446,12 @@ pub async fn reply(
|
||||
Ok(SseResponse::new(stream))
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||
pub struct PermissionConfirmationRequest {
|
||||
id: String,
|
||||
#[serde(default = "default_principal_type")]
|
||||
principal_type: PrincipalType,
|
||||
action: String,
|
||||
session_id: String,
|
||||
}
|
||||
|
||||
fn default_principal_type() -> PrincipalType {
|
||||
PrincipalType::Tool
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/confirm",
|
||||
request_body = PermissionConfirmationRequest,
|
||||
responses(
|
||||
(status = 200, description = "Permission action is confirmed", body = Value),
|
||||
(status = 401, description = "Unauthorized - invalid secret key"),
|
||||
(status = 500, description = "Internal server error")
|
||||
)
|
||||
)]
|
||||
pub async fn confirm_permission(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(request): Json<PermissionConfirmationRequest>,
|
||||
) -> Result<Json<Value>, StatusCode> {
|
||||
let agent = state.get_agent_for_route(request.session_id).await?;
|
||||
let permission = match request.action.as_str() {
|
||||
"always_allow" => Permission::AlwaysAllow,
|
||||
"allow_once" => Permission::AllowOnce,
|
||||
"deny" => Permission::DenyOnce,
|
||||
_ => Permission::DenyOnce,
|
||||
};
|
||||
|
||||
agent
|
||||
.handle_confirmation(
|
||||
request.id.clone(),
|
||||
PermissionConfirmation {
|
||||
principal_type: request.principal_type,
|
||||
permission,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
Ok(Json(Value::Object(serde_json::Map::new())))
|
||||
}
|
||||
|
||||
pub fn routes(state: Arc<AppState>) -> Router {
|
||||
Router::new()
|
||||
.route(
|
||||
"/reply",
|
||||
post(reply).layer(DefaultBodyLimit::max(50 * 1024 * 1024)),
|
||||
)
|
||||
.route("/confirm", post(confirm_permission))
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user