diff --git a/apps/studio/components/interfaces/Linter/LintDetail.tsx b/apps/studio/components/interfaces/Linter/LintDetail.tsx index 334353bd703..a88d3bb6b45 100644 --- a/apps/studio/components/interfaces/Linter/LintDetail.tsx +++ b/apps/studio/components/interfaces/Linter/LintDetail.tsx @@ -2,15 +2,15 @@ import Link from 'next/link' import ReactMarkdown from 'react-markdown' import { createLintSummaryPrompt, lintInfoMap } from 'components/interfaces/Linter/Linter.utils' +import { SIDEBAR_KEYS } from 'components/layouts/ProjectLayout/LayoutSidebar/LayoutSidebarProvider' import { Lint } from 'data/lint/lint-query' import { DOCS_URL } from 'lib/constants' import { useTrack } from 'lib/telemetry/track' import { ExternalLink } from 'lucide-react' import { useAiAssistantStateSnapshot } from 'state/ai-assistant-state' +import { useSidebarManagerSnapshot } from 'state/sidebar-manager-state' import { AiIconAnimation, Button } from 'ui' import { EntityTypeIcon, LintCTA, LintEntity } from './Linter.utils' -import { SIDEBAR_KEYS } from 'components/layouts/ProjectLayout/LayoutSidebar/LayoutSidebarProvider' -import { useSidebarManagerSnapshot } from 'state/sidebar-manager-state' interface LintDetailProps { lint: Lint @@ -35,7 +35,7 @@ const LintDetail = ({ lint, projectRef, onAskAssistant }: LintDetailProps) => { openSidebar(SIDEBAR_KEYS.AI_ASSISTANT) snap.newChat({ name: 'Summarize lint', - initialInput: createLintSummaryPrompt(lint), + initialMessage: createLintSummaryPrompt(lint), }) } diff --git a/apps/studio/components/ui/AIAssistantPanel/AIAssistant.tsx b/apps/studio/components/ui/AIAssistantPanel/AIAssistant.tsx index c075f57000b..3f353fb3712 100644 --- a/apps/studio/components/ui/AIAssistantPanel/AIAssistant.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/AIAssistant.tsx @@ -1,6 +1,6 @@ import type { UIMessage as MessageType } from '@ai-sdk/react' import { useChat } from '@ai-sdk/react' -import { DefaultChatTransport, lastAssistantMessageIsCompleteWithToolCalls } from 'ai' +import { lastAssistantMessageIsCompleteWithToolCalls } from 'ai' import { AnimatePresence, motion } from 'framer-motion' import { Eraser, Info, Pencil, X } from 'lucide-react' import { useRouter } from 'next/router' @@ -12,7 +12,6 @@ import { Markdown } from 'components/interfaces/Markdown' import { SIDEBAR_KEYS } from 'components/layouts/ProjectLayout/LayoutSidebar/LayoutSidebarProvider' import { useCheckOpenAIKeyQuery } from 'data/ai/check-api-key-query' import { useRateMessageMutation } from 'data/ai/rate-message-mutation' -import { constructHeaders } from 'data/fetchers' import { useTablesQuery } from 'data/tables/tables-query' import { useSendEventMutation } from 'data/telemetry/send-event-mutation' import { useLocalStorageQuery } from 'hooks/misc/useLocalStorage' @@ -20,11 +19,10 @@ import { useOrgAiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' import { useSelectedOrganizationQuery } from 'hooks/misc/useSelectedOrganization' import { useSelectedProjectQuery } from 'hooks/misc/useSelectedProject' import { useHotKey } from 'hooks/ui/useHotKey' -import { prepareMessagesForAPI } from 'lib/ai/message-utils' -import { BASE_PATH, IS_PLATFORM } from 'lib/constants' +import { IS_PLATFORM } from 'lib/constants' import { uuidv4 } from 'lib/helpers' import type { AssistantModel } from 'state/ai-assistant-state' -import { useAiAssistantStateSnapshot } from 'state/ai-assistant-state' +import { useAiAssistantState, useAiAssistantStateSnapshot } from 'state/ai-assistant-state' import { useSidebarManagerSnapshot } from 'state/sidebar-manager-state' import { useSqlEditorV2StateSnapshot } from 'state/sql-editor-v2' import { Button, cn, KeyboardShortcut } from 'ui' @@ -61,6 +59,7 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { const disablePrompts = useFlag('disableAssistantPrompts') const { snippets } = useSqlEditorV2StateSnapshot() const snap = useAiAssistantStateSnapshot() + const state = useAiAssistantState() const { closeSidebar, activeSidebar } = useSidebarManagerSnapshot() const isPaidPlan = selectedOrganization?.plan?.id !== 'free' @@ -125,41 +124,16 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { const currentSchema = searchParams?.get('schema') ?? 'public' const currentChat = snap.activeChat?.name - const { mutate: sendEvent } = useSendEventMutation() - - const updateMessage = useCallback( - (updatedMessage: MessageType) => { - snap.updateMessage(updatedMessage) - }, - [snap] - ) - - // Handle completion of the assistant's response - const handleChatFinish = useCallback( - ({ message }: { message: MessageType }) => { - if (lastUserMessageRef.current) { - snap.saveMessage([lastUserMessageRef.current, message]) - lastUserMessageRef.current = null - } else { - updateMessage(message) - } - }, - [snap, updateMessage] - ) - - // TODO(refactor): This useChat hook should be moved down into each chat session. - // That way we won't have to disable switching chats while the chat is loading, - // and don't run the risk of messages getting mixed up between chats. - // Sanitize messages to remove Valtio proxy wrappers that can't be cloned - const sanitizedMessages = useMemo(() => { - if (!snap.activeChat?.messages) return undefined - - return snap.activeChat.messages.map((msg: any) => { - // Convert proxy objects to plain objects - const plainMessage = JSON.parse(JSON.stringify(msg)) - return plainMessage + // Update context in state + useEffect(() => { + state.setContext({ + projectRef: project?.ref, + orgSlug: selectedOrganizationRef.current?.slug, + connectionString: project?.connectionString || undefined, }) - }, [snap.activeChat?.messages]) + }, [project?.ref, project?.connectionString, selectedOrganizationRef.current?.slug, state]) + + const { mutate: sendEvent } = useSendEventMutation() const { messages: chatMessages, @@ -172,60 +146,11 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { regenerate, } = useChat({ id: snap.activeChatId, + ...(snap.activeChatId && snap.chatInstances[snap.activeChatId] + ? { chat: snap.chatInstances[snap.activeChatId] } + : {}), sendAutomaticallyWhen: lastAssistantMessageIsCompleteWithToolCalls, - messages: sanitizedMessages, - async onToolCall({ toolCall }) { - if (toolCall.dynamic) { - return - } - - if (toolCall.toolName === 'rename_chat') { - const { newName } = toolCall.input as { newName: string } - - if (snap.activeChatId && newName?.trim()) { - snap.renameChat(snap.activeChatId, newName.trim()) - - addToolResult({ - tool: toolCall.toolName, - toolCallId: toolCall.toolCallId, - output: 'Chat renamed', - }) - } else { - addToolResult({ - tool: toolCall.toolName, - toolCallId: toolCall.toolCallId, - output: 'Failed to rename chat: Invalid chat or name', - }) - } - } - }, - transport: new DefaultChatTransport({ - api: `${BASE_PATH}/api/ai/sql/generate-v4`, - async prepareSendMessagesRequest({ messages, ...options }) { - const cleanedMessages = prepareMessagesForAPI(messages) - - const headerData = await constructHeaders() - const authorizationHeader = headerData.get('Authorization') - - return { - ...options, - body: { - messages: cleanedMessages, - aiOptInLevel, - projectRef: project?.ref, - connectionString: project?.connectionString, - schema: currentSchema, - table: currentTable?.name, - chatName: currentChat, - orgSlug: selectedOrganizationRef.current?.slug, - model: selectedModel, - }, - ...(IS_PLATFORM ? { headers: { Authorization: authorizationHeader ?? '' } } : {}), - } - }, - }), onError: onErrorChat, - onFinish: handleChatFinish, }) const isChatLoading = chatStatus === 'submitted' || chatStatus === 'streaming' @@ -389,7 +314,12 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { snap.clearSqlSnippets() lastUserMessageRef.current = payload - sendMessage(payload) + sendMessage(payload, { + body: { + schema: currentSchema, + table: currentTable?.name, + }, + }) setValue('') if (finalContent.includes('Help me to debug')) { @@ -412,6 +342,7 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { } const handleClearMessages = () => { + if (isChatLoading) stop() snap.clearMessages() setMessages([]) lastUserMessageRef.current = null @@ -649,7 +580,7 @@ export const AIAssistant = ({ className }: AIAssistantProps) => { // to save partial responses from the AI const lastMessage = chatMessages[chatMessages.length - 1] if (lastMessage && lastMessage.role === 'assistant') { - handleChatFinish({ message: lastMessage }) + state.updateMessage(lastMessage) } }} sqlSnippets={snap.sqlSnippets as SqlSnippet[] | undefined} diff --git a/apps/studio/components/ui/AIAssistantPanel/AIAssistantHeader.tsx b/apps/studio/components/ui/AIAssistantPanel/AIAssistantHeader.tsx index 6299e05b72f..76724362782 100644 --- a/apps/studio/components/ui/AIAssistantPanel/AIAssistantHeader.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/AIAssistantHeader.tsx @@ -47,7 +47,7 @@ export const AIAssistantHeader = ({ - +
@@ -57,7 +57,6 @@ export const AIAssistantHeader = ({ icon={} onClick={onNewChat} className="h-7 w-7 p-0" - disabled={isChatLoading} tooltip={{ content: { side: 'bottom', text: 'New chat' } }} /> activeChatId?: string model: AssistantModel + context: AiAssistantContext } // Data structure stored in IndexedDB @@ -53,6 +65,7 @@ const INITIAL_AI_ASSISTANT: AiAssistantData = { chats: {}, activeChatId: undefined, model: 'gpt-5', + context: {}, } const DB_NAME = 'ai-assistant-db' @@ -211,12 +224,79 @@ function ensureActiveChatOrInitialize(state: AiAssistantState) { } } +function createChatInstance( + state: AiAssistantState, + options: { id: string; initialMessages: MessageType[] } +) { + return new Chat({ + id: options.id, + messages: options.initialMessages, + sendAutomaticallyWhen: lastAssistantMessageIsCompleteWithToolCalls, + transport: new DefaultChatTransport({ + api: `${BASE_PATH}/api/ai/sql/generate-v4`, + async prepareSendMessagesRequest({ messages, ...opts }) { + const cleanedMessages = prepareMessagesForAPI(messages) + const headerData = await constructHeaders() + const authorizationHeader = headerData.get('Authorization') + + // Get the chat specific to this request to ensure we have the correct name + const chat = state.chats[options.id] + + return { + ...opts, + body: { + messages: cleanedMessages, + projectRef: state.context.projectRef, + connectionString: state.context.connectionString, + chatName: chat?.name, + orgSlug: state.context.orgSlug, + context: state.context, + model: state.model, + ...opts.body, + }, + ...(IS_PLATFORM ? { headers: { Authorization: authorizationHeader ?? '' } } : {}), + } + }, + }), + async onToolCall({ toolCall }) { + if (toolCall.dynamic) { + return + } + + if (toolCall.toolName === 'rename_chat') { + const { newName } = toolCall.input as { newName: string } + + if (options.id && newName?.trim()) { + state.renameChat(options.id, newName.trim()) + } + } + }, + onFinish(result) { + // Sync messages back to state + const chatInstance = state.chatInstances[options.id] + if (chatInstance) { + const messages = chatInstance.messages + const chat = state.chats[options.id] + if (chat) { + chat.messages = messages as AssistantMessageType[] + chat.updatedAt = new Date() + } + } + }, + }) +} + export const createAiAssistantState = (): AiAssistantState => { // Initialize with defaults, loading happens asynchronously in the provider const initialState = { ...INITIAL_AI_ASSISTANT } const state: AiAssistantState = proxy({ ...initialState, // Spread initial values directly + chatInstances: {}, + + setContext: (context: Partial) => { + state.context = { ...state.context, ...context } + }, resetAiAssistantPanel: () => { Object.assign(state, INITIAL_AI_ASSISTANT) @@ -232,7 +312,7 @@ export const createAiAssistantState = (): AiAssistantState => { }, newChat: ( - options?: { name?: string } & Partial< + options?: { name?: string; initialMessage?: string } & Partial< Pick > ) => { @@ -251,6 +331,18 @@ export const createAiAssistantState = (): AiAssistantState => { } state.activeChatId = chatId + // Create new chat instance + const chatInstance = createChatInstance(state, { id: chatId, initialMessages: [] }) + + state.chatInstances[chatId] = ref(chatInstance) + + // If initialMessage is provided, append it to the chat instance + if (options?.initialMessage) { + chatInstance.sendMessage({ + text: options.initialMessage, + }) + } + // Update non-chat related state based on options, falling back to current state, then initial state.initialInput = options?.initialInput ?? INITIAL_AI_ASSISTANT.initialInput state.sqlSnippets = options?.sqlSnippets ?? INITIAL_AI_ASSISTANT.sqlSnippets @@ -263,6 +355,14 @@ export const createAiAssistantState = (): AiAssistantState => { selectChat: (id: string) => { if (id !== state.activeChatId) { state.activeChatId = id + const chat = state.chats[id] + if (chat) { + if (!state.chatInstances[id]) { + state.chatInstances[id] = ref( + createChatInstance(state, { id, initialMessages: chat.messages }) + ) + } + } } }, @@ -273,6 +373,15 @@ export const createAiAssistantState = (): AiAssistantState => { if (id === state.activeChatId) { const remainingChatIds = Object.keys(remainingChats) state.activeChatId = remainingChatIds.length > 0 ? remainingChatIds[0] : undefined + + if (state.activeChatId) { + const chat = state.chats[state.activeChatId] + if (!state.chatInstances[state.activeChatId]) { + state.chatInstances[state.activeChatId] = ref( + createChatInstance(state, { id: state.activeChatId, initialMessages: chat.messages }) + ) + } + } } }, @@ -382,6 +491,20 @@ export const createAiAssistantState = (): AiAssistantState => { state.newChat() } } + + // Initialize chat instance for the active chat + if ( + state.activeChatId && + state.chats[state.activeChatId] && + !state.chatInstances[state.activeChatId] + ) { + state.chatInstances[state.activeChatId] = ref( + createChatInstance(state, { + id: state.activeChatId, + initialMessages: state.chats[state.activeChatId].messages, + }) + ) + } }, clearStorage: async () => { @@ -395,9 +518,11 @@ export const createAiAssistantState = (): AiAssistantState => { export type AiAssistantState = AiAssistantData & { resetAiAssistantPanel: () => void activeChat: ChatSession | undefined + chatInstances: Record> + setContext: (context: Partial) => void setModel: (model: AssistantModel) => void newChat: ( - options?: { name?: string } & Partial< + options?: { name?: string; initialMessage?: string } & Partial< Pick > ) => string @@ -508,3 +633,8 @@ export const useAiAssistantStateSnapshot = (options?: Parameters { + const state = useContext(AiAssistantStateContext) + return state +}