fix: improve and simplify tool call chain rendering (#5704)
This commit is contained in:
@@ -4,12 +4,6 @@ import { extractImagePaths, removeImagePathsFromText } from '../utils/imageUtils
|
||||
import { formatMessageTimestamp } from '../utils/timeUtils';
|
||||
import MarkdownContent from './MarkdownContent';
|
||||
import ToolCallWithResponse from './ToolCallWithResponse';
|
||||
import ToolCallChain from './ToolCallChain';
|
||||
import {
|
||||
identifyConsecutiveToolCalls,
|
||||
shouldHideMessage,
|
||||
getChainForMessage,
|
||||
} from '../utils/toolCallChaining';
|
||||
import {
|
||||
getTextContent,
|
||||
getToolRequests,
|
||||
@@ -22,6 +16,7 @@ import ToolCallConfirmation from './ToolCallConfirmation';
|
||||
import MessageCopyLink from './MessageCopyLink';
|
||||
import { NotificationEvent } from '../hooks/useMessageStream';
|
||||
import { cn } from '../utils';
|
||||
import { identifyConsecutiveToolCalls, shouldHideTimestamp } from '../utils/toolCallChaining';
|
||||
|
||||
interface GooseMessageProps {
|
||||
// messages up to this index are presumed to be "history" from a resumed session, this is used to track older tool confirmation requests
|
||||
@@ -48,15 +43,10 @@ export default function GooseMessage({
|
||||
isStreaming = false,
|
||||
}: GooseMessageProps) {
|
||||
const contentRef = useRef<HTMLDivElement | null>(null);
|
||||
// Track which tool confirmations we've already handled to prevent infinite loops
|
||||
const handledToolConfirmations = useRef<Set<string>>(new Set());
|
||||
|
||||
// Extract text content from the message
|
||||
let textContent = getTextContent(message);
|
||||
|
||||
// Utility to split Chain-of-Thought (CoT) from the visible assistant response.
|
||||
// If the text contains a <think>...</think> block, everything inside is treated as the
|
||||
// CoT and removed from the user-visible text.
|
||||
const splitChainOfThought = (text: string): { visibleText: string; cotText: string | null } => {
|
||||
const regex = /<think>([\s\S]*?)<\/think>/i;
|
||||
const match = text.match(regex);
|
||||
@@ -73,78 +63,30 @@ export default function GooseMessage({
|
||||
};
|
||||
};
|
||||
|
||||
// Split out Chain-of-Thought
|
||||
const { visibleText, cotText } = splitChainOfThought(textContent);
|
||||
|
||||
// Extract image paths from the message content
|
||||
const imagePaths = extractImagePaths(visibleText);
|
||||
|
||||
// Remove image paths from text for display
|
||||
const displayText =
|
||||
imagePaths.length > 0 ? removeImagePathsFromText(visibleText, imagePaths) : visibleText;
|
||||
|
||||
// Memoize the timestamp
|
||||
const timestamp = useMemo(() => formatMessageTimestamp(message.created), [message.created]);
|
||||
|
||||
// Get tool requests from the message
|
||||
const toolRequests = getToolRequests(message);
|
||||
|
||||
// Get current message index
|
||||
const messageIndex = messages.findIndex((msg) => msg.id === message.id);
|
||||
|
||||
// Enhanced chain detection that works during streaming
|
||||
const toolCallChains = useMemo(() => {
|
||||
// Always run chain detection, but handle streaming messages specially
|
||||
const chains = identifyConsecutiveToolCalls(messages);
|
||||
|
||||
// If this message is streaming and has tool calls but no text,
|
||||
// check if it should extend an existing chain
|
||||
if (isStreaming && toolRequests.length > 0 && !displayText.trim()) {
|
||||
// Look for an existing chain that this message could extend
|
||||
const previousMessage = messageIndex > 0 ? messages[messageIndex - 1] : null;
|
||||
if (previousMessage) {
|
||||
const prevToolRequests = getToolRequests(previousMessage);
|
||||
|
||||
// If previous message has tool calls (with or without text), extend its chain
|
||||
if (prevToolRequests.length > 0) {
|
||||
// Find if previous message is part of a chain
|
||||
const prevChain = chains.find((chain) => chain.includes(messageIndex - 1));
|
||||
if (prevChain) {
|
||||
// Extend the existing chain to include this streaming message
|
||||
const extendedChains = chains.map((chain) =>
|
||||
chain === prevChain ? [...chain, messageIndex] : chain
|
||||
);
|
||||
return extendedChains;
|
||||
} else {
|
||||
// Create a new chain with previous and current message
|
||||
return [...chains, [messageIndex - 1, messageIndex]];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return chains;
|
||||
}, [messages, isStreaming, messageIndex, toolRequests, displayText]);
|
||||
|
||||
// Check if this message should be hidden (part of chain but not first)
|
||||
const shouldHide = shouldHideMessage(messageIndex, toolCallChains);
|
||||
|
||||
// Get the chain this message belongs to
|
||||
const messageChain = getChainForMessage(messageIndex, toolCallChains);
|
||||
const toolConfirmationContent = getToolConfirmationContent(message);
|
||||
const toolCallChains = useMemo(() => identifyConsecutiveToolCalls(messages), [messages]);
|
||||
const hideTimestamp = useMemo(
|
||||
() => shouldHideTimestamp(messageIndex, toolCallChains),
|
||||
[messageIndex, toolCallChains]
|
||||
);
|
||||
const hasToolConfirmation = toolConfirmationContent !== undefined;
|
||||
|
||||
// Find tool responses that correspond to the tool requests in this message
|
||||
const toolResponsesMap = useMemo(() => {
|
||||
const responseMap = new Map();
|
||||
|
||||
// Look for tool responses in subsequent messages
|
||||
if (messageIndex !== undefined && messageIndex >= 0) {
|
||||
for (let i = messageIndex + 1; i < messages.length; i++) {
|
||||
const responses = getToolResponses(messages[i]);
|
||||
|
||||
for (const response of responses) {
|
||||
// Check if this response matches any of our tool requests
|
||||
const matchingRequest = toolRequests.find((req) => req.id === response.id);
|
||||
if (matchingRequest) {
|
||||
responseMap.set(response.id, response);
|
||||
@@ -157,21 +99,17 @@ export default function GooseMessage({
|
||||
}, [messages, messageIndex, toolRequests]);
|
||||
|
||||
useEffect(() => {
|
||||
// If the message is the last message in the resumed session and has tool confirmation, it means the tool confirmation
|
||||
// is broken or cancelled, to contonue use the session, we need to append a tool response to avoid mismatch tool result error.
|
||||
if (
|
||||
messageIndex === messageHistoryIndex - 1 &&
|
||||
hasToolConfirmation &&
|
||||
toolConfirmationContent &&
|
||||
!handledToolConfirmations.current.has(toolConfirmationContent.id)
|
||||
) {
|
||||
// Only append the error message if there isn't already a response for this tool confirmation
|
||||
const hasExistingResponse = messages.some((msg) =>
|
||||
getToolResponses(msg).some((response) => response.id === toolConfirmationContent.id)
|
||||
);
|
||||
|
||||
if (!hasExistingResponse) {
|
||||
// Mark this tool confirmation as handled to prevent infinite loop
|
||||
handledToolConfirmations.current.add(toolConfirmationContent.id);
|
||||
|
||||
appendMessage(
|
||||
@@ -188,14 +126,6 @@ export default function GooseMessage({
|
||||
appendMessage,
|
||||
]);
|
||||
|
||||
// If this message should be hidden (part of chain but not first), don't render it
|
||||
if (shouldHide) {
|
||||
return null;
|
||||
}
|
||||
|
||||
// Determine rendering logic based on chain membership and content
|
||||
const isFirstInChain = messageChain && messageChain[0] === messageIndex;
|
||||
|
||||
return (
|
||||
<div className="goose-message flex w-[90%] justify-start min-w-0">
|
||||
<div className="flex flex-col w-full min-w-0">
|
||||
@@ -216,7 +146,6 @@ export default function GooseMessage({
|
||||
<MarkdownContent content={displayText} />
|
||||
</div>
|
||||
|
||||
{/* Image previews */}
|
||||
{imagePaths.length > 0 && (
|
||||
<div className="mt-4">
|
||||
{imagePaths.map((imagePath, index) => (
|
||||
@@ -244,39 +173,28 @@ export default function GooseMessage({
|
||||
|
||||
{toolRequests.length > 0 && (
|
||||
<div className={cn(displayText && 'mt-2')}>
|
||||
{isFirstInChain ? (
|
||||
<ToolCallChain
|
||||
messages={messages}
|
||||
chainIndices={messageChain}
|
||||
toolCallNotifications={toolCallNotifications}
|
||||
toolResponsesMap={toolResponsesMap}
|
||||
messageHistoryIndex={messageHistoryIndex}
|
||||
isStreaming={isStreaming}
|
||||
/>
|
||||
) : !messageChain ? (
|
||||
<div className="relative flex flex-col w-full">
|
||||
<div className="flex flex-col gap-3">
|
||||
{toolRequests.map((toolRequest) => (
|
||||
<div className="goose-message-tool" key={toolRequest.id}>
|
||||
<ToolCallWithResponse
|
||||
isCancelledMessage={
|
||||
messageIndex < messageHistoryIndex &&
|
||||
toolResponsesMap.get(toolRequest.id) == undefined
|
||||
}
|
||||
toolRequest={toolRequest}
|
||||
toolResponse={toolResponsesMap.get(toolRequest.id)}
|
||||
notifications={toolCallNotifications.get(toolRequest.id)}
|
||||
isStreamingMessage={isStreaming}
|
||||
append={append}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
<div className="text-xs text-text-muted pt-1 transition-all duration-200 group-hover:-translate-y-4 group-hover:opacity-0">
|
||||
{!isStreaming && timestamp}
|
||||
</div>
|
||||
<div className="relative flex flex-col w-full">
|
||||
<div className="flex flex-col gap-3">
|
||||
{toolRequests.map((toolRequest) => (
|
||||
<div className="goose-message-tool" key={toolRequest.id}>
|
||||
<ToolCallWithResponse
|
||||
isCancelledMessage={
|
||||
messageIndex < messageHistoryIndex &&
|
||||
toolResponsesMap.get(toolRequest.id) == undefined
|
||||
}
|
||||
toolRequest={toolRequest}
|
||||
toolResponse={toolResponsesMap.get(toolRequest.id)}
|
||||
notifications={toolCallNotifications.get(toolRequest.id)}
|
||||
isStreamingMessage={isStreaming}
|
||||
append={append}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
) : null}
|
||||
<div className="text-xs text-text-muted transition-all duration-200 group-hover:-translate-y-4 group-hover:opacity-0 pt-1">
|
||||
{!isStreaming && !hideTimestamp && timestamp}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
* - Configurable batch size and delay
|
||||
*/
|
||||
|
||||
import { useCallback, useEffect, useRef, useState } from 'react';
|
||||
import { useCallback, useEffect, useMemo, useRef, useState } from 'react';
|
||||
import { Message } from '../api';
|
||||
import GooseMessage from './GooseMessage';
|
||||
import UserMessage from './UserMessage';
|
||||
@@ -22,6 +22,7 @@ import { SystemNotificationInline } from './context_management/SystemNotificatio
|
||||
import { NotificationEvent } from '../hooks/useMessageStream';
|
||||
import LoadingGoose from './LoadingGoose';
|
||||
import { ChatType } from '../types/chat';
|
||||
import { identifyConsecutiveToolCalls, isInChain } from '../utils/toolCallChaining';
|
||||
|
||||
interface ProgressiveMessageListProps {
|
||||
messages: Message[];
|
||||
@@ -161,6 +162,9 @@ export default function ProgressiveMessageList({
|
||||
return () => window.removeEventListener('keydown', handleKeyDown);
|
||||
}, [isLoading, messages.length]);
|
||||
|
||||
// Detect tool call chains
|
||||
const toolCallChains = useMemo(() => identifyConsecutiveToolCalls(messages), [messages]);
|
||||
|
||||
// Render messages up to the current rendered count
|
||||
const renderMessages = useCallback(() => {
|
||||
const messagesToRender = messages.slice(0, renderedCount);
|
||||
@@ -195,11 +199,12 @@ export default function ProgressiveMessageList({
|
||||
}
|
||||
|
||||
const isUser = isUserMessage(message);
|
||||
const messageIsInChain = isInChain(index, toolCallChains);
|
||||
|
||||
return (
|
||||
<div
|
||||
key={message.id && `${message.id}-${message.content.length}`}
|
||||
className={`relative ${index === 0 ? 'mt-0' : 'mt-4'} ${isUser ? 'user' : 'assistant'}`}
|
||||
className={`relative ${index === 0 ? 'mt-0' : 'mt-4'} ${isUser ? 'user' : 'assistant'} ${messageIsInChain ? 'in-chain' : ''}`}
|
||||
data-testid="message-container"
|
||||
>
|
||||
{isUser ? (
|
||||
@@ -238,6 +243,7 @@ export default function ProgressiveMessageList({
|
||||
toolCallNotifications,
|
||||
isStreamingMessage,
|
||||
onMessageUpdate,
|
||||
toolCallChains,
|
||||
]);
|
||||
|
||||
return (
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
import { formatMessageTimestamp } from '../utils/timeUtils';
|
||||
import { Message } from '../api';
|
||||
import { getToolRequests } from '../types/message';
|
||||
import { NotificationEvent } from '../hooks/useMessageStream';
|
||||
import ToolCallWithResponse from './ToolCallWithResponse';
|
||||
|
||||
interface ToolCallChainProps {
|
||||
messages: Message[];
|
||||
chainIndices: number[];
|
||||
toolCallNotifications: Map<string, NotificationEvent[]>;
|
||||
toolResponsesMap: Map<string, import('../types/message').ToolResponseMessageContent>;
|
||||
messageHistoryIndex: number;
|
||||
isStreaming?: boolean;
|
||||
}
|
||||
|
||||
export default function ToolCallChain({
|
||||
messages,
|
||||
chainIndices,
|
||||
toolCallNotifications,
|
||||
toolResponsesMap,
|
||||
messageHistoryIndex,
|
||||
isStreaming = false,
|
||||
}: ToolCallChainProps) {
|
||||
const lastMessageIndex = chainIndices[chainIndices.length - 1];
|
||||
const lastMessage = messages[lastMessageIndex];
|
||||
const timestamp = lastMessage ? formatMessageTimestamp(lastMessage.created) : '';
|
||||
|
||||
return (
|
||||
<div className="relative flex flex-col w-full">
|
||||
<div className="flex flex-col gap-3">
|
||||
{chainIndices.map((messageIndex) => {
|
||||
const message = messages[messageIndex];
|
||||
const toolRequests = getToolRequests(message);
|
||||
|
||||
return toolRequests.map((toolRequest) => (
|
||||
<div key={toolRequest.id} className="goose-message-tool">
|
||||
<ToolCallWithResponse
|
||||
isCancelledMessage={
|
||||
messageIndex < messageHistoryIndex &&
|
||||
toolResponsesMap.get(toolRequest.id) == undefined
|
||||
}
|
||||
toolRequest={toolRequest}
|
||||
toolResponse={toolResponsesMap.get(toolRequest.id)}
|
||||
notifications={toolCallNotifications.get(toolRequest.id)}
|
||||
isStreamingMessage={isStreaming}
|
||||
/>
|
||||
</div>
|
||||
));
|
||||
})}
|
||||
</div>
|
||||
|
||||
<div className="text-xs text-text-muted pt-1 transition-all duration-200 group-hover:-translate-y-4 group-hover:opacity-0">
|
||||
{!isStreaming && timestamp}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -247,24 +247,12 @@ function ToolCallView({
|
||||
}
|
||||
}, [toolResponse, startTime]);
|
||||
|
||||
const toolResults: { result: Content; isExpandToolResults: boolean }[] =
|
||||
const toolResults: Content[] =
|
||||
loadingStatus === 'success' && Array.isArray(toolResponse?.toolResult.value)
|
||||
? toolResponse!.toolResult.value
|
||||
.filter((item) => {
|
||||
const audience = item.annotations?.audience as string[] | undefined;
|
||||
return !audience || audience.includes('user');
|
||||
})
|
||||
.map((item) => {
|
||||
// Use user preference for detailed/concise, but still respect high priority items
|
||||
const priority = (item.annotations?.priority as number | undefined) ?? -1;
|
||||
const isHighPriority = priority >= 0.5;
|
||||
const shouldExpandBasedOnStyle = responseStyle === 'detailed' || responseStyle === null;
|
||||
|
||||
return {
|
||||
result: item,
|
||||
isExpandToolResults: isHighPriority || shouldExpandBasedOnStyle,
|
||||
};
|
||||
})
|
||||
? toolResponse!.toolResult.value.filter((item) => {
|
||||
const audience = item.annotations?.audience as string[] | undefined;
|
||||
return !audience || audience.includes('user');
|
||||
})
|
||||
: [];
|
||||
|
||||
const logs = notifications
|
||||
@@ -290,17 +278,6 @@ function ToolCallView({
|
||||
const isRenderingProgress =
|
||||
loadingStatus === 'loading' && (progressEntries.length > 0 || (logs || []).length > 0);
|
||||
|
||||
// Determine if the main tool call should be expanded
|
||||
const isShouldExpand = (() => {
|
||||
// Always expand if there are high priority results that need to be shown
|
||||
const hasHighPriorityResults = toolResults.some((v) => v.isExpandToolResults);
|
||||
|
||||
// Also expand based on user preference for detailed mode
|
||||
const shouldExpandBasedOnStyle = responseStyle === 'detailed' || responseStyle === null;
|
||||
|
||||
return hasHighPriorityResults || shouldExpandBasedOnStyle;
|
||||
})();
|
||||
|
||||
// Function to create a descriptive representation of what the tool is doing
|
||||
const getToolDescription = (): string | null => {
|
||||
const args = toolCall.arguments as Record<string, ToolCallArgumentValue>;
|
||||
@@ -488,7 +465,7 @@ function ToolCallView({
|
||||
return (
|
||||
<ToolCallExpandable
|
||||
isStartExpanded={isRenderingProgress}
|
||||
isForceExpand={isShouldExpand}
|
||||
isForceExpand={false}
|
||||
label={
|
||||
extensionTooltip ? (
|
||||
<TooltipWrapper tooltipContent={extensionTooltip} side="top" align="start">
|
||||
@@ -529,13 +506,11 @@ function ToolCallView({
|
||||
{/* Tool Output */}
|
||||
{!isCancelledMessage && (
|
||||
<>
|
||||
{toolResults.map(({ result, isExpandToolResults }, index) => {
|
||||
return (
|
||||
<div key={index} className={cn('border-t border-borderSubtle')}>
|
||||
<ToolResultView result={result} isStartExpanded={isExpandToolResults} />
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
{toolResults.map((result, index) => (
|
||||
<div key={index} className={cn('border-t border-borderSubtle')}>
|
||||
<ToolResultView result={result} isStartExpanded={false} />
|
||||
</div>
|
||||
))}
|
||||
</>
|
||||
)}
|
||||
</ToolCallExpandable>
|
||||
|
||||
@@ -555,7 +555,7 @@ p > code.bg-inline-code {
|
||||
scrollbar-width: thin;
|
||||
}
|
||||
|
||||
.assistant:has(+ .user) .goose-message {
|
||||
.assistant:has(+ .user):not(.in-chain) .goose-message {
|
||||
padding-bottom: 24px;
|
||||
}
|
||||
|
||||
|
||||
@@ -56,6 +56,20 @@ export function shouldHideMessage(messageIndex: number, chains: number[][]): boo
|
||||
return false;
|
||||
}
|
||||
|
||||
export function shouldHideTimestamp(messageIndex: number, chains: number[][]): boolean {
|
||||
for (const chain of chains) {
|
||||
if (chain.includes(messageIndex)) {
|
||||
// Hide timestamp for all but the last message in the chain
|
||||
return chain[chain.length - 1] !== messageIndex;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
export function isInChain(messageIndex: number, chains: number[][]): boolean {
|
||||
return chains.some((chain) => chain.includes(messageIndex));
|
||||
}
|
||||
|
||||
export function getChainForMessage(messageIndex: number, chains: number[][]): number[] | null {
|
||||
return chains.find((chain) => chain.includes(messageIndex)) || null;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user