feat: use the same permission flow for enable extensions (#2302)

This commit is contained in:
Yingjie He
2025-04-23 08:54:04 -07:00
committed by GitHub
parent 08682507d9
commit cc755100f0
16 changed files with 77 additions and 420 deletions
+1 -38
View File
@@ -28,7 +28,6 @@ import {
ToolRequestMessageContent,
ToolResponseMessageContent,
ToolConfirmationRequestMessageContent,
ExtensionRequestMessageContent,
} from '../types/message';
export interface ChatType {
@@ -48,9 +47,6 @@ const isUserMessage = (message: Message): boolean => {
if (message.content.every((c) => c.type === 'toolConfirmationRequest')) {
return false;
}
if (message.content.every((c) => c.type === 'extensionRequest')) {
return false;
}
return true;
};
@@ -260,14 +256,6 @@ export default function ChatView({
return [content.id, toolCall];
}
});
const extensionRequests = lastMessage.content
.filter(
(content): content is ExtensionRequestMessageContent =>
content.type === 'extensionRequest'
)
.map((content) => {
return [content.id, content.extensionCall];
});
if (toolRequests.length !== 0) {
// This means we were interrupted during a tool request
@@ -297,30 +285,6 @@ export default function ChatView({
// Use an immutable update to add the response message to the messages array
setMessages([...messages, responseMessage]);
}
// do the same for enable extension requests
// leverages toolResponse to send the error notification
if (extensionRequests.length !== 0) {
let responseMessage: Message = {
role: 'user',
created: Date.now(),
content: [],
};
const notification = 'Interrupted by the user to make a correction';
// generate a response saying it was interrupted for each extension request
for (const [reqId, _] of extensionRequests) {
const toolResponse: ToolResponseMessageContent = {
type: 'toolResponse',
id: reqId,
toolResult: {
status: 'error',
error: notification,
},
};
responseMessage.content.push(toolResponse);
}
setMessages([...messages, responseMessage]);
}
}
};
@@ -338,9 +302,8 @@ export default function ChatView({
(c) => c.type === 'toolConfirmationRequest'
);
const hasExtensionRequest = message.content.every((c) => c.type === 'extensionRequest');
// Keep the message if it has text content or tool confirmation or is not just tool responses
return hasTextContent || !hasOnlyToolResponses || hasToolConfirmation || hasExtensionRequest;
return hasTextContent || !hasOnlyToolResponses || hasToolConfirmation;
}
return true;
@@ -1,109 +0,0 @@
import React, { useState } from 'react';
import { snakeToTitleCase } from '../utils';
import { confirmPermission } from '../api';
interface ExtensionConfirmationProps {
isCancelledMessage: boolean;
isClicked: boolean;
extensionConfirmationId: string;
extensionName: string;
toolName: string;
}
export default function ExtensionConfirmation({
isCancelledMessage,
isClicked,
extensionConfirmationId,
extensionName,
toolName,
}: ExtensionConfirmationProps) {
const [clicked, setClicked] = useState(isClicked);
const [status, setStatus] = useState('unknown');
const extensionAction = toolName.toLowerCase().includes('enable') ? 'enable' : 'disable';
const handleButtonClick = async (confirmed: boolean) => {
setClicked(true);
setStatus(confirmed ? 'approved' : 'denied');
try {
const response = await confirmPermission({
body: {
id: extensionConfirmationId,
action: confirmed ? 'allow_once' : 'deny',
principal_type: 'Extension',
},
});
if (response.error) {
console.error('Failed to confirm permission: ', response.error);
}
} catch (err) {
console.error('Error fetching tools:', err);
}
};
return isCancelledMessage ? (
<div className="goose-message-content bg-bgSubtle rounded-2xl px-4 py-2 text-textStandard">
Extension {extensionAction} is cancelled.
</div>
) : (
<>
<div className="goose-message-content bg-bgSubtle rounded-2xl px-4 py-2 rounded-b-none text-textStandard">
Goose would like to {extensionAction} the following extension. Allow?
</div>
{clicked ? (
<div className="goose-message-tool bg-bgApp border border-borderSubtle dark:border-gray-700 rounded-b-2xl px-4 pt-4 pb-2 flex gap-4 mt-1">
<div className="flex items-center">
{status === 'approved' && (
<svg
className="w-5 h-5 text-gray-500"
xmlns="http://www.w3.org/2000/svg"
fill="none"
viewBox="0 0 24 24"
stroke="currentColor"
strokeWidth={2}
>
<path strokeLinecap="round" strokeLinejoin="round" d="M5 13l4 4L19 7" />
</svg>
)}
{status === 'denied' && (
<svg
className="w-5 h-5 text-gray-500"
xmlns="http://www.w3.org/2000/svg"
fill="none"
viewBox="0 0 24 24"
stroke="currentColor"
strokeWidth={2}
>
<path strokeLinecap="round" strokeLinejoin="round" d="M6 18L18 6M6 6l12 12" />
</svg>
)}
<span className="ml-2 text-textStandard">
{isClicked
? 'Extension enablement is not available'
: `${snakeToTitleCase(extensionName.includes('__') ? extensionName.split('__').pop() || extensionName : extensionName)} is ${status}`}{' '}
</span>
</div>
</div>
) : (
<div className="goose-message-tool bg-bgApp border border-borderSubtle dark:border-gray-700 rounded-b-2xl px-4 pt-4 pb-2 flex gap-4 mt-1">
<button
className={
'bg-black text-white dark:bg-white dark:text-black rounded-full px-6 py-2 transition'
}
onClick={() => handleButtonClick(true)}
>
{extensionAction.charAt(0).toUpperCase() + extensionAction.slice(1).toLowerCase()}{' '}
extension
</button>
<button
className={
'bg-white text-black dark:bg-black dark:text-white border border-gray-300 dark:border-gray-700 rounded-full px-6 py-2 transition'
}
onClick={() => handleButtonClick(false)}
>
Deny
</button>
</div>
)}
</>
);
}
@@ -12,11 +12,9 @@ import {
getToolResponses,
getToolConfirmationContent,
createToolErrorResponseMessage,
getExtensionContent,
} from '../types/message';
import ToolCallConfirmation from './ToolCallConfirmation';
import MessageCopyLink from './MessageCopyLink';
import ExtensionConfirmation from './ExtensionConfirmation';
interface GooseMessageProps {
messageHistoryIndex: number;
@@ -58,9 +56,6 @@ export default function GooseMessage({
const toolConfirmationContent = getToolConfirmationContent(message);
const hasToolConfirmation = toolConfirmationContent !== undefined;
const extensionContent = getExtensionContent(message);
const hasExtensionRequest = extensionContent !== undefined;
// Find tool responses that correspond to the tool requests in this message
const toolResponsesMap = useMemo(() => {
const responseMap = new Map();
@@ -95,23 +90,12 @@ export default function GooseMessage({
createToolErrorResponseMessage(toolConfirmationContent.id, 'The tool call is cancelled.')
);
}
if (messageIndex == messageHistoryIndex - 1 && hasExtensionRequest) {
appendMessage(
createToolErrorResponseMessage(
extensionContent.id,
'The extension enablement is cancelled.'
)
);
}
}, [
messageIndex,
messageHistoryIndex,
hasToolConfirmation,
toolConfirmationContent,
appendMessage,
hasExtensionRequest,
// Only include enableExtensionContent if it exists
extensionContent?.id,
]);
return (
@@ -172,16 +156,6 @@ export default function GooseMessage({
toolName={toolConfirmationContent.toolName}
/>
)}
{hasExtensionRequest && (
<ExtensionConfirmation
isCancelledMessage={messageIndex == messageHistoryIndex - 1}
isClicked={messageIndex < messageHistoryIndex - 1}
extensionConfirmationId={extensionContent.id}
extensionName={extensionContent.extensionName}
toolName={extensionContent.toolName}
/>
)}
</div>
{/* TODO(alexhancock): Re-enable link previews once styled well again */}
@@ -34,9 +34,7 @@ export default function PermissionModal({ extensionName, onClose }: PermissionMo
} else {
const filteredTools = (response.data || []).filter(
(tool) =>
tool.name !== 'platform__enable_extension' &&
tool.name !== 'platform__read_resource' &&
tool.name !== 'platform__list_resources'
tool.name !== 'platform__read_resource' && tool.name !== 'platform__list_resources'
);
setTools(filteredTools);
}
@@ -46,6 +46,11 @@ export default function PermissionSettingsView({ onClose }: { onClose: () => voi
const extensionsList = await getExtensions(true); // Force refresh
// Filter out disabled extensions
const enabledExtensions = extensionsList.filter((extension) => extension.enabled);
enabledExtensions.push({
name: 'platform',
type: 'builtin',
enabled: true,
});
// Sort extensions by name to maintain consistent order
const sortedExtensions = [...enabledExtensions].sort((a, b) => {
// First sort by builtin
+1 -41
View File
@@ -41,13 +41,6 @@ export interface ToolResponse {
toolResult: ToolCallResult<Content[]>;
}
export interface ToolConfirmationRequest {
id: string;
toolName: string;
arguments: Record<string, unknown>;
prompt?: string;
}
export interface ToolRequestMessageContent {
type: 'toolRequest';
id: string;
@@ -80,33 +73,12 @@ export interface ExtensionCallResult<T> {
error?: string;
}
export interface ExtensionRequest {
id: string;
extensionCall: ExtensionCallResult<ExtensionCall>;
}
export interface ExtensionConfirmationRequest {
id: string;
extensionName: string;
arguments: Record<string, unknown>;
prompt?: string;
}
export interface ExtensionRequestMessageContent {
type: 'extensionRequest';
id: string;
extensionCall: ExtensionCallResult<ExtensionCall>;
extensionName: string;
toolName: string;
}
export type MessageContent =
| TextContent
| ImageContent
| ToolRequestMessageContent
| ToolResponseMessageContent
| ToolConfirmationRequestMessageContent
| ExtensionRequestMessageContent;
| ToolConfirmationRequestMessageContent;
export interface Message {
id?: string;
@@ -220,12 +192,6 @@ export function getToolResponses(message: Message): ToolResponseMessageContent[]
);
}
export function getExtensionRequests(message: Message): ExtensionRequestMessageContent[] {
return message.content.filter(
(content): content is ExtensionRequestMessageContent => content.type === 'extensionRequest'
);
}
export function getToolConfirmationContent(
message: Message
): ToolConfirmationRequestMessageContent {
@@ -235,12 +201,6 @@ export function getToolConfirmationContent(
);
}
export function getExtensionContent(message: Message): ExtensionRequestMessageContent {
return message.content.find(
(content): content is ExtensionRequestMessageContent => content.type === 'extensionRequest'
);
}
export function hasCompletedToolCalls(message: Message): boolean {
const toolRequests = getToolRequests(message);
if (toolRequests.length === 0) return false;