Change it all

This commit is contained in:
Douwe Osinga
2025-10-07 21:06:54 -04:00
parent e14fa24ee7
commit be4567ecfa
31 changed files with 264 additions and 450 deletions
+23 -10
View File
@@ -39,28 +39,44 @@ where
Ok(content)
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, ToSchema)]
#[serde(tag = "status", rename_all = "lowercase")]
pub enum ToolCallResult<T> {
Success { value: T },
Error { error: String },
}
impl<T> From<ToolResult<T>> for ToolCallResult<T> {
fn from(result: ToolResult<T>) -> 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<CallToolRequestParam>,
pub tool_call: ToolCallResult<CallToolRequestParam>,
}
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(|_| "<<invalid json>>".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<Vec<Content>>,
}
@@ -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<CallToolRequestParam>,
}
+1 -1
View File
@@ -1,4 +1,4 @@
pub use rmcp::model::ErrorData;
/// Type alias for tool results
pub type ToolResult<T> = std::result::Result<T, ErrorData>;
pub type ToolResult<T> = Result<T, ErrorData>;
+3 -3
View File
@@ -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"
}
}
},
+1 -1
View File
@@ -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);
+13 -1
View File
@@ -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) => {};
+1 -1
View File
@@ -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';
+1 -1
View File
@@ -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';
@@ -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: {
@@ -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';
+2 -1
View File
@@ -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';
@@ -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({
@@ -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<string, unknown> };
if (!toolCall) {
return null;
}
@@ -53,11 +54,12 @@ export default function ToolCallWithResponse({
/>
</div>
{/* 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 (
<div key={`${content.type}-${index}`} className="mt-3">
<div key={`${index}`} className="mt-3">
<MCPUIResourceRenderer content={content} appendPromptToChat={append} />
<div className="mt-3 p-4 py-3 border border-borderSubtle rounded-lg bg-background-muted flex items-center">
<FlaskConical className="mr-2" size={20} />
@@ -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<number | null>(null);
@@ -547,30 +549,34 @@ interface ToolResultViewProps {
}
function ToolResultView({ result, isStartExpanded }: ToolResultViewProps) {
const resultAny = result as any;
return (
<ToolCallExpandable
label={<span className="pl-4 py-1 font-sans text-sm">Output</span>}
isStartExpanded={isStartExpanded}
>
<div className="pl-4 pr-4 py-4">
{result.type === 'text' && result.text && (
{'text' in resultAny && resultAny.text && (
<MarkdownContent
content={result.text}
content={resultAny.text}
className="whitespace-pre-wrap max-w-full overflow-x-auto"
/>
)}
{result.type === 'image' && (
<img
src={`data:${result.mimeType};base64,${result.data}`}
alt="Tool result"
className="max-w-full h-auto rounded-md my-2"
onError={(e) => {
console.error('Failed to load image');
e.currentTarget.style.display = 'none';
}}
/>
)}
{result.type === 'resource' && (
{'mimeType' in resultAny &&
'data' in resultAny &&
resultAny.mimeType?.startsWith('image') && (
<img
src={`data:${resultAny.mimeType};base64,${resultAny.data}`}
alt="Tool result"
className="max-w-full h-auto rounded-md my-2"
onError={(e) => {
console.error('Failed to load image');
e.currentTarget.style.display = 'none';
}}
/>
)}
{'resource' in resultAny && (
<pre className="font-sans text-sm">{JSON.stringify(result, null, 2)}</pre>
)}
</div>
+3 -2
View File
@@ -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';
@@ -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<CompactionMarkerProps> = ({ 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';
@@ -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);
@@ -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', () => {
@@ -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 () => {
@@ -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<ContextManageResponse> {
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<string, unknown> | undefined,
};
} else if (apiContent.type === 'image') {
return {
type: 'image',
data: apiContent.data,
mimeType: apiContent.mimeType,
annotations: apiContent.annotations as Record<string, unknown> | undefined,
};
} else if (apiContent.type === 'toolRequest') {
// Ensure the toolCall has the correct type structure
const toolCall = apiContent.toolCall as unknown as ToolCallResult<ToolCall>;
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<FrontendContent[]>;
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<string, unknown>,
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;
}
@@ -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<SessionHistoryViewProps> = ({
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');
@@ -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<SessionMessagesProps> = ({
) : 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);
+2 -12
View File
@@ -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,
+6 -5
View File
@@ -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', () => ({
+7 -26
View File
@@ -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<ToolCall>][] = lastMessage.content
const toolRequests: [string, Record<string, unknown>][] = 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<ToolCall> = {
const toolCall = {
status: 'success',
value: {
name: content.toolName,
@@ -391,8 +373,7 @@ export const useChatEngine = ({
return filteredMessages
.reduce<string[]>((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);
}
+114
View File
@@ -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>(ChatState.Idle);
const abortControllerRef = useRef<AbortController | null>(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,
};
}
+3 -1
View File
@@ -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';
+3 -1
View File
@@ -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,
+1 -1
View File
@@ -1,5 +1,5 @@
import { Message } from './types/message';
import { safeJsonParse } from './utils/jsonUtils';
import { Message } from './api';
export interface SharedSessionDetails {
share_token: string;
+1 -1
View File
@@ -1,5 +1,5 @@
import { Message } from './message';
import { Recipe } from '../recipe';
import { Message } from '../api';
export interface ChatType {
sessionId: string;
+12 -134
View File
@@ -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<string, unknown>;
}
export interface ImageContent {
type: 'image';
data: string;
mimeType: string;
annotations?: Record<string, unknown>;
}
export interface ResourceContent {
type: 'resource';
resource: {
uri: string;
mimeType: string;
text?: string;
blob?: string;
};
annotations?: Record<string, unknown>;
}
export type Content = TextContent | ImageContent | ResourceContent;
export interface ToolCall {
name: string;
arguments: Record<string, unknown>;
}
export interface ToolCallResult<T> {
status: 'success' | 'error';
value?: T;
error?: string;
}
export interface ToolRequest {
id: string;
toolCall: ToolCallResult<ToolCall>;
}
export interface ToolResponse {
id: string;
toolResult: ToolCallResult<Content[]>;
}
export interface ToolRequestMessageContent {
type: 'toolRequest';
id: string;
toolCall: ToolCallResult<ToolCall>;
}
export interface ToolResponseMessageContent {
type: 'toolResponse';
id: string;
toolResult: ToolCallResult<Content[]>;
}
export interface ToolConfirmationRequestMessageContent {
type: 'toolConfirmationRequest';
id: string;
toolName: string;
arguments: Record<string, unknown>;
prompt?: string;
}
export interface ExtensionCall {
name: string;
arguments: Record<string, unknown>;
extensionName: string;
}
export interface ExtensionCallResult<T> {
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;
}
+2 -3
View File
@@ -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
+2 -1
View File
@@ -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[][] = [];