diff --git a/crates/goose/src/conversation/message.rs b/crates/goose/src/conversation/message.rs index b3c6ab6b7b..b5846d9b0f 100644 --- a/crates/goose/src/conversation/message.rs +++ b/crates/goose/src/conversation/message.rs @@ -39,28 +39,44 @@ where Ok(content) } +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, ToSchema)] +#[serde(tag = "status", rename_all = "lowercase")] +pub enum ToolCallResult { + Success { value: T }, + Error { error: String }, +} + +impl From> for ToolCallResult { + fn from(result: ToolResult) -> Self { + match result { + Ok(value) => ToolCallResult::Success { value }, + Err(error) => ToolCallResult::Error { + error: error.to_string(), + }, + } + } +} + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] #[derive(ToSchema)] pub struct ToolRequest { pub id: String, - #[serde(with = "tool_result_serde")] - #[schema(value_type = Object)] - pub tool_call: ToolResult, + pub tool_call: ToolCallResult, } impl ToolRequest { pub fn to_readable_string(&self) -> String { match &self.tool_call { - Ok(tool_call) => { + ToolCallResult::Success { value } => { format!( "Tool: {}, Args: {}", - tool_call.name, - serde_json::to_string_pretty(&tool_call.arguments) + value.name, + serde_json::to_string_pretty(&value.arguments) .unwrap_or_else(|_| "<>".to_string()) ) } - Err(e) => format!("Invalid tool call: {}", e), + ToolCallResult::Error { error } => format!("Invalid tool call: {}", error), } } } @@ -71,7 +87,6 @@ impl ToolRequest { pub struct ToolResponse { pub id: String, #[serde(with = "tool_result_serde")] - #[schema(value_type = Object)] pub tool_result: ToolResult>, } @@ -100,8 +115,6 @@ pub struct RedactedThinkingContent { #[serde(rename_all = "camelCase")] pub struct FrontendToolRequest { pub id: String, - #[serde(with = "tool_result_serde")] - #[schema(value_type = Object)] pub tool_call: ToolResult, } diff --git a/crates/goose/src/mcp_utils.rs b/crates/goose/src/mcp_utils.rs index 9f319e34a5..7360fa2b51 100644 --- a/crates/goose/src/mcp_utils.rs +++ b/crates/goose/src/mcp_utils.rs @@ -1,4 +1,4 @@ pub use rmcp::model::ErrorData; /// Type alias for tool results -pub type ToolResult = std::result::Result; +pub type ToolResult = Result; diff --git a/ui/desktop/openapi.json b/ui/desktop/openapi.json index a8398f713d..4877cef6f0 100644 --- a/ui/desktop/openapi.json +++ b/ui/desktop/openapi.json @@ -2754,7 +2754,7 @@ "type": "string" }, "toolCall": { - "type": "object" + "$ref": "#/components/schemas/ToolResult" } } }, @@ -4339,7 +4339,7 @@ "type": "string" }, "toolCall": { - "type": "object" + "$ref": "#/components/schemas/ToolResult" } } }, @@ -4354,7 +4354,7 @@ "type": "string" }, "toolResult": { - "type": "object" + "$ref": "#/components/schemas/ToolResult" } } }, diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index 07231583ab..c03c3e87be 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -61,10 +61,10 @@ import { useChatEngine } from '../hooks/useChatEngine'; import { useRecipeManager } from '../hooks/useRecipeManager'; import { useFileDrop } from '../hooks/useFileDrop'; import { useCostTracking } from '../hooks/useCostTracking'; -import { Message } from '../types/message'; import { ChatState } from '../types/chatState'; import { ChatType } from '../types/chat'; import { useToolCount } from './alerts/useToolCount'; +import { Message } from '../api'; // Context for sharing current model info const CurrentModelContext = createContext<{ model: string; mode: string } | null>(null); diff --git a/ui/desktop/src/components/BaseChat2.tsx b/ui/desktop/src/components/BaseChat2.tsx index 9e045ea185..5ffaa64ec6 100644 --- a/ui/desktop/src/components/BaseChat2.tsx +++ b/ui/desktop/src/components/BaseChat2.tsx @@ -53,12 +53,13 @@ import { MainPanelLayout } from './Layout/MainPanelLayout'; import ChatInput from './ChatInput'; import { ScrollArea, ScrollAreaHandle } from './ui/scroll-area'; import { useFileDrop } from '../hooks/useFileDrop'; -import { Message } from '../types/message'; +import { Message } from '../api'; import { ChatState } from '../types/chatState'; import { ChatType } from '../types/chat'; import { useIsMobile } from '../hooks/use-mobile'; import { useSidebar } from './ui/sidebar'; import { cn } from '../utils'; +import { useChatStream } from '../hooks/useChatStream'; interface BaseChatProps { chat: ChatType | null; @@ -114,6 +115,17 @@ function BaseChatContent({ // session: sessionMetadata, // }); + const { chatState, handleSubmit, stopStreaming } = useChatStream({ + sessionId: chat?.sessionId || '', + messages, + setMessages: (newMessages) => { + if (chat) { + setChat({ ...chat, messages: newMessages }); + } + }, + onStreamFinish: onMessageStreamFinish, + }); + // TODO(Douwe): send this to the chatbox instead, possibly autosubmit? or backend const append = (_txt: string) => {}; diff --git a/ui/desktop/src/components/ChatInput.tsx b/ui/desktop/src/components/ChatInput.tsx index 92763ddea8..33916c6cb4 100644 --- a/ui/desktop/src/components/ChatInput.tsx +++ b/ui/desktop/src/components/ChatInput.tsx @@ -8,7 +8,7 @@ import { Attach, Send, Close, Microphone } from './icons'; import { ChatState } from '../types/chatState'; import debounce from 'lodash/debounce'; import { LocalMessageStorage } from '../utils/localMessageStorage'; -import { Message } from '../types/message'; +import { Message } from '../api'; import { DirSwitcher } from './bottom_menu/DirSwitcher'; import ModelsBottomBar from './settings/models/bottom_bar/ModelsBottomBar'; import { BottomMenuModeSelection } from './bottom_menu/BottomMenuModeSelection'; diff --git a/ui/desktop/src/components/GooseMessage.tsx b/ui/desktop/src/components/GooseMessage.tsx index cbc8aa6039..6c75bb03d1 100644 --- a/ui/desktop/src/components/GooseMessage.tsx +++ b/ui/desktop/src/components/GooseMessage.tsx @@ -11,13 +11,13 @@ import { getChainForMessage, } from '../utils/toolCallChaining'; import { - Message, getTextContent, getToolRequests, getToolResponses, getToolConfirmationContent, createToolErrorResponseMessage, } from '../types/message'; +import { Message } from '../api'; import ToolCallConfirmation from './ToolCallConfirmation'; import MessageCopyLink from './MessageCopyLink'; import { NotificationEvent } from '../hooks/useMessageStream'; diff --git a/ui/desktop/src/components/MCPUIResourceRenderer.tsx b/ui/desktop/src/components/MCPUIResourceRenderer.tsx index 3e5ad494a5..c73aab418b 100644 --- a/ui/desktop/src/components/MCPUIResourceRenderer.tsx +++ b/ui/desktop/src/components/MCPUIResourceRenderer.tsx @@ -7,13 +7,14 @@ import { UIActionResultToolCall, } from '@mcp-ui/client'; import { useState, useEffect } from 'react'; -import { ResourceContent } from '../types/message'; import { toast } from 'react-toastify'; +import { EmbeddedResource } from '../api'; interface MCPUIResourceRendererProps { - content: ResourceContent; + content: EmbeddedResource & { type: 'resource' }; appendPromptToChat?: (value: string) => void; } + type UISizeChange = { type: 'ui-size-change'; payload: { diff --git a/ui/desktop/src/components/ProgressiveMessageList.tsx b/ui/desktop/src/components/ProgressiveMessageList.tsx index 0b4e87f6dc..9bdd0ab689 100644 --- a/ui/desktop/src/components/ProgressiveMessageList.tsx +++ b/ui/desktop/src/components/ProgressiveMessageList.tsx @@ -15,7 +15,7 @@ */ import { useCallback, useEffect, useRef, useState } from 'react'; -import { Message } from '../types/message'; +import { Message } from '../api'; import GooseMessage from './GooseMessage'; import UserMessage from './UserMessage'; import { CompactionMarker } from './context_management/CompactionMarker'; diff --git a/ui/desktop/src/components/ToolCallChain.tsx b/ui/desktop/src/components/ToolCallChain.tsx index ea02e343c4..a952c3771d 100644 --- a/ui/desktop/src/components/ToolCallChain.tsx +++ b/ui/desktop/src/components/ToolCallChain.tsx @@ -1,5 +1,6 @@ import { formatMessageTimestamp } from '../utils/timeUtils'; -import { Message, getToolRequests } from '../types/message'; +import { Message } from '../api'; +import { getToolRequests } from '../types/message'; import { NotificationEvent } from '../hooks/useMessageStream'; import ToolCallWithResponse from './ToolCallWithResponse'; diff --git a/ui/desktop/src/components/ToolCallConfirmation.tsx b/ui/desktop/src/components/ToolCallConfirmation.tsx index f58de292e8..ab0b080ed7 100644 --- a/ui/desktop/src/components/ToolCallConfirmation.tsx +++ b/ui/desktop/src/components/ToolCallConfirmation.tsx @@ -2,7 +2,7 @@ import { useState, useEffect } from 'react'; import { snakeToTitleCase } from '../utils'; import PermissionModal from './settings/permission/PermissionModal'; import { ChevronRight } from 'lucide-react'; -import { confirmPermission } from '../api'; +import { confirmPermission, ToolConfirmationRequest } from '../api'; import { Button } from './ui/button'; const ALLOW_ONCE = 'allow_once'; @@ -20,13 +20,11 @@ const toolConfirmationState = new Map< } >(); -import { ToolConfirmationRequestMessageContent } from '../types/message'; - interface ToolConfirmationProps { sessionId: string; isCancelledMessage: boolean; isClicked: boolean; - toolConfirmationContent: ToolConfirmationRequestMessageContent; + toolConfirmationContent: ToolConfirmationRequest & { type: 'toolConfirmationRequest' }; } export default function ToolConfirmation({ diff --git a/ui/desktop/src/components/ToolCallWithResponse.tsx b/ui/desktop/src/components/ToolCallWithResponse.tsx index da52d2b00f..a2c5fc3bc6 100644 --- a/ui/desktop/src/components/ToolCallWithResponse.tsx +++ b/ui/desktop/src/components/ToolCallWithResponse.tsx @@ -4,7 +4,7 @@ import React, { useEffect, useRef, useState } from 'react'; import { Button } from './ui/button'; import { ToolCallArguments, ToolCallArgumentValue } from './ToolCallArguments'; import MarkdownContent from './MarkdownContent'; -import { Content, ToolRequestMessageContent, ToolResponseMessageContent } from '../types/message'; +import { ToolRequestMessageContent, ToolResponseMessageContent } from '../types/message'; import { cn, snakeToTitleCase } from '../utils'; import { LoadingStatus } from './ui/Dot'; import { NotificationEvent } from '../hooks/useMessageStream'; @@ -12,6 +12,7 @@ import { ChevronRight, FlaskConical } from 'lucide-react'; import { TooltipWrapper } from './settings/providers/subcomponents/buttons/TooltipWrapper'; import MCPUIResourceRenderer from './MCPUIResourceRenderer'; import { isUIResource } from '@mcp-ui/client'; +import { Content } from '../api'; interface ToolCallWithResponseProps { isCancelledMessage: boolean; @@ -27,10 +28,10 @@ export default function ToolCallWithResponse({ toolRequest, toolResponse, notifications, - isStreamingMessage = false, + isStreamingMessage, append, }: ToolCallWithResponseProps) { - const toolCall = toolRequest.toolCall.status === 'success' ? toolRequest.toolCall.value : null; + const toolCall = toolRequest.toolCall as { name: string; arguments: Record }; if (!toolCall) { return null; } @@ -53,11 +54,12 @@ export default function ToolCallWithResponse({ /> {/* MCP UI — Inline */} - {toolResponse?.toolResult?.value && - toolResponse.toolResult.value.map((content, index) => { + {toolResponse?.toolResult && + Array.isArray((toolResponse.toolResult as any).value) && + (toolResponse.toolResult as any).value.map((content: Content, index: number) => { if (isUIResource(content)) { return ( -
+
@@ -211,7 +213,7 @@ function ToolCallView({ ? shouldShowAsComplete ? 'success' : 'loading' - : toolResponse.toolResult.status; + : (toolResponse.toolResult as any).status || 'success'; // Tool call timing tracking const [startTime, setStartTime] = useState(null); @@ -547,30 +549,34 @@ interface ToolResultViewProps { } function ToolResultView({ result, isStartExpanded }: ToolResultViewProps) { + const resultAny = result as any; + return ( Output} isStartExpanded={isStartExpanded} >
- {result.type === 'text' && result.text && ( + {'text' in resultAny && resultAny.text && ( )} - {result.type === 'image' && ( - Tool result { - console.error('Failed to load image'); - e.currentTarget.style.display = 'none'; - }} - /> - )} - {result.type === 'resource' && ( + {'mimeType' in resultAny && + 'data' in resultAny && + resultAny.mimeType?.startsWith('image') && ( + Tool result { + console.error('Failed to load image'); + e.currentTarget.style.display = 'none'; + }} + /> + )} + {'resource' in resultAny && (
{JSON.stringify(result, null, 2)}
)}
diff --git a/ui/desktop/src/components/UserMessage.tsx b/ui/desktop/src/components/UserMessage.tsx index 036a7c5b55..1292d4fbeb 100644 --- a/ui/desktop/src/components/UserMessage.tsx +++ b/ui/desktop/src/components/UserMessage.tsx @@ -1,8 +1,9 @@ -import { useRef, useMemo, useState, useEffect, useCallback } from 'react'; +import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; import ImagePreview from './ImagePreview'; import { extractImagePaths, removeImagePathsFromText } from '../utils/imageUtils'; import MarkdownContent from './MarkdownContent'; -import { Message, getTextContent } from '../types/message'; +import { getTextContent } from '../types/message'; +import { Message } from '../api'; import MessageCopyLink from './MessageCopyLink'; import { formatMessageTimestamp } from '../utils/timeUtils'; import Edit from './icons/Edit'; diff --git a/ui/desktop/src/components/context_management/CompactionMarker.tsx b/ui/desktop/src/components/context_management/CompactionMarker.tsx index a08563855a..f453aa9b06 100644 --- a/ui/desktop/src/components/context_management/CompactionMarker.tsx +++ b/ui/desktop/src/components/context_management/CompactionMarker.tsx @@ -1,5 +1,5 @@ import React from 'react'; -import { Message, SummarizationRequestedContent } from '../../types/message'; +import { Message, SummarizationRequested } from '../../api'; interface CompactionMarkerProps { message: Message; @@ -7,8 +7,9 @@ interface CompactionMarkerProps { export const CompactionMarker: React.FC = ({ message }) => { const compactionContent = message.content.find( - (content) => content.type === 'summarizationRequested' - ) as SummarizationRequestedContent | undefined; + (content): content is SummarizationRequested & { type: 'summarizationRequested' } => + content.type === 'summarizationRequested' + ); const markerText = compactionContent?.msg || 'Conversation compacted'; diff --git a/ui/desktop/src/components/context_management/ContextManager.tsx b/ui/desktop/src/components/context_management/ContextManager.tsx index 4d5ad9d9e1..54c606bd88 100644 --- a/ui/desktop/src/components/context_management/ContextManager.tsx +++ b/ui/desktop/src/components/context_management/ContextManager.tsx @@ -1,6 +1,6 @@ import React, { createContext, useContext, useState, useCallback } from 'react'; -import { Message } from '../../types/message'; -import { manageContextFromBackend, convertApiMessageToFrontendMessage } from './index'; +import { manageContextFromBackend } from './index'; +import { Message } from '../../api'; // Define the context management interface interface ContextManagerState { @@ -53,21 +53,14 @@ export const ContextManagerProvider: React.FC<{ children: React.ReactNode }> = ( sessionId: sessionId, }); - // Convert API messages to frontend messages - // The server now handles all visibility - we just display what we receive - const convertedMessages = summaryResponse.messages.map((apiMessage) => - convertApiMessageToFrontendMessage(apiMessage) - ); - - // Replace messages with the server-provided messages - setMessages(convertedMessages); + setMessages(summaryResponse.messages); // Only automatically submit the continuation message for auto-compaction (context limit reached) // Manual compaction should just compact without continuing the conversation if (!isManual) { // Automatically submit the continuation message to continue the conversation // This should be the third message (index 2) which contains the "I ran into a context length exceeded error..." text - const continuationMessage = convertedMessages[2]; + const continuationMessage = summaryResponse.messages[2]; if (continuationMessage) { setTimeout(() => { append(continuationMessage); diff --git a/ui/desktop/src/components/context_management/__tests__/CompactionMarker.test.tsx b/ui/desktop/src/components/context_management/__tests__/CompactionMarker.test.tsx index 8e1691ba24..3244bbdfce 100644 --- a/ui/desktop/src/components/context_management/__tests__/CompactionMarker.test.tsx +++ b/ui/desktop/src/components/context_management/__tests__/CompactionMarker.test.tsx @@ -1,7 +1,7 @@ import { describe, it, expect } from 'vitest'; import { render, screen } from '@testing-library/react'; import { CompactionMarker } from '../CompactionMarker'; -import { Message } from '../../../types/message'; +import { Message } from '../../../api'; describe('CompactionMarker', () => { it('should render default message when no summarizationRequested content found', () => { diff --git a/ui/desktop/src/components/context_management/__tests__/ContextManager.test.tsx b/ui/desktop/src/components/context_management/__tests__/ContextManager.test.tsx index 53ac40b7a0..fe2f4134dc 100644 --- a/ui/desktop/src/components/context_management/__tests__/ContextManager.test.tsx +++ b/ui/desktop/src/components/context_management/__tests__/ContextManager.test.tsx @@ -1,9 +1,8 @@ import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest'; import { renderHook, act } from '@testing-library/react'; import { ContextManagerProvider, useContextManager } from '../ContextManager'; -import { Message } from '../../../types/message'; import * as contextManagement from '../index'; -import { ContextManageResponse } from '../../../api'; +import { ContextManageResponse, Message } from '../../../api'; // Mock the context management functions vi.mock('../index', () => ({ @@ -12,9 +11,6 @@ vi.mock('../index', () => ({ })); const mockManageContextFromBackend = vi.mocked(contextManagement.manageContextFromBackend); -const mockConvertApiMessageToFrontendMessage = vi.mocked( - contextManagement.convertApiMessageToFrontendMessage -); describe('ContextManager', () => { const mockMessages: Message[] = [ @@ -157,12 +153,6 @@ describe('ContextManager', () => { ], }; - // Mock the conversion function to return different messages based on call order - mockConvertApiMessageToFrontendMessage - .mockReturnValueOnce(mockCompactionMarker) // First call - compaction marker - .mockReturnValueOnce(mockSummaryMessage) // Second call - summary - .mockReturnValueOnce(mockContinuationMessage); // Third call - continuation - const { result } = renderContextManager(); await act(async () => { @@ -180,33 +170,6 @@ describe('ContextManager', () => { sessionId: 'test-session-id', }); - // Verify conversion calls with correct parameters - expect(mockConvertApiMessageToFrontendMessage).toHaveBeenNthCalledWith( - 1, - expect.objectContaining({ - content: [ - { type: 'summarizationRequested', msg: 'Conversation compacted and summarized' }, - ], - }) - ); - expect(mockConvertApiMessageToFrontendMessage).toHaveBeenNthCalledWith( - 2, - expect.objectContaining({ - content: [{ type: 'text', text: 'Summary content' }], - }) - ); - expect(mockConvertApiMessageToFrontendMessage).toHaveBeenNthCalledWith( - 3, - expect.objectContaining({ - content: [ - { - type: 'text', - text: expect.stringContaining('The previous message contains a summary'), - }, - ], - }) - ); - // Expect setMessages to be called with all 3 converted messages expect(mockSetMessages).toHaveBeenCalledWith([ mockCompactionMarker, @@ -290,8 +253,6 @@ describe('ContextManager', () => { tokenCounts: [100, 50], }); - mockConvertApiMessageToFrontendMessage.mockReturnValue(mockSummaryMessage); - await act(async () => { await promise; }); @@ -382,11 +343,6 @@ describe('ContextManager', () => { ], }; - mockConvertApiMessageToFrontendMessage - .mockReturnValueOnce(mockCompactionMarker) - .mockReturnValueOnce(mockSummaryMessage) - .mockReturnValueOnce(mockContinuationMessage); - const { result } = renderContextManager(); await act(async () => { @@ -431,8 +387,6 @@ describe('ContextManager', () => { tokenCounts: [100, 50], }); - mockConvertApiMessageToFrontendMessage.mockReturnValue(mockSummaryMessage); - const { result } = renderContextManager(); await act(async () => { @@ -500,11 +454,6 @@ describe('ContextManager', () => { ], }; - mockConvertApiMessageToFrontendMessage - .mockReturnValueOnce(mockCompactionMarker) - .mockReturnValueOnce(mockSummaryMessage) - .mockReturnValueOnce(mockContinuationMessage); - const { result } = renderContextManager(); await act(async () => { @@ -571,8 +520,6 @@ describe('ContextManager', () => { content: [{ type: 'toolResponse', id: 'test', toolResult: { status: 'success' } }], }; - mockConvertApiMessageToFrontendMessage.mockReturnValue(mockMessageWithoutText); - const { result } = renderContextManager(); await act(async () => { diff --git a/ui/desktop/src/components/context_management/index.ts b/ui/desktop/src/components/context_management/index.ts index c54d95eb0f..c4f7062196 100644 --- a/ui/desktop/src/components/context_management/index.ts +++ b/ui/desktop/src/components/context_management/index.ts @@ -1,149 +1,29 @@ -import { - Message as FrontendMessage, - Content as FrontendContent, - MessageContent as FrontendMessageContent, - ToolCallResult, - ToolCall, - Role, -} from '../../types/message'; -import { - ContextManageRequest, - ContextManageResponse, - manageContext, - Message as ApiMessage, - MessageContent as ApiMessageContent, -} from '../../api'; -import { generateId } from 'ai'; +import { ContextManageRequest, ContextManageResponse, manageContext, Message } from '../../api'; export async function manageContextFromBackend({ messages, manageAction, sessionId, }: { - messages: FrontendMessage[]; + messages: Message[]; manageAction: 'truncation' | 'summarize'; sessionId: string; }): Promise { - try { - const contextManagementRequest = { manageAction, messages, sessionId }; + const contextManagementRequest = { manageAction, messages, sessionId }; - // Cast to the API-expected type - const result = await manageContext({ - body: contextManagementRequest as unknown as ContextManageRequest, - }); + // Cast to the API-expected type + const result = await manageContext({ + body: contextManagementRequest as unknown as ContextManageRequest, + }); - // Check for errors in the result - if (result.error) { - throw new Error(`Context management failed: ${result.error}`); - } - - // Extract the actual data from the result - if (!result.data) { - throw new Error('Context management returned no data'); - } - - return result.data; - } catch (error) { - console.error(`Context management failed: ${error || 'Unknown error'}`); - throw new Error( - `Context management failed: ${error || 'Unknown error'}\n\nStart a new session.` - ); - } -} - -// Function to convert API Message to frontend Message -export function convertApiMessageToFrontendMessage(apiMessage: ApiMessage): FrontendMessage { - return { - id: generateId(), - role: apiMessage.role as Role, - created: apiMessage.created ?? Math.floor(Date.now() / 1000), - content: apiMessage.content - .map((apiContent) => mapApiContentToFrontendMessageContent(apiContent)) - .filter((content): content is FrontendMessageContent => content !== null), - }; -} - -// Function to convert API MessageContent to frontend MessageContent -function mapApiContentToFrontendMessageContent( - apiContent: ApiMessageContent -): FrontendMessageContent | null { - // Handle each content type specifically based on its "type" property - if (apiContent.type === 'text') { - return { - type: 'text', - text: apiContent.text, - annotations: apiContent.annotations as Record | undefined, - }; - } else if (apiContent.type === 'image') { - return { - type: 'image', - data: apiContent.data, - mimeType: apiContent.mimeType, - annotations: apiContent.annotations as Record | undefined, - }; - } else if (apiContent.type === 'toolRequest') { - // Ensure the toolCall has the correct type structure - const toolCall = apiContent.toolCall as unknown as ToolCallResult; - - return { - type: 'toolRequest', - id: apiContent.id, - toolCall: toolCall, - }; - } else if (apiContent.type === 'toolResponse') { - // Ensure the toolResult has the correct type structure - const toolResult = apiContent.toolResult as unknown as ToolCallResult; - - return { - type: 'toolResponse', - id: apiContent.id, - toolResult: toolResult, - }; - } else if (apiContent.type === 'toolConfirmationRequest') { - return { - type: 'toolConfirmationRequest', - id: apiContent.id, - toolName: apiContent.toolName, - arguments: apiContent.arguments as Record, - prompt: apiContent.prompt === null ? undefined : apiContent.prompt, - }; - } else if (apiContent.type === 'contextLengthExceeded') { - return { - type: 'contextLengthExceeded', - msg: apiContent.msg, - }; - } else if (apiContent.type === 'summarizationRequested') { - return { - type: 'summarizationRequested', - msg: apiContent.msg, - }; + // Check for errors in the result + if (result.error) { + throw new Error(`Context management failed: ${result.error}`); } - // For types that exist in API but not in frontend, either skip or convert - console.warn(`Skipping unsupported content type: ${apiContent.type}`); - return null; -} - -export function createSummarizationRequestMessage( - messages: FrontendMessage[], - requestMessage: string -): FrontendMessage { - // Get the last message - const lastMessage = messages[messages.length - 1]; - - // Determine the next role (opposite of the last message) - const nextRole: Role = lastMessage.role === 'user' ? 'assistant' : 'user'; - - // Create the new message with SummarizationRequestedContent - return { - id: generateId(), - role: nextRole, - created: Math.floor(Date.now() / 1000), - content: [ - { - type: 'summarizationRequested', - msg: requestMessage, - }, - ], - }; + if (!result.data) { + throw new Error('Context management returned no data'); + } + + return result.data; } diff --git a/ui/desktop/src/components/sessions/SessionHistoryView.tsx b/ui/desktop/src/components/sessions/SessionHistoryView.tsx index 781e42d841..4d297ae9d1 100644 --- a/ui/desktop/src/components/sessions/SessionHistoryView.tsx +++ b/ui/desktop/src/components/sessions/SessionHistoryView.tsx @@ -29,11 +29,9 @@ import { import ProgressiveMessageList from '../ProgressiveMessageList'; import { SearchView } from '../conversation/SearchView'; import { ContextManagerProvider } from '../context_management/ContextManager'; -import { Message } from '../../types/message'; import BackButton from '../ui/BackButton'; import { Tooltip, TooltipContent, TooltipTrigger } from '../ui/Tooltip'; -import { Session } from '../../api'; -import { convertApiMessageToFrontendMessage } from '../context_management'; +import { Message, Session } from '../../api'; // Helper function to determine if a message is a user message (same as useChatEngine) const isUserMessage = (message: Message): boolean => { @@ -153,7 +151,7 @@ const SessionHistoryView: React.FC = ({ const [isCopied, setIsCopied] = useState(false); const [canShare, setCanShare] = useState(false); - const messages = (session.conversation || []).map(convertApiMessageToFrontendMessage); + const messages = session.conversation || []; useEffect(() => { const savedSessionConfig = localStorage.getItem('session_sharing_config'); diff --git a/ui/desktop/src/components/sessions/SessionViewComponents.tsx b/ui/desktop/src/components/sessions/SessionViewComponents.tsx index efd217ca05..e9a6fe450a 100644 --- a/ui/desktop/src/components/sessions/SessionViewComponents.tsx +++ b/ui/desktop/src/components/sessions/SessionViewComponents.tsx @@ -8,13 +8,13 @@ import MarkdownContent from '../MarkdownContent'; import ToolCallWithResponse from '../ToolCallWithResponse'; import ImagePreview from '../ImagePreview'; import { + getTextContent, ToolRequestMessageContent, ToolResponseMessageContent, - TextContent, } from '../../types/message'; -import { type Message } from '../../types/message'; import { formatMessageTimestamp } from '../../utils/timeUtils'; import { extractImagePaths, removeImagePathsFromText } from '../../utils/imageUtils'; +import { Message } from '../../api'; /** * Get tool responses map from messages @@ -111,12 +111,7 @@ export const SessionMessages: React.FC = ({ ) : messages?.length > 0 ? ( messages .map((message, index) => { - // Extract text content from the message - let textContent = message.content - .filter((c): c is TextContent => c.type === 'text') - .map((c) => c.text) - .join('\n'); - + const textContent = getTextContent(message); // Extract image paths from the message const imagePaths = extractImagePaths(textContent); diff --git a/ui/desktop/src/hooks/useAgent.ts b/ui/desktop/src/hooks/useAgent.ts index dc588c1fa1..ecffcee8a9 100644 --- a/ui/desktop/src/hooks/useAgent.ts +++ b/ui/desktop/src/hooks/useAgent.ts @@ -6,7 +6,6 @@ import { initializeCostDatabase } from '../utils/costDatabase'; import { backupConfig, initConfig, - Message as ApiMessage, readAllConfig, Recipe, recoverConfig, @@ -15,7 +14,6 @@ import { validateConfig, } from '../api'; import { COST_TRACKING_ENABLED } from '../updates'; -import { convertApiMessageToFrontendMessage } from '../components/context_management'; export enum AgentState { UNINITIALIZED = 'uninitialized', @@ -78,9 +76,7 @@ export function useAgent(): UseAgentReturn { sessionId: agentSession.id, title: agentSession.recipe?.title || agentSession.description, messageHistoryIndex: 0, - messages: messages?.map((message: ApiMessage) => - convertApiMessageToFrontendMessage(message) - ), + messages, recipe: agentSession.recipe, recipeParameters: agentSession.user_recipe_values || null, }; @@ -160,13 +156,7 @@ export function useAgent(): UseAgentReturn { const conversation = agentSession.conversation || []; // If we're loading a recipe from initContext (new recipe load), start with empty messages // Otherwise, use the messages from the session - const messages = - initContext.recipe && !initContext.resumeSessionId - ? [] - : conversation.map((message: ApiMessage) => - convertApiMessageToFrontendMessage(message) - ); - + const messages = initContext.recipe && !initContext.resumeSessionId ? [] : conversation; let initChat: ChatType = { sessionId: agentSession.id, title: agentSession.recipe?.title || agentSession.description, diff --git a/ui/desktop/src/hooks/useChatEngine.test.ts b/ui/desktop/src/hooks/useChatEngine.test.ts index 068ebaeb05..198b37625e 100644 --- a/ui/desktop/src/hooks/useChatEngine.test.ts +++ b/ui/desktop/src/hooks/useChatEngine.test.ts @@ -1,12 +1,13 @@ /** * @vitest-environment jsdom */ -import { renderHook, act } from '@testing-library/react'; -import { describe, it, expect, vi, beforeEach } from 'vitest'; -import { useChatEngine } from './useChatEngine'; -import { Message, getTextContent } from '../types/message'; -import { ChatType } from '../types/chat'; +import { act, renderHook } from '@testing-library/react'; import type { Mock } from 'vitest'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { useChatEngine } from './useChatEngine'; +import { getTextContent } from '../types/message'; +import { Message } from '../api'; +import { ChatType } from '../types/chat'; // Mock the useMessageStream hook which is a dependency of useChatEngine vi.mock('./useMessageStream', () => ({ diff --git a/ui/desktop/src/hooks/useChatEngine.ts b/ui/desktop/src/hooks/useChatEngine.ts index f0dd7fc381..6e99a519ad 100644 --- a/ui/desktop/src/hooks/useChatEngine.ts +++ b/ui/desktop/src/hooks/useChatEngine.ts @@ -2,20 +2,10 @@ import { useCallback, useEffect, useMemo, useState } from 'react'; import { getApiUrl } from '../config'; import { useMessageStream } from './useMessageStream'; import { LocalMessageStorage } from '../utils/localMessageStorage'; -import { - Message, - createUserMessage, - ToolCall, - ToolCallResult, - ToolRequestMessageContent, - ToolResponseMessageContent, - ToolConfirmationRequestMessageContent, - getTextContent, - TextContent, -} from '../types/message'; +import { createUserMessage, getTextContent, ToolResponseMessageContent } from '../types/message'; +import { getSession, Message } from '../api'; import { ChatType } from '../types/chat'; import { ChatState } from '../types/chatState'; -import { getSession } from '../api'; // Helper function to determine if a message is a user message const isUserMessage = (message: Message): boolean => { @@ -308,11 +298,7 @@ export const useChatEngine = ({ // isUserMessage also checks if the message is a toolConfirmationRequest // check if the last message is a real user's message if (lastMessage && isUserMessage(lastMessage) && !isToolResponse) { - // Get the text content from the last message before removing it - const textContent = lastMessage.content.find((c): c is TextContent => c.type === 'text'); - const textValue = textContent?.text || ''; - - // Set the text back to the input field + const textValue = getTextContent(lastMessage); _setInput(textValue); // Also add to local storage history as a backup so cmd+up can retrieve it @@ -327,19 +313,15 @@ export const useChatEngine = ({ setMessages([]); } } else if (!isUserMessage(lastMessage)) { - // the last message was an assistant message - // check if we have any tool requests or tool confirmation requests - const toolRequests: [string, ToolCallResult][] = lastMessage.content + const toolRequests: [string, Record][] = lastMessage.content .filter( - (content): content is ToolRequestMessageContent | ToolConfirmationRequestMessageContent => - content.type === 'toolRequest' || content.type === 'toolConfirmationRequest' + (content) => content.type === 'toolRequest' || content.type === 'toolConfirmationRequest' ) .map((content) => { if (content.type === 'toolRequest') { return [content.id, content.toolCall]; } else { - // extract tool call from confirmation - const toolCall: ToolCallResult = { + const toolCall = { status: 'success', value: { name: content.toolName, @@ -391,8 +373,7 @@ export const useChatEngine = ({ return filteredMessages .reduce((history, message) => { if (isUserMessage(message)) { - const textContent = message.content.find((c): c is TextContent => c.type === 'text'); - const text = textContent?.text?.trim(); + const text = getTextContent(message).trim(); if (text) { history.push(text); } diff --git a/ui/desktop/src/hooks/useChatStream.ts b/ui/desktop/src/hooks/useChatStream.ts new file mode 100644 index 0000000000..7d5c321c55 --- /dev/null +++ b/ui/desktop/src/hooks/useChatStream.ts @@ -0,0 +1,114 @@ +import { useState, useCallback, useRef } from 'react'; +import { ChatState } from '../types/chatState'; +import { Message } from '../api'; + +const TextDecoder = globalThis.TextDecoder; + +interface UseChatStreamProps { + sessionId: string; + messages: Message[]; + setMessages: (messages: Message[]) => void; + onStreamFinish?: () => void; +} + +export function useChatStream({ + sessionId, + messages, + setMessages, + onStreamFinish, +}: UseChatStreamProps) { + const [chatState, setChatState] = useState(ChatState.Idle); + const abortControllerRef = useRef(null); + + const handleSubmit = useCallback( + async (userMessage: string) => { + const newMessage: Message = { + role: 'user', + content: [{ type: 'text', text: userMessage }], + created: Date.now(), + }; + + const updatedMessages = [...messages, newMessage]; + setMessages(updatedMessages); + setChatState(ChatState.Streaming); + + abortControllerRef.current = new AbortController(); + + try { + const response = await fetch('/reply', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + session_id: sessionId, + messages: updatedMessages.map((m) => ({ + role: m.role, + content: m.content, + })), + }), + signal: abortControllerRef.current.signal, + }); + + if (!response.ok) throw new Error(`HTTP ${response.status}`); + if (!response.body) throw new Error('No response body'); + + const reader = response.body.getReader(); + const decoder = new TextDecoder(); + + while (true) { + const { done, value } = await reader.read(); + if (done) break; + + const chunk = decoder.decode(value); + const lines = chunk.split('\n'); + + for (const line of lines) { + if (!line.startsWith('data: ')) continue; + + const data = line.slice(6); + if (data === '[DONE]') continue; + + try { + const event = JSON.parse(data); + + if (event.message) { + const msg = event.message as Message; + setMessages([...updatedMessages, msg]); + } + + if (event.error) { + console.error('Stream error:', event.error); + setChatState(ChatState.Idle); + return; + } + + if (event.finish) { + setChatState(ChatState.Idle); + onStreamFinish?.(); + return; + } + } catch (e) { + console.error('Failed to parse SSE:', e); + } + } + } + } catch (error: any) { + if (error.name !== 'AbortError') { + console.error('Stream error:', error); + } + setChatState(ChatState.Idle); + } + }, + [sessionId, messages, setMessages, onStreamFinish] + ); + + const stopStreaming = useCallback(() => { + abortControllerRef.current?.abort(); + setChatState(ChatState.Idle); + }, []); + + return { + chatState, + handleSubmit, + stopStreaming, + }; +} diff --git a/ui/desktop/src/hooks/useMessageStream.ts b/ui/desktop/src/hooks/useMessageStream.ts index 5058b7b231..1b3781839e 100644 --- a/ui/desktop/src/hooks/useMessageStream.ts +++ b/ui/desktop/src/hooks/useMessageStream.ts @@ -1,6 +1,8 @@ import { useCallback, useEffect, useId, useReducer, useRef, useState } from 'react'; import useSWR from 'swr'; -import { createUserMessage, hasCompletedToolCalls, Message, Role } from '../types/message'; +import { createUserMessage, hasCompletedToolCalls } from '../types/message'; +import { Message, Role } from '../api'; + import { getSession, Session } from '../api'; import { ChatState } from '../types/chatState'; diff --git a/ui/desktop/src/hooks/useRecipeManager.ts b/ui/desktop/src/hooks/useRecipeManager.ts index e95d613073..695e775157 100644 --- a/ui/desktop/src/hooks/useRecipeManager.ts +++ b/ui/desktop/src/hooks/useRecipeManager.ts @@ -1,6 +1,8 @@ import { useEffect, useMemo, useState, useRef } from 'react'; import { Recipe, scanRecipe } from '../recipe'; -import { Message, createUserMessage } from '../types/message'; +import { createUserMessage } from '../types/message'; +import { Message } from '../api'; + import { updateSystemPromptWithParameters, substituteParameters, diff --git a/ui/desktop/src/sharedSessions.ts b/ui/desktop/src/sharedSessions.ts index 2dd469aa2a..4e6b7184ca 100644 --- a/ui/desktop/src/sharedSessions.ts +++ b/ui/desktop/src/sharedSessions.ts @@ -1,5 +1,5 @@ -import { Message } from './types/message'; import { safeJsonParse } from './utils/jsonUtils'; +import { Message } from './api'; export interface SharedSessionDetails { share_token: string; diff --git a/ui/desktop/src/types/chat.ts b/ui/desktop/src/types/chat.ts index 006e943ffb..70130ebcf9 100644 --- a/ui/desktop/src/types/chat.ts +++ b/ui/desktop/src/types/chat.ts @@ -1,5 +1,5 @@ -import { Message } from './message'; import { Recipe } from '../recipe'; +import { Message } from '../api'; export interface ChatType { sessionId: string; diff --git a/ui/desktop/src/types/message.ts b/ui/desktop/src/types/message.ts index 578636d9c1..c5273c10c9 100644 --- a/ui/desktop/src/types/message.ts +++ b/ui/desktop/src/types/message.ts @@ -1,116 +1,8 @@ -/** - * Message types that match the Rust message structures - * for direct serialization between client and server - */ +import { Content, Message, ToolConfirmationRequest, ToolRequest, ToolResponse } from '../api'; -export type Role = 'user' | 'assistant'; +export type ToolRequestMessageContent = ToolRequest & { type: 'toolRequest' }; +export type ToolResponseMessageContent = ToolResponse & { type: 'toolResponse' }; -export interface TextContent { - type: 'text'; - text: string; - annotations?: Record; -} - -export interface ImageContent { - type: 'image'; - data: string; - mimeType: string; - annotations?: Record; -} - -export interface ResourceContent { - type: 'resource'; - resource: { - uri: string; - mimeType: string; - text?: string; - blob?: string; - }; - annotations?: Record; -} - -export type Content = TextContent | ImageContent | ResourceContent; - -export interface ToolCall { - name: string; - arguments: Record; -} - -export interface ToolCallResult { - status: 'success' | 'error'; - value?: T; - error?: string; -} - -export interface ToolRequest { - id: string; - toolCall: ToolCallResult; -} - -export interface ToolResponse { - id: string; - toolResult: ToolCallResult; -} - -export interface ToolRequestMessageContent { - type: 'toolRequest'; - id: string; - toolCall: ToolCallResult; -} - -export interface ToolResponseMessageContent { - type: 'toolResponse'; - id: string; - toolResult: ToolCallResult; -} - -export interface ToolConfirmationRequestMessageContent { - type: 'toolConfirmationRequest'; - id: string; - toolName: string; - arguments: Record; - prompt?: string; -} - -export interface ExtensionCall { - name: string; - arguments: Record; - extensionName: string; -} - -export interface ExtensionCallResult { - status: 'success' | 'error'; - value?: T; - error?: string; -} - -export interface ContextLengthExceededContent { - type: 'contextLengthExceeded'; - msg: string; -} - -export interface SummarizationRequestedContent { - type: 'summarizationRequested'; - msg: string; -} - -export type MessageContent = - | TextContent - | ImageContent - | ToolRequestMessageContent - | ToolResponseMessageContent - | ToolConfirmationRequestMessageContent - | ContextLengthExceededContent - | SummarizationRequestedContent; - -export interface Message { - id?: string; - role: Role; - created: number; - content: MessageContent[]; -} - -// Helper functions to create messages export function createUserMessage(text: string): Message { return { id: generateId(), @@ -190,56 +82,42 @@ export function createToolErrorResponseMessage(id: string, error: string): Messa }; } -// Generate a unique ID for messages function generateId(): string { return Math.random().toString(36).substring(2, 10); } -// Helper functions to extract content from messages export function getTextContent(message: Message): string { return message.content - .filter( - (content): content is TextContent | ContextLengthExceededContent => - content.type === 'text' || content.type === 'contextLengthExceeded' - ) .map((content) => { - if (content.type === 'text') { - return content.text; - } else if (content.type === 'contextLengthExceeded') { - return content.msg; - } + if (content.type === 'text') return content.text; + if (content.type === 'contextLengthExceeded') return content.msg; return ''; }) .join(''); } -export function getToolRequests(message: Message): ToolRequestMessageContent[] { +export function getToolRequests(message: Message): (ToolRequest & { type: 'toolRequest' })[] { return message.content.filter( - (content): content is ToolRequestMessageContent => content.type === 'toolRequest' + (content): content is ToolRequest & { type: 'toolRequest' } => content.type === 'toolRequest' ); } -export function getToolResponses(message: Message): ToolResponseMessageContent[] { +export function getToolResponses(message: Message): (ToolResponse & { type: 'toolResponse' })[] { return message.content.filter( - (content): content is ToolResponseMessageContent => content.type === 'toolResponse' + (content): content is ToolResponse & { type: 'toolResponse' } => content.type === 'toolResponse' ); } export function getToolConfirmationContent( message: Message -): ToolConfirmationRequestMessageContent | undefined { +): (ToolConfirmationRequest & { type: 'toolConfirmationRequest' }) | undefined { return message.content.find( - (content): content is ToolConfirmationRequestMessageContent => + (content): content is ToolConfirmationRequest & { type: 'toolConfirmationRequest' } => content.type === 'toolConfirmationRequest' ); } export function hasCompletedToolCalls(message: Message): boolean { const toolRequests = getToolRequests(message); - if (toolRequests.length === 0) return false; - - // For now, we'll assume all tool calls are completed when this is checked - // In a real implementation, you'd need to check if all tool requests have responses - // by looking through subsequent messages - return true; + return toolRequests.length > 0; } diff --git a/ui/desktop/src/utils/timeUtils.ts b/ui/desktop/src/utils/timeUtils.ts index 2625e796d8..c7a9abb5ed 100644 --- a/ui/desktop/src/utils/timeUtils.ts +++ b/ui/desktop/src/utils/timeUtils.ts @@ -1,6 +1,5 @@ -export function formatMessageTimestamp(timestamp: number): string { - // Convert from Unix timestamp (seconds) to milliseconds - const date = new Date(timestamp * 1000); +export function formatMessageTimestamp(timestamp?: number): string { + const date = timestamp ? new Date(timestamp * 1000) : new Date(); const now = new Date(); // Format time as HH:MM AM/PM diff --git a/ui/desktop/src/utils/toolCallChaining.ts b/ui/desktop/src/utils/toolCallChaining.ts index 80f36383cb..7e5715f642 100644 --- a/ui/desktop/src/utils/toolCallChaining.ts +++ b/ui/desktop/src/utils/toolCallChaining.ts @@ -1,4 +1,5 @@ -import { Message, getToolRequests, getTextContent, getToolResponses } from '../types/message'; +import { getToolRequests, getTextContent, getToolResponses } from '../types/message'; +import { Message } from '../api'; export function identifyConsecutiveToolCalls(messages: Message[]): number[][] { const chains: number[][] = [];