mirror of
https://github.com/supabase/supabase.git
synced 2026-10-05 09:25:06 +03:00
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:
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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in new issue
Block a user