import {__} from '@wordpress/i18n' import { useState, } from '@wordpress/element' import { useQuery, useQueryClient, } from '@tanstack/react-query' import { useDispatch, } from '@wordpress/data' import { store as noticesStore, } from '@wordpress/notices' import { useSkaBlocksOptions, } from '@ska/shared' import { debugError, debugInfo, } from '@ska/utils' import { openRouterRequest, } from '../api' import { getHTML, } from '../util' import type { AISession, Message, NonStreamingChoice, OpenRouterErrorResponse, OpenRouterResponse, } from '../types' const noop = () => {} export interface UseAIOptions extends AISession { /** Custom model to use. */ model?: string /** Called when a request to AI has been made. */ onRequestBegin?: () => void /** When a response value was received. */ onResponse?: (response: string) => void /** When a response error was received. */ onError?: (error: OpenRouterErrorResponse) => void /** When a response was resolved. */ onRequestFinish?: () => void } const useAI = (opts: UseAIOptions = {}) => { const { mode = 'html', systemPrompt = '', model, onRequestBegin = noop, onResponse, onError = noop, onRequestFinish = noop, } = opts const { createWarningNotice, createErrorNotice, } = useDispatch(noticesStore) const { openRouterApiKey, openRouterModel, } = useSkaBlocksOptions() const systemMessage: Message = {role: 'system', content: systemPrompt} const [messages, setMessages] = useState(systemPrompt ? [systemMessage] : []) const request = { models: [model || openRouterModel], messages, } const takeResponse = (data: OpenRouterResponse) => { const hasChoices = ( data.choices && Array.isArray(data.choices) && data.choices.length > 0 ) if(!hasChoices) { return } for(const choice of data.choices) { if(typeof choice !== 'object' || !('message' in choice)) { continue } const {message} = choice as NonStreamingChoice const {role, content} = message if(role && content) { const receivedContent = mode === 'html' ? getHTML(content) : content.trim() if(!receivedContent) { continue } const msg = {role, content: receivedContent} as Message setMessages(messages => [...messages, msg]) onResponse && onResponse(receivedContent) return } } createWarningNotice(__(`Couldn't extract an AI response.`, 'ska-blocks'), {type: 'snackbar'}) } const userMessages = messages.filter(({role}) => role === 'user') const queryKey = ['ai', request.models[0], ...userMessages.map(({content}, index) => typeof content === 'string' ? content : index)] debugInfo('AI query key', queryKey) const { data, ...query } = useQuery({ queryKey, queryFn: async ({signal}) => { debugInfo('useAI query') onRequestBegin() const response = await openRouterRequest(request, openRouterApiKey, {mode, signal}) onRequestFinish() if('code' in response) { debugError('OpenRouter responded with error code', response) onError(response) const msg = response?.message || 'OpenRouter error.' createErrorNotice(msg, {type: 'snackbar'}) throw new Error(msg) } takeResponse(response) return response }, enabled: !!request.models[0] && userMessages.length > 0 && !!systemPrompt, refetchOnMount: false, refetchOnWindowFocus: false, refetchOnReconnect: false, retry: false, }) const queryClient = useQueryClient() return { ...query, messages, sendMessage: (content: string) => { setMessages([ ...messages, { role: 'user', content: content || 'Try again.', }, ]) }, abort: () => { debugInfo('useAI cancelQueries.') queryClient.cancelQueries({queryKey}) }, } } export default useAI