Move chat state to project context (#40673)

* global chat manage

* support multiple chats

* multiple chat and initial message

* fix

* fix some issues

* prettier

* prettier

* Update apps/studio/state/ai-assistant-state.tsx

Co-authored-by: Alaister Young <alaister@users.noreply.github.com>

* combine ifs

---------

Co-authored-by: Alaister Young <alaister@users.noreply.github.com>
This commit is contained in:
Saxon FletcherandAlaister Young authored and GitHub committed 2025-11-25 15:43:39 +10:00
1 parent 31a026e6f0
commit cd4d091fc0
4 files changed
+162 -102

No files matched your search

@@ -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),
})
}
@@ -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}
@@ -47,7 +47,7 @@ export const AIAssistantHeader = ({
<path d="M16 3.549L7.12 20.600" />
</svg>
</span>
<AIAssistantChatSelector disabled={isChatLoading} />
<AIAssistantChatSelector />
</div>
<div className="flex items-center gap-x-4">
<div className="flex items-center">
@@ -57,7 +57,6 @@ export const AIAssistantHeader = ({
icon={<Plus strokeWidth={1.5} />}
onClick={onNewChat}
className="h-7 w-7 p-0"
disabled={isChatLoading}
tooltip={{ content: { side: 'bottom', text: 'New chat' } }}
/>
<ButtonTooltip
+134 -4
View File
@@ -1,9 +1,14 @@
import type { UIMessage as MessageType } from '@ai-sdk/react'
import { Chat, type UIMessage as MessageType } from '@ai-sdk/react'
import { DefaultChatTransport, lastAssistantMessageIsCompleteWithToolCalls } from 'ai'
import { DBSchema, IDBPDatabase, openDB } from 'idb'
import { debounce } from 'lodash'
import { createContext, PropsWithChildren, useContext, useEffect, useState } from 'react'
import { v4 as uuidv4 } from 'uuid'
import { proxy, snapshot, subscribe, useSnapshot } from 'valtio'
import { proxy, ref, snapshot, subscribe, useSnapshot } from 'valtio'
import { constructHeaders } from 'data/fetchers'
import { prepareMessagesForAPI } from 'lib/ai/message-utils'
import { BASE_PATH, IS_PLATFORM } from 'lib/constants'
import { LOCAL_STORAGE_KEYS } from 'common'
import { useSelectedProjectQuery } from 'hooks/misc/useSelectedProject'
@@ -27,6 +32,12 @@ type ChatSession = {
updatedAt: Date
}
export type AiAssistantContext = {
projectRef?: string
orgSlug?: string
connectionString?: string
}
type AiAssistantData = {
initialInput: string
sqlSnippets?: SqlSnippet[]
@@ -35,6 +46,7 @@ type AiAssistantData = {
chats: Record<string, ChatSession>
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<MessageType>({
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<AiAssistantContext>) => {
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<AiAssistantData, 'initialInput' | 'sqlSnippets' | 'suggestions' | 'tables'>
>
) => {
@@ -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<string, Chat<MessageType>>
setContext: (context: Partial<AiAssistantContext>) => void
setModel: (model: AssistantModel) => void
newChat: (
options?: { name?: string } & Partial<
options?: { name?: string; initialMessage?: string } & Partial<
Pick<AiAssistantData, 'initialInput' | 'sqlSnippets' | 'suggestions' | 'tables'>
>
) => string
@@ -508,3 +633,8 @@ export const useAiAssistantStateSnapshot = (options?: Parameters<typeof useSnaps
const state = useContext(AiAssistantStateContext)
return useSnapshot(state, options)
}
export const useAiAssistantState = () => {
const state = useContext(AiAssistantStateContext)
return state
}