chore: generalize extension request (#2213)
This commit is contained in:
@@ -17,6 +17,7 @@ use completion::GooseCompleter;
|
||||
use etcetera::choose_app_strategy;
|
||||
use etcetera::AppStrategy;
|
||||
use goose::agents::extension::{Envs, ExtensionConfig};
|
||||
use goose::agents::platform_tools::PLATFORM_ENABLE_EXTENSION_TOOL_NAME;
|
||||
use goose::agents::{Agent, SessionConfig};
|
||||
use goose::config::Config;
|
||||
use goose::message::{Message, MessageContent};
|
||||
@@ -620,13 +621,19 @@ impl Session {
|
||||
principal_type: PrincipalType::Tool,
|
||||
permission,
|
||||
},).await;
|
||||
} else if let Some(MessageContent::EnableExtensionRequest(enable_extension_request)) = message.content.first() {
|
||||
} else if let Some(MessageContent::ExtensionRequest(enable_extension_request)) = message.content.first() {
|
||||
output::hide_thinking();
|
||||
|
||||
let prompt = "Goose would like to install the following extension, do you approve?".to_string();
|
||||
let extension_action = if enable_extension_request.tool_name == PLATFORM_ENABLE_EXTENSION_TOOL_NAME {
|
||||
"enable"
|
||||
} else {
|
||||
"disable"
|
||||
};
|
||||
|
||||
let prompt = format!("Goose would like to {} the following extension, do you approve?", extension_action);
|
||||
let confirmed = cliclack::select(prompt)
|
||||
.item(true, "Yes, for this session", "Enable the extension for this session")
|
||||
.item(false, "No", "Do not enable the extension")
|
||||
.item(true, "Yes, for this session", format!("{} the extension for this session", extension_action))
|
||||
.item(false, "No", format!("Do not {} the extension", extension_action))
|
||||
.interact()?;
|
||||
let permission = if confirmed {
|
||||
Permission::AllowOnce
|
||||
|
||||
@@ -61,9 +61,10 @@ pub struct ToolConfirmationRequest {
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct EnableExtensionRequest {
|
||||
pub struct ExtensionRequest {
|
||||
pub id: String,
|
||||
pub extension_name: String,
|
||||
pub tool_name: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
@@ -94,7 +95,7 @@ pub enum MessageContent {
|
||||
ToolRequest(ToolRequest),
|
||||
ToolResponse(ToolResponse),
|
||||
ToolConfirmationRequest(ToolConfirmationRequest),
|
||||
EnableExtensionRequest(EnableExtensionRequest),
|
||||
ExtensionRequest(ExtensionRequest),
|
||||
FrontendToolRequest(FrontendToolRequest),
|
||||
Thinking(ThinkingContent),
|
||||
RedactedThinking(RedactedThinkingContent),
|
||||
@@ -144,10 +145,15 @@ impl MessageContent {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn enable_extension_request<S: Into<String>>(id: S, extension_name: String) -> Self {
|
||||
MessageContent::EnableExtensionRequest(EnableExtensionRequest {
|
||||
pub fn extension_request<S: Into<String>>(
|
||||
id: S,
|
||||
extension_name: String,
|
||||
tool_name: String,
|
||||
) -> Self {
|
||||
MessageContent::ExtensionRequest(ExtensionRequest {
|
||||
id: id.into(),
|
||||
extension_name,
|
||||
tool_name,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -192,9 +198,9 @@ impl MessageContent {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn as_enable_extension_request(&self) -> Option<&EnableExtensionRequest> {
|
||||
if let MessageContent::EnableExtensionRequest(ref enable_extension_request) = self {
|
||||
Some(enable_extension_request)
|
||||
pub fn as_extension_request(&self) -> Option<&ExtensionRequest> {
|
||||
if let MessageContent::ExtensionRequest(ref extension_request) = self {
|
||||
Some(extension_request)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
@@ -359,12 +365,17 @@ impl Message {
|
||||
))
|
||||
}
|
||||
|
||||
pub fn with_enable_extension_request<S: Into<String>>(
|
||||
pub fn with_extension_request<S: Into<String>>(
|
||||
self,
|
||||
id: S,
|
||||
extension_name: String,
|
||||
tool_name: String,
|
||||
) -> Self {
|
||||
self.with_content(MessageContent::enable_extension_request(id, extension_name))
|
||||
self.with_content(MessageContent::extension_request(
|
||||
id,
|
||||
extension_name,
|
||||
tool_name,
|
||||
))
|
||||
}
|
||||
|
||||
pub fn with_frontend_tool_request<S: Into<String>>(
|
||||
|
||||
@@ -160,7 +160,7 @@ pub fn get_confirmation_message(request_id: &str, tool_call: ToolCall) -> (Princ
|
||||
if tool_call.name == PLATFORM_ENABLE_EXTENSION_TOOL_NAME {
|
||||
(
|
||||
PrincipalType::Extension,
|
||||
Message::user().with_enable_extension_request(
|
||||
Message::user().with_extension_request(
|
||||
request_id,
|
||||
tool_call
|
||||
.arguments
|
||||
@@ -168,6 +168,7 @@ pub fn get_confirmation_message(request_id: &str, tool_call: ToolCall) -> (Princ
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
tool_call.name.clone(),
|
||||
),
|
||||
)
|
||||
} else {
|
||||
|
||||
@@ -60,8 +60,8 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||
MessageContent::ToolConfirmationRequest(_tool_confirmation_request) => {
|
||||
// Skip tool confirmation requests
|
||||
}
|
||||
MessageContent::EnableExtensionRequest(_enable_extension_request) => {
|
||||
// Skip enable extension requests
|
||||
MessageContent::ExtensionRequest(_extension_request) => {
|
||||
// Skip extension requests
|
||||
}
|
||||
MessageContent::Thinking(thinking) => {
|
||||
content.push(json!({
|
||||
|
||||
@@ -31,7 +31,7 @@ pub fn to_bedrock_message_content(content: &MessageContent) -> Result<bedrock::C
|
||||
MessageContent::ToolConfirmationRequest(_tool_confirmation_request) => {
|
||||
bedrock::ContentBlock::Text("".to_string())
|
||||
}
|
||||
MessageContent::EnableExtensionRequest(_enable_extension_request) => {
|
||||
MessageContent::ExtensionRequest(_extension_request) => {
|
||||
bedrock::ContentBlock::Text("".to_string())
|
||||
}
|
||||
MessageContent::Image(_) => {
|
||||
|
||||
@@ -179,7 +179,7 @@ pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<
|
||||
MessageContent::ToolConfirmationRequest(_) => {
|
||||
// Skip tool confirmation requests
|
||||
}
|
||||
MessageContent::EnableExtensionRequest(_) => {
|
||||
MessageContent::ExtensionRequest(_) => {
|
||||
// Skip enable extension requests
|
||||
}
|
||||
MessageContent::Image(image) => {
|
||||
|
||||
@@ -147,7 +147,7 @@ pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<
|
||||
MessageContent::ToolConfirmationRequest(_) => {
|
||||
// Skip tool confirmation requests
|
||||
}
|
||||
MessageContent::EnableExtensionRequest(_) => {
|
||||
MessageContent::ExtensionRequest(_) => {
|
||||
// Skip enable extension requests
|
||||
}
|
||||
MessageContent::Image(image) => {
|
||||
|
||||
Reference in New Issue
Block a user