import { setRouteCache } from '@/atoms/route-cache'; import routeCachekey from '@/config/route-cachekey'; import { HEADER_HEIGHT } from '@/config/settings'; import { useCancelToken } from '@/hooks/use-request-token'; import { readAudioFile } from '@/utils/load-audio-file'; import { SendOutlined } from '@ant-design/icons'; import { AlertInfo, AudioAnimation, AudioPlayer, CopyButton, IconFont, UploadAudio, useOverlayScroller } from '@gpustack/core-ui'; import { useIntl } from '@umijs/max'; import { Button, Spin, Tooltip } from 'antd'; import _ from 'lodash'; import React, { forwardRef, useCallback, useEffect, useImperativeHandle, useMemo, useRef, useState } from 'react'; import { AUDIO_SPEECH_TO_TEXT_API } from '../apis'; import AudioInput from '../components/audio-input'; import RightContainer from '../components/right-container'; import ViewCommonCode from '../components/view-common-code'; import { SpeechToTextFormat } from '../config'; import '../style/ground-llm.less'; import '../style/speech-to-text.less'; import '../style/system-message-wrap.less'; import { speechToTextCode } from '../view-code/audio'; import STTForm from './forms/stt-form'; import { useNonStreamSTT, useStreamSTT } from './hooks'; interface MessageProps { modelList: Global.BaseOption[]; ref?: any; } const GroundSTT: React.FC = forwardRef((props, ref) => { const intl = useIntl(); const { modelList } = props; const messageId = useRef(0); const [messageList, setMessageList] = useState< { uid: number; content: string }[] >([]); const [parameters, setParams] = useState({ model: '', language: 'auto' }); const [show, setShow] = useState(false); const [loading, setLoading] = useState(false); const [tokenResult, setTokenResult] = useState(null); const [collapse, setCollapse] = useState(false); const scroller = useRef(null); const [audioPermissionOn, setAudioPermissionOn] = useState(true); const [audioData, setAudioData] = useState(null); const [audioChunks, setAudioChunks] = useState({ data: [], analyser: null }); const [isRecording, setIsRecording] = useState(false); const formRef = useRef(null); const { updateCancelToken, getCanceltToken, cancelRequest } = useCancelToken(); const { initialize, updateScrollerPosition } = useOverlayScroller(); const nonStreamSTT = useNonStreamSTT({ onSuccess: (result) => { setMessageList([ { content: result.text, uid: messageId.current } ]); }, onError: (error) => { setTokenResult({ error: true, errorMessage: error }); } }); const streamSTT = useStreamSTT({ onChunk: (text) => { setMessageList([ { content: text, uid: messageId.current } ]); }, onComplete: () => { console.log('Stream transcription completed'); }, onError: (error) => { setTokenResult({ error: true, errorMessage: error?.message || _.toString(error) }); } }); useImperativeHandle(ref, () => { return { viewCode() { setShow(true); }, setCollapse() { setCollapse(!collapse); }, collapse: collapse }; }); const setMessageId = () => { messageId.current = messageId.current + 1; }; const viewCodeContent = useMemo(() => { return speechToTextCode({ api: AUDIO_SPEECH_TO_TEXT_API, parameters: { ...parameters } }); }, [parameters]); const handleStopConversation = () => { cancelRequest(); streamSTT.abort(); setLoading(false); }; const submitMessage = async () => { try { await formRef.current?.form.validateFields(); if (!parameters.model) return; setLoading(true); setMessageId(); setTokenResult(null); setMessageList([]); updateCancelToken(); setRouteCache(routeCachekey['/playground/speech'], true); const params = { ...parameters, file: new File([audioData.data], audioData.name, { type: audioData.type }) }; if (parameters.stream) { await streamSTT.generate(params); } else { await nonStreamSTT.generate(params, getCanceltToken()); } } catch (error: any) { setTokenResult({ error: true, errorMessage: error?.message || _.toString(error) }); } finally { setLoading(false); setIsRecording(false); setRouteCache(routeCachekey['/playground/speech'], false); } }; const handleClear = () => { setMessageId(); setMessageList([]); setTokenResult(null); }; const handleCloseViewCode = () => { setShow(false); }; const handleOnAudioData = useCallback( (data: { chunks: Blob[]; url: string; name: string; duration: number; type: string; }) => { setAudioData(() => { return { url: data.url, name: data.name, data: data.chunks, type: data.type, duration: data.duration }; }); }, [] ); const handleOnAudioPermission = useCallback((permission: boolean) => { setAudioPermissionOn(permission); }, []); const handleUploadChange = useCallback( async (data: { file: any; fileList: any }) => { try { const res = await readAudioFile(data.file); setAudioData(res); setTokenResult(null); } catch (error) {} }, [] ); const handleOnAnalyse = useCallback((data: any, analyser: any) => { setAudioChunks(() => { return { data: data, analyser: analyser }; }); }, []); const handleOnRecord = useCallback((val: boolean) => { setIsRecording(val); setAudioData(null); setTokenResult(null); setMessageList([]); }, []); const handleOnGenerate = async () => { if (loading) { handleStopConversation(); return; } submitMessage(); }; const renderAniamtion = () => { if (!audioPermissionOn) { return (
{intl.formatMessage({ id: 'playground.audio.enablemic' })}
); } if (isRecording) { return ( ); } return (
{intl.formatMessage({ id: 'playground.audio.speechtotext.tips' })}
); }; const updateParams = (values: any) => { setParams((pre: Record) => { return { ...pre, ...values }; }); }; useEffect(() => { if (scroller.current) { initialize(scroller.current); } }, [initialize]); useEffect(() => { if (loading) { updateScrollerPosition(); } }, [messageList, loading]); return (
{!isRecording && ( )}
{audioData ? (
{ } } >
) : ( renderAniamtion() )}
{messageList?.length > 0 && ( )}
{!tokenResult && (
{messageList.length ? ( messageList[0]?.content ) : ( {intl.formatMessage({ id: 'playground.audio.generating.tips' })} )}
)} {tokenResult && (
)}
{loading && (
)}
); }); export default GroundSTT;