useAiChat.ts244 lines · main
| 1 | import type OpenAI from 'openai' |
| 2 | import type { Dispatch, SetStateAction } from 'react' |
| 3 | import { useCallback, useReducer, useRef, useState } from 'react' |
| 4 | import { SSE } from 'sse.js' |
| 5 | |
| 6 | import { BASE_PATH } from '../shared/constants' |
| 7 | import type { Message, MessageAction, SourceLink } from './utils' |
| 8 | import { MessageRole, MessageStatus } from './utils' |
| 9 | |
| 10 | export function parseSourcesFromContent(content: string): { |
| 11 | cleanedContent: string |
| 12 | sources: SourceLink[] |
| 13 | } { |
| 14 | // Only match Sources section at the very end of the message |
| 15 | const sourcesMatch = content.match(/### Sources\s*(?:\n((?:- [^\n]+\n?)*))?\s*$/) |
| 16 | |
| 17 | let cleanedContent = content |
| 18 | const sources: SourceLink[] = [] |
| 19 | |
| 20 | if (sourcesMatch) { |
| 21 | // Extract sources |
| 22 | const sourcesText = sourcesMatch[1] || '' |
| 23 | const sourceLines = sourcesText.split('\n').filter((line) => line.trim().startsWith('- ')) |
| 24 | |
| 25 | for (const sourceLine of sourceLines) { |
| 26 | const path = sourceLine.replace(/^- /, '').trim() |
| 27 | // Only include paths that start with '/' |
| 28 | if (path && path.startsWith('/')) { |
| 29 | sources.push({ |
| 30 | path, |
| 31 | url: `https://supabase.com/docs${path}`, |
| 32 | }) |
| 33 | } |
| 34 | } |
| 35 | |
| 36 | // Remove sources section from content |
| 37 | const sourcesIndex = content.lastIndexOf('### Sources') |
| 38 | if (sourcesIndex !== -1) { |
| 39 | cleanedContent = content.substring(0, sourcesIndex).trim() |
| 40 | } |
| 41 | } |
| 42 | |
| 43 | return { cleanedContent, sources } |
| 44 | } |
| 45 | |
| 46 | const messageReducer = (state: Message[], messageAction: MessageAction) => { |
| 47 | let current = [...state] |
| 48 | const { type } = messageAction |
| 49 | |
| 50 | switch (type) { |
| 51 | case 'new': { |
| 52 | const { message } = messageAction |
| 53 | current.push(message) |
| 54 | break |
| 55 | } |
| 56 | case 'update': { |
| 57 | const { index, message } = messageAction |
| 58 | if (current[index]) { |
| 59 | current[index] = Object.assign({}, current[index], message) |
| 60 | } |
| 61 | break |
| 62 | } |
| 63 | case 'append-content': { |
| 64 | const { index, content, idempotencyKey } = messageAction |
| 65 | |
| 66 | const messageToEdit = current[index] |
| 67 | if (!messageToEdit || messageToEdit.idempotencyKey === idempotencyKey) { |
| 68 | break |
| 69 | } |
| 70 | messageToEdit.idempotencyKey = idempotencyKey |
| 71 | |
| 72 | current[index] = Object.assign({}, messageToEdit, { |
| 73 | content: (messageToEdit.content += content), |
| 74 | }) |
| 75 | break |
| 76 | } |
| 77 | case 'finalize-with-sources': { |
| 78 | const { index } = messageAction |
| 79 | const messageToFinalize = current[index] |
| 80 | if (messageToFinalize && messageToFinalize.content) { |
| 81 | const { cleanedContent, sources } = parseSourcesFromContent(messageToFinalize.content) |
| 82 | |
| 83 | current[index] = Object.assign({}, messageToFinalize, { |
| 84 | status: MessageStatus.Complete, |
| 85 | content: cleanedContent, |
| 86 | sources: sources.length > 0 ? sources : undefined, |
| 87 | }) |
| 88 | } else { |
| 89 | current[index] = Object.assign({}, messageToFinalize, { |
| 90 | status: MessageStatus.Complete, |
| 91 | }) |
| 92 | } |
| 93 | break |
| 94 | } |
| 95 | case 'reset': { |
| 96 | current = [] |
| 97 | break |
| 98 | } |
| 99 | default: { |
| 100 | throw new Error(`Unknown message action '${type}'`) |
| 101 | } |
| 102 | } |
| 103 | |
| 104 | return current |
| 105 | } |
| 106 | |
| 107 | interface UseAiChatOptions { |
| 108 | messageTemplate?: (message: string) => string |
| 109 | setIsLoading?: Dispatch<SetStateAction<boolean>> |
| 110 | } |
| 111 | |
| 112 | const useAiChat = ({ messageTemplate = (message) => message, setIsLoading }: UseAiChatOptions) => { |
| 113 | const eventSourceRef = useRef<SSE | undefined>(undefined) |
| 114 | const messageIdempotencyKey = useRef(0) |
| 115 | |
| 116 | const [isResponding, setIsResponding] = useState(false) |
| 117 | const [hasError, setHasError] = useState(false) |
| 118 | |
| 119 | const [currentMessageIndex, setCurrentMessageIndex] = useState(1) |
| 120 | const [messages, dispatchMessage] = useReducer(messageReducer, []) |
| 121 | |
| 122 | const submit = useCallback( |
| 123 | async (query: string) => { |
| 124 | dispatchMessage({ |
| 125 | type: 'new', |
| 126 | message: { |
| 127 | status: MessageStatus.Complete, |
| 128 | role: MessageRole.User, |
| 129 | content: query, |
| 130 | }, |
| 131 | }) |
| 132 | dispatchMessage({ |
| 133 | type: 'new', |
| 134 | message: { |
| 135 | status: MessageStatus.Pending, |
| 136 | role: MessageRole.Assistant, |
| 137 | content: '', |
| 138 | }, |
| 139 | }) |
| 140 | setIsResponding(false) |
| 141 | setHasError(false) |
| 142 | setIsLoading?.(true) |
| 143 | |
| 144 | const eventSource = new SSE(`${BASE_PATH}/api/ai/docs`, { |
| 145 | headers: { |
| 146 | apikey: process.env.NEXT_PUBLIC_BRIVEN_ANON_KEY ?? '', |
| 147 | Authorization: `Bearer ${process.env.NEXT_PUBLIC_BRIVEN_ANON_KEY}`, |
| 148 | 'Content-Type': 'application/json', |
| 149 | }, |
| 150 | payload: JSON.stringify({ |
| 151 | messages: messages |
| 152 | .filter(({ status }) => status === MessageStatus.Complete) |
| 153 | .map(({ role, content }) => ({ role, content })) |
| 154 | .concat({ role: MessageRole.User, content: messageTemplate(query) }), |
| 155 | }), |
| 156 | }) |
| 157 | |
| 158 | function handleError<T>(err: T) { |
| 159 | setIsLoading?.(false) |
| 160 | setIsResponding(false) |
| 161 | setHasError(true) |
| 162 | console.error(err) |
| 163 | } |
| 164 | |
| 165 | function handleMessage(e: MessageEvent) { |
| 166 | try { |
| 167 | setIsLoading?.(false) |
| 168 | |
| 169 | if (e.data === '[DONE]') { |
| 170 | setIsResponding(false) |
| 171 | // Parse sources from the content and clean the message |
| 172 | dispatchMessage({ |
| 173 | type: 'finalize-with-sources', |
| 174 | index: currentMessageIndex, |
| 175 | }) |
| 176 | setCurrentMessageIndex((x) => x + 2) |
| 177 | return |
| 178 | } |
| 179 | |
| 180 | dispatchMessage({ |
| 181 | type: 'update', |
| 182 | index: currentMessageIndex, |
| 183 | message: { |
| 184 | status: MessageStatus.InProgress, |
| 185 | }, |
| 186 | }) |
| 187 | |
| 188 | setIsResponding(true) |
| 189 | |
| 190 | const data = JSON.parse(e.data) |
| 191 | const completionChunk: OpenAI.Chat.Completions.ChatCompletionChunk = data |
| 192 | const [ |
| 193 | { |
| 194 | delta: { content }, |
| 195 | }, |
| 196 | ] = completionChunk.choices |
| 197 | |
| 198 | if (content) { |
| 199 | dispatchMessage({ |
| 200 | type: 'append-content', |
| 201 | index: currentMessageIndex, |
| 202 | idempotencyKey: messageIdempotencyKey.current++, |
| 203 | content, |
| 204 | }) |
| 205 | } |
| 206 | } catch (err) { |
| 207 | handleError(err) |
| 208 | } |
| 209 | } |
| 210 | |
| 211 | eventSource.addEventListener('error', handleError) |
| 212 | eventSource.addEventListener('message', handleMessage) |
| 213 | |
| 214 | eventSource.stream() |
| 215 | |
| 216 | eventSourceRef.current = eventSource |
| 217 | |
| 218 | setIsLoading?.(true) |
| 219 | }, |
| 220 | [currentMessageIndex, messages, messageTemplate] |
| 221 | ) |
| 222 | |
| 223 | function reset() { |
| 224 | eventSourceRef.current?.close() |
| 225 | eventSourceRef.current = undefined |
| 226 | setIsResponding(false) |
| 227 | setHasError(false) |
| 228 | setCurrentMessageIndex(1) |
| 229 | dispatchMessage({ |
| 230 | type: 'reset', |
| 231 | }) |
| 232 | } |
| 233 | |
| 234 | return { |
| 235 | submit, |
| 236 | reset, |
| 237 | messages, |
| 238 | isResponding, |
| 239 | hasError, |
| 240 | } |
| 241 | } |
| 242 | |
| 243 | export { useAiChat } |
| 244 | export type { Message, UseAiChatOptions } |