From f1ab927a7ecad184dd0457ec9f95b81a2a6a928f Mon Sep 17 00:00:00 2001 From: jialin Date: Wed, 19 Feb 2025 20:45:33 +0800 Subject: [PATCH] chore: chat hooks --- src/components/alert-info/block.less | 58 +++++ src/components/alert-info/block.tsx | 48 ++++ src/components/logs-viewer/index.tsx | 147 ----------- src/components/logs-viewer/virtual-inner.tsx | 149 ----------- src/components/seal-table/types.ts | 1 - .../llmodels/components/view-logs-modal.tsx | 30 ++- .../playground/components/ground-left.tsx | 222 ++-------------- .../playground/components/message-input.tsx | 32 +-- .../components/model-parameters.tsx | 106 ++++++++ .../components/multiple-chat/index.tsx | 10 - .../components/multiple-chat/model-item.tsx | 239 +++--------------- .../components/reference-params.tsx | 2 + .../playground/hooks/use-chat-completion.ts | 218 ++++++++++++++++ src/pages/playground/hooks/use-init-meta.ts | 117 +++++++++ src/pages/playground/hooks/use-text-image.ts | 199 +++++++++++++++ src/utils/fetch-chunk-data.ts | 5 +- 16 files changed, 827 insertions(+), 756 deletions(-) create mode 100644 src/components/alert-info/block.less create mode 100644 src/components/alert-info/block.tsx delete mode 100644 src/components/logs-viewer/index.tsx delete mode 100644 src/components/logs-viewer/virtual-inner.tsx create mode 100644 src/pages/playground/components/model-parameters.tsx create mode 100644 src/pages/playground/hooks/use-chat-completion.ts create mode 100644 src/pages/playground/hooks/use-init-meta.ts create mode 100644 src/pages/playground/hooks/use-text-image.ts diff --git a/src/components/alert-info/block.less b/src/components/alert-info/block.less new file mode 100644 index 00000000..7f1ce39b --- /dev/null +++ b/src/components/alert-info/block.less @@ -0,0 +1,58 @@ +.alert-info-block { + padding: 6px 10px; + position: relative; + padding-top: 35px; + text-align: center; + border-radius: var(--border-radius-base); + margin: 0; + border: 1px solid transparent; + + .ant-typography { + margin-bottom: 0; + } + + &.danger { + border-color: var(--ant-color-error-border); + background-color: var(--ant-color-error-bg); + } + + &.warning { + border-color: var(--ant-color-warning-border); + background-color: var(--ant-color-warning-bg); + } + + .title { + position: absolute; + left: 0; + top: 0; + display: flex; + height: 30px; + width: 100%; + padding: 5px 10px; + border-radius: var(--border-radius-base) var(--border-radius-base) 0 0; + + &.danger { + background-color: var(--ant-color-error-bg-hover); + } + + &.warning { + background-color: var(--ant-color-warning-bg-hover); + } + + .info-icon { + margin-right: 8px; + + &.danger { + color: var(--ant-color-error); + } + + &.warning { + color: var(--ant-color-warning); + } + } + + .text { + font-weight: var(--font-weight-bold); + } + } +} diff --git a/src/components/alert-info/block.tsx b/src/components/alert-info/block.tsx new file mode 100644 index 00000000..04f6ab38 --- /dev/null +++ b/src/components/alert-info/block.tsx @@ -0,0 +1,48 @@ +import { WarningFilled } from '@ant-design/icons'; +import { Typography } from 'antd'; +import classNames from 'classnames'; +import React from 'react'; +import './block.less'; +interface AlertInfoProps { + type: 'danger' | 'warning'; + message: string; + rows?: number; + icon?: React.ReactNode; + ellipsis?: boolean; + style?: React.CSSProperties; + title: React.ReactNode; +} + +const AlertInfo: React.FC = (props) => { + const { message, type, rows = 1, ellipsis, style, title } = props; + + return ( + <> + {message ? ( +
+ +
+ + {title} +
+ {message} +
+
+ ) : null} + + ); +}; + +export default React.memo(AlertInfo); diff --git a/src/components/logs-viewer/index.tsx b/src/components/logs-viewer/index.tsx deleted file mode 100644 index abc42d69..00000000 --- a/src/components/logs-viewer/index.tsx +++ /dev/null @@ -1,147 +0,0 @@ -import useSetChunkFetch from '@/hooks/use-chunk-fetch'; -import useOverlayScroller from '@/hooks/use-overlay-scroller'; -import '@xterm/xterm/css/xterm.css'; -import classNames from 'classnames'; -import _ from 'lodash'; -import { memo, useCallback, useEffect, useRef, useState } from 'react'; -import { replaceLineRegex } from './config'; -import useParseAnsi from './parse-ansi'; -import './styles/index.less'; - -interface LogsViewerProps { - height: number; - content?: string; - url: string; - params?: object; -} -const LogsViewer: React.FC = (props) => { - const { height, url } = props; - const { - initialize, - updateScrollerPosition, - generateInstance, - scrollEventElement, - instance, - initialized - } = useOverlayScroller({ - options: { - theme: 'os-theme-light' - } - }); - const { isClean, parseAnsi } = useParseAnsi(); - const { setChunkFetch } = useSetChunkFetch(); - const chunkRequedtRef = useRef(null); - const scroller = useRef({}); - const cacheDataRef = useRef(''); - const uidRef = useRef(0); - const [logs, setLogs] = useState([]); - const stopScroll = useRef(false); - - const setId = () => { - uidRef.current += 1; - return uidRef.current; - }; - - const updateContent = useCallback( - (inputStr: string) => { - const data = inputStr.replace(replaceLineRegex, '\n'); - if (isClean(data)) { - cacheDataRef.current = data; - } else { - cacheDataRef.current += data; - } - const res = parseAnsi(cacheDataRef.current, setId); - console.log('res===', res); - setLogs(res); - }, - [setLogs, setId] - ); - - const createChunkConnection = async () => { - chunkRequedtRef.current?.current?.abort?.(); - chunkRequedtRef.current = setChunkFetch({ - url, - params: { - ...props.params, - watch: true - }, - contentType: 'text', - handler: updateContent - }); - }; - - const debounceResetStopScroll = _.debounce(() => { - stopScroll.current = false; - }, 30000); - - const handleOnWheel = useCallback( - (e: any) => { - const scrollTop = scrollEventElement?.scrollTop; - const scrollHeight = scrollEventElement?.scrollHeight; - const clientHeight = scrollEventElement?.clientHeight; - if (scrollTop + clientHeight >= scrollHeight) { - stopScroll.current = false; - } else { - stopScroll.current = true; - } - debounceResetStopScroll(); - }, - [debounceResetStopScroll, scrollEventElement] - ); - - const debounceUpdateScrollerPosition = _.debounce(() => { - generateInstance(); - updateScrollerPosition(0); - }, 200); - - useEffect(() => { - createChunkConnection(); - return () => { - chunkRequedtRef.current?.current?.abort?.(); - }; - }, [url, props.params]); - - useEffect(() => { - if (scroller.current) { - initialize(scroller.current); - } - }, [scroller.current, initialize]); - - useEffect(() => { - if (logs.length && !stopScroll.current && instance) { - updateScrollerPosition(0); - } else if (logs.length && !stopScroll.current && scroller.current) { - if (!initialized) { - initialize(scroller.current); - } - if (!instance) { - debounceUpdateScrollerPosition(); - } else { - updateScrollerPosition(0); - } - } - }, [logs, stopScroll.current, instance, scroller.current]); - - return ( -
-
-
- {_.map(logs, (item: any, index: number) => { - return ( -
- {item.content} -
- ); - })} -
-
-
- ); -}; - -export default memo(LogsViewer); diff --git a/src/components/logs-viewer/virtual-inner.tsx b/src/components/logs-viewer/virtual-inner.tsx deleted file mode 100644 index 727879f1..00000000 --- a/src/components/logs-viewer/virtual-inner.tsx +++ /dev/null @@ -1,149 +0,0 @@ -import _, { throttle } from 'lodash'; -import List from 'rc-virtual-list'; -import React, { - forwardRef, - useCallback, - useEffect, - useImperativeHandle, - useRef, - useState -} from 'react'; -import './styles/logs-list.less'; - -interface LogsInnerProps { - ref?: any; - data: { content: string; uid: number }[]; - onScroll?: (data: { isTop: boolean; isBottom: boolean }) => void; - diffHeight?: number; -} - -const LogsInner: React.FC = forwardRef((props, ref) => { - const { data, diffHeight = 96, onScroll } = props; - const viewportHeight = window.innerHeight; - const viewHeight = viewportHeight - diffHeight; - const [innerHieght, setInnerHeight] = useState(viewHeight); - const scroller = useRef(null); - const stopScroll = useRef(false); - const logsWrapper = useRef(null); - - const RC_VIRTUAL_LIST_HOLDER_CLASS = '.rc-virtual-list-holder'; - - const updataPositionToBottom = useCallback( - throttle(() => { - if (!stopScroll.current && data.length > 0) { - scroller.current?.scrollTo?.({ - index: data.length - 1, - align: 'bottom' - }); - } - }, 200), - [data, scroller.current, stopScroll.current] - ); - - const updataPositionToTop = useCallback( - throttle(() => { - if (!stopScroll.current && data.length > 0) { - scroller.current?.scrollTo?.({ - index: 0, - align: 'bottom' - }); - } - }, 200), - [data, scroller.current, stopScroll.current] - ); - - const updatePositionToTop = useCallback( - _.throttle((data: { isTop: boolean; isBottom: boolean }) => { - props.onScroll?.(data); - }, 200), - [props.onScroll] - ); - - const isScrollBottom = useCallback((root: HTMLElement) => { - const virtualList = root.querySelector(RC_VIRTUAL_LIST_HOLDER_CLASS); - if (!virtualList) { - return false; - } - return ( - virtualList.scrollTop >= - virtualList.scrollHeight - virtualList.clientHeight - ); - }, []); - - const isScrollTop = useCallback((root: HTMLElement) => { - const virtualList = root.querySelector(RC_VIRTUAL_LIST_HOLDER_CLASS); - if (!virtualList) { - return false; - } - return virtualList.scrollTop <= 0; - }, []); - - const handleOnScroll = useCallback((e: any) => { - const isBottom = isScrollBottom(logsWrapper.current); - const isTop = isScrollTop(logsWrapper.current); - if (isBottom) { - stopScroll.current = false; - } else { - stopScroll.current = true; - } - updatePositionToTop({ - isTop, - isBottom - }); - }, []); - - useImperativeHandle(ref, () => ({ - scrollToBottom() { - updataPositionToBottom(); - }, - scrollToTop() { - updataPositionToTop(); - } - })); - - useEffect(() => { - updataPositionToBottom(); - }, [updataPositionToBottom]); - - useEffect(() => { - const handleResize = throttle(() => { - const viewportHeight = window.innerHeight; - const viewHeight = viewportHeight - diffHeight; - setInnerHeight(viewHeight); - }, 100); - window.addEventListener('resize', handleResize); - return () => { - window.removeEventListener('resize', handleResize); - }; - }, [diffHeight]); - return ( -
- - {(item: any, index: number) => ( -
- {item.content} -
- )} -
-
- ); -}); - -export default React.memo(LogsInner); diff --git a/src/components/seal-table/types.ts b/src/components/seal-table/types.ts index ddb5bb00..c641c3c1 100644 --- a/src/components/seal-table/types.ts +++ b/src/components/seal-table/types.ts @@ -28,7 +28,6 @@ export interface TableHeaderProps { sortOrder?: 'ascend' | 'descend' | null; dataIndex: string; onSort?: (dataIndex: string, order: 'ascend' | 'descend') => void; - children?: React.ReactElement; title: React.ReactNode; style?: React.CSSProperties; firstCell?: boolean; diff --git a/src/pages/llmodels/components/view-logs-modal.tsx b/src/pages/llmodels/components/view-logs-modal.tsx index 69504adb..ca7a35cd 100644 --- a/src/pages/llmodels/components/view-logs-modal.tsx +++ b/src/pages/llmodels/components/view-logs-modal.tsx @@ -27,6 +27,7 @@ const ViewCodeModal: React.FC = (props) => { const [enableScorllLoad, setEnableScorllLoad] = useState(true); const logsViewerRef = React.useRef(null); const requestRef = React.useRef(null); + const contentRef = React.useRef(null); const handleCancel = useCallback(() => { logsViewerRef.current?.abort(); @@ -43,6 +44,29 @@ const ViewCodeModal: React.FC = (props) => { } }; + useEffect(() => { + const handleKeyDown = (e: any) => { + if ((e.ctrlKey || e.metaKey) && e.key === 'a') { + e.preventDefault(); + if (contentRef.current) { + const range = document.createRange(); + range.selectNodeContents(contentRef.current); + const selection = window.getSelection(); + selection?.removeAllRanges(); + selection?.addRange(range); + } + } + }; + + if (open) { + document.addEventListener('keydown', handleKeyDown); + } + + return () => { + document.removeEventListener('keydown', handleKeyDown); + }; + }, [open]); + useEffect(() => { if (!props.id) return; if (open) { @@ -51,10 +75,10 @@ const ViewCodeModal: React.FC = (props) => { url: `${MODELS_API}/${props.modelId}/instances`, handler: updateHandler }); - } else { - logsViewerRef.current?.abort(); } + return () => { + logsViewerRef.current?.abort(); requestRef.current?.current?.cancel?.(); }; }, [props.id, open]); @@ -85,7 +109,7 @@ const ViewCodeModal: React.FC = (props) => { width={modalSize.width} footer={null} > -
+
= forwardRef((props, ref) => { const { modelList } = props; - const messageId = useRef(0); - const [messageList, setMessageList] = useState([]); - const intl = useIntl(); const [searchParams] = useSearchParams(); const selectModel = searchParams.get('model') || ''; const [parameters, setParams] = useState({}); const [systemMessage, setSystemMessage] = useState(''); const [show, setShow] = useState(false); - const [loading, setLoading] = useState(false); - const [tokenResult, setTokenResult] = useState(null); const [collapse, setCollapse] = useState(false); - const contentRef = useRef(''); - const controllerRef = useRef(null); const scroller = useRef(null); - const currentMessageRef = useRef(null); const paramsRef = useRef(null); - const messageListLengthCache = useRef(0); - const reasonContentRef = useRef(''); const [actions, setActions] = useState([ 'upload', 'delete', 'copy' ]); - const { initialize, updateScrollerPosition } = useOverlayScroller(); const { initialize: innitializeParams } = useOverlayScroller(); + const { + submitMessage, + handleStopConversation, + handleAddNewMessage, + handleClear, + setMessageList, + tokenResult, + messageList, + loading + } = useChatCompletion(scroller); useImperativeHandle(ref, () => { return { @@ -81,159 +78,16 @@ const GroundLeft: React.FC = forwardRef((props, ref) => { ]); }, [messageList, systemMessage]); - const setMessageId = () => { - messageId.current = messageId.current + 1; - }; - - const formatContent = (data: { - content: string; - reasoningContent: string; - }) => { - if (data.reasoningContent && !data.content) { - return `${data.reasoningContent}`; - } - if (data.reasoningContent && data.content) { - return `${data.reasoningContent}${data.content}`; - } - return data.content; - }; - - const handleNewMessage = (message?: { role: string; content: string }) => { - const newMessage = message || { - role: - _.last(messageList)?.role === Roles.User ? Roles.Assistant : Roles.User, - content: '' - }; - messageList.push({ - ...newMessage, - uid: messageId.current + 1 - }); - setMessageId(); - setMessageList([...messageList]); - }; - - const joinMessage = (chunk: any) => { - console.log('chunk:', chunk); - setTokenResult({ - ...(chunk?.usage ?? {}) - }); - - if (!chunk || !_.get(chunk, 'choices', []).length) { - return; - } - - reasonContentRef.current = - reasonContentRef.current + - _.get(chunk, 'choices.0.delta.reasoning_content', ''); - contentRef.current = - contentRef.current + _.get(chunk, 'choices.0.delta.content', ''); - - const content = formatContent({ - content: contentRef.current, - reasoningContent: reasonContentRef.current - }); - - setMessageList([ - ...messageList, - ...currentMessageRef.current, - { - role: Roles.Assistant, - content: content, - uid: messageId.current - } - ]); - }; - const handleStopConversation = () => { - controllerRef.current?.abort?.(); - setLoading(false); - }; - - const submitMessage = async (current?: { role: string; content: string }) => { - if (!parameters.model) return; - try { - setLoading(true); - setMessageId(); - setTokenResult(null); - - controllerRef.current?.abort?.(); - controllerRef.current = new AbortController(); - const signal = controllerRef.current.signal; - currentMessageRef.current = current - ? [ - { - ...current, - uid: messageId.current - } - ] - : []; - - contentRef.current = ''; - reasonContentRef.current = ''; - setMessageList((pre) => { - return [...pre, ...currentMessageRef.current]; - }); - - const messageParams = [ - { role: Roles.System, content: systemMessage }, - ...messageList, - ...currentMessageRef.current - ]; - - const messages = generateMessages(messageParams); - - const chatParams = { - messages: messages, - ...parameters, - stream: true, - stream_options: { - include_usage: true - } - }; - const result: any = await fetchChunkedData({ - data: chatParams, - url: CHAT_API, - signal - }); - if (result?.error) { - setTokenResult({ - error: true, - errorMessage: - result?.data?.error?.message || result?.data?.message || '' - }); - return; - } - setMessageId(); - const { reader, decoder } = result; - await readStreamData(reader, decoder, (chunk: any) => { - if (chunk?.error) { - setTokenResult({ - error: true, - errorMessage: chunk?.error?.message || chunk?.message || '' - }); - return; - } - joinMessage(chunk); - }); - } catch (error) { - console.log('error:', error); - } finally { - setLoading(false); - } - }; - const handleClear = () => { - if (!messageList.length) { - return; - } - setMessageId(); - setMessageList([]); - setTokenResult(null); - }; - const handleSendMessage = (message: Omit) => { - console.log('message:', message); const currentMessage = message.content || message.imgs?.length ? message : undefined; - submitMessage(currentMessage); + submitMessage({ + system: systemMessage + ? { role: Roles.System, content: systemMessage } + : undefined, + current: currentMessage, + parameters + }); }; const handleCloseViewCode = () => { @@ -242,21 +96,6 @@ const GroundLeft: React.FC = forwardRef((props, ref) => { const handleSelectModel = () => {}; - const handlePresetPrompt = (list: { role: string; content: string }[]) => { - const sysMsg = list.filter((item) => item.role === 'system'); - const userMsg = list - .filter((item) => item.role === 'user') - .map((item) => { - setMessageId(); - return { - ...item, - uid: messageId.current - }; - }); - setSystemMessage(sysMsg[0]?.content || ''); - setMessageList(userMsg); - }; - const handleOnCheck = (e: any) => { const checked = e.target.checked; if (checked) { @@ -266,34 +105,12 @@ const GroundLeft: React.FC = forwardRef((props, ref) => { } }; - const throttleUpdatePosition = _.throttle(updateScrollerPosition, 100); - - useEffect(() => { - if (scroller.current) { - initialize(scroller.current); - } - }, [scroller.current, initialize]); - useEffect(() => { if (paramsRef.current) { innitializeParams(paramsRef.current); } }, [paramsRef.current, innitializeParams]); - useEffect(() => { - if (loading) { - console.log('loading:', loading); - updateScrollerPosition(); - } - }, [messageList, loading]); - - useEffect(() => { - if (messageList.length > messageListLengthCache.current) { - updateScrollerPosition(); - } - messageListLengthCache.current = messageList.length; - }, [messageList.length]); - return (
@@ -357,11 +174,10 @@ const GroundLeft: React.FC = forwardRef((props, ref) => { disabled={!parameters.model} isEmpty={!messageList.length} handleSubmit={handleSendMessage} - addMessage={handleNewMessage} + addMessage={handleAddNewMessage} handleAbortFetch={handleStopConversation} clearAll={handleClear} setModelSelections={handleSelectModel} - presetPrompt={handlePresetPrompt} />
diff --git a/src/pages/playground/components/message-input.tsx b/src/pages/playground/components/message-input.tsx index 7f151e58..6cd78841 100644 --- a/src/pages/playground/components/message-input.tsx +++ b/src/pages/playground/components/message-input.tsx @@ -16,7 +16,6 @@ import { useHotkeys } from 'react-hotkeys-hook'; import { Roles } from '../config'; import { MessageItem } from '../config/types'; import '../style/message-input.less'; -import PromptModal from './prompt-modal'; import ThumbImg from './thumb-img'; import UploadImg from './upload-img'; @@ -83,7 +82,6 @@ interface MessageInputProps { ) => void; onCheck?: (e: any) => void; submitIcon?: React.ReactNode; - presetPrompt?: (list: CurrentMessage[]) => void; addMessage?: (message: CurrentMessage) => void; onInputChange?: (e: any) => void; title?: React.ReactNode; @@ -105,7 +103,6 @@ const MessageInput: React.FC = forwardRef( { handleSubmit, handleAbortFetch, - presetPrompt, clearAll, updateLayout, addMessage, @@ -129,7 +126,6 @@ const MessageInput: React.FC = forwardRef( ) => { const { TextArea } = Input; const intl = useIntl(); - const [open, setOpen] = useState(false); const [focused, setFocused] = useState(false); const [message, setMessage] = useState({ role: Roles.User, @@ -203,6 +199,7 @@ const MessageInput: React.FC = forwardRef( const getPasteContent = useCallback( async (event: any) => { + // @ts-ignore const clipboardData = event.clipboardData || window.clipboardData; const items = clipboardData.items; const imgPromises: Promise[] = []; @@ -292,26 +289,6 @@ const MessageInput: React.FC = forwardRef( } }, [message.imgs, handleDeleteImg]); - const handleKeyDown = useCallback( - (event: any) => { - if ( - event.key === 'Backspace' && - message.content === '' && - message.imgs && - message.imgs?.length > 0 - ) { - // inputref blur - event.preventDefault(); - handleDeleteLastImage(); - } - }, - [message, handleDeleteLastImage] - ); - - const handleSelectPrompt = (list: CurrentMessage[]) => { - presetPrompt?.(list); - }; - useImperativeHandle(ref, () => ({ handleInputChange: handleInputChange })); @@ -363,6 +340,8 @@ const MessageInput: React.FC = forwardRef(
- setOpen(false)} - onSelect={handleSelectPrompt} - >
); } diff --git a/src/pages/playground/components/model-parameters.tsx b/src/pages/playground/components/model-parameters.tsx new file mode 100644 index 00000000..5cb3684b --- /dev/null +++ b/src/pages/playground/components/model-parameters.tsx @@ -0,0 +1,106 @@ +import FieldComponent from '@/components/seal-form/field-component'; +import { useIntl } from '@umijs/max'; +import { Form } from 'antd'; +import _ from 'lodash'; +import { + forwardRef, + memo, + useCallback, + useEffect, + useId, + useImperativeHandle, + useMemo +} from 'react'; +import { ParamsSchema } from '../config/types'; + +type ParamsSettingsProps = { + ref?: any; + style?: React.CSSProperties; + onValuesChange?: (changeValues: any, value: Record) => void; + paramsConfig?: ParamsSchema[]; + initialValues?: Record; + extra?: React.ReactNode; +}; + +const ParamsSettings: React.FC = forwardRef( + ({ onValuesChange, style, paramsConfig, initialValues, extra }, ref) => { + const intl = useIntl(); + const [form] = Form.useForm(); + const formId = useId(); + + useImperativeHandle(ref, () => ({ + form + })); + + useEffect(() => { + form.setFieldsValue({ + ...initialValues + }); + }, [initialValues]); + + const handleOnFinish = (values: any) => { + console.log('handleOnFinish', values); + }; + + const handleOnFinishFailed = (errorInfo: any) => { + console.log('handleOnFinishFailed', errorInfo); + }; + + const handleValuesChange = useCallback( + (changedValues: any, allValues: any) => { + onValuesChange?.(changedValues, allValues); + }, + [onValuesChange] + ); + + const renderFields = useMemo(() => { + if (!paramsConfig) { + return null; + } + const formValues = form?.getFieldsValue(); + return paramsConfig?.map((item: ParamsSchema) => { + return ( + + + + ); + }); + }, [paramsConfig, intl]); + + return ( +
+
+ {renderFields} + {extra} +
+
+ ); + } +); + +export default memo(ParamsSettings); diff --git a/src/pages/playground/components/multiple-chat/index.tsx b/src/pages/playground/components/multiple-chat/index.tsx index 221cc81a..1215a7d4 100644 --- a/src/pages/playground/components/multiple-chat/index.tsx +++ b/src/pages/playground/components/multiple-chat/index.tsx @@ -57,15 +57,6 @@ const MultiCompare: React.FC = ({ modelList, loaded }) => { return modelRefList.some((instanceId: symbol) => loadingStatus[instanceId]); }, [loadingStatus]); - const modelFullList = useMemo(() => { - return modelList.map((item) => { - return { - ...item, - disabled: modelSelections.some((model) => model.value === item.value) - }; - }); - }, [modelList, modelSelections]); - const setModelCounter = (model: string) => { modelsCounterMap.current[model] = _.add(modelsCounterMap.current[model], 1); return modelsCounterMap.current[model]; @@ -320,7 +311,6 @@ const MultiCompare: React.FC = ({ modelList, loaded }) => { clearAll={handleClearAll} updateLayout={updateLayout} setModelSelections={handleUpdateModelSelections} - presetPrompt={handlePresetPrompt} actions={[ 'clear', 'layout', diff --git a/src/pages/playground/components/multiple-chat/model-item.tsx b/src/pages/playground/components/multiple-chat/model-item.tsx index 090f771b..59cf5953 100644 --- a/src/pages/playground/components/multiple-chat/model-item.tsx +++ b/src/pages/playground/components/multiple-chat/model-item.tsx @@ -1,7 +1,5 @@ import AutoTooltip from '@/components/auto-tooltip'; import IconFont from '@/components/icon-font'; -import useOverlayScroller from '@/hooks/use-overlay-scroller'; -import { fetchChunkedData, readStreamData } from '@/utils/fetch-chunk-data'; import { ClearOutlined, DeleteOutlined, @@ -23,10 +21,10 @@ import React, { useState } from 'react'; import 'simplebar-react/dist/simplebar.min.css'; -import { CHAT_API } from '../../apis'; import { OpenAIViewCode, Roles, generateMessages } from '../../config'; import CompareContext from '../../config/compare-context'; import { MessageItem, ModelSelectionItem } from '../../config/types'; +import useChatCompletion from '../../hooks/use-chat-completion'; import '../../style/model-item.less'; import ParamsSettings from '../params-settings'; import ReferenceParams from '../reference-params'; @@ -50,7 +48,6 @@ const ModelItem: React.FC = forwardRef( handleDeleteModel, handleApplySystemChangeToAll, modelFullList, - loadingStatus, actions } = useContext(CompareContext); const intl = useIntl(); @@ -59,19 +56,19 @@ const ModelItem: React.FC = forwardRef( const [params, setParams] = useState>({ model: model }); - const [loading, setLoading] = useState(false); - const messageId = useRef(0); - const [messageList, setMessageList] = useState([]); - const [tokenResult, setTokenResult] = useState(null); const [show, setShow] = useState(false); - const contentRef = useRef(''); - const controllerRef = useRef(null); - const currentMessageRef = useRef([]); - const modelScrollRef = useRef(null); - const messageListLengthCache = useRef(0); - const reasonContentRef = useRef(''); + const scroller = useRef(null); - const { initialize, updateScrollerPosition } = useOverlayScroller(); + const { + submitMessage, + handleAddNewMessage, + handleClear, + setMessageList, + handleStopConversation, + tokenResult, + messageList, + loading + } = useChatCompletion(scroller); const viewCodeMessage = useMemo(() => { return generateMessages([ @@ -80,134 +77,11 @@ const ModelItem: React.FC = forwardRef( ]); }, [messageList, systemMessage]); - const setMessageId = () => { - messageId.current = messageId.current + 1; - }; - const abortFetch = () => { - controllerRef.current?.abort?.(); + handleStopConversation(); setLoadingStatus(instanceId, false); }; - const formatContent = (data: { - content: string; - reasoningContent: string; - }) => { - if (data.reasoningContent && !data.content) { - return `${data.reasoningContent}`; - } - if (data.reasoningContent && data.content) { - return `${data.reasoningContent}${data.content}`; - } - return data.content; - }; - - const joinMessage = (chunk: any) => { - setTokenResult({ - ...(chunk?.usage ?? {}) - }); - - if (!chunk || !_.get(chunk, 'choices', [].length)) { - return; - } - - reasonContentRef.current = - reasonContentRef.current + - _.get(chunk, 'choices.0.delta.reasoning_content', ''); - - contentRef.current = - contentRef.current + _.get(chunk, 'choices.0.delta.content', ''); - - const content = formatContent({ - content: contentRef.current, - reasoningContent: reasonContentRef.current - }); - setMessageList([ - ...messageList, - ...currentMessageRef.current, - { - role: Roles.Assistant, - content, - uid: messageId.current - } - ]); - }; - - const submitMessage = async (currentMessage?: Omit) => { - if (!params.model) return; - try { - setLoadingStatus(instanceId, true); - setMessageId(); - - controllerRef.current?.abort?.(); - controllerRef.current = new AbortController(); - const signal = controllerRef.current.signal; - currentMessageRef.current = currentMessage - ? [ - { - ...currentMessage, - uid: messageId.current - } - ] - : []; - setMessageList((preList) => { - return [...preList, ...currentMessageRef.current]; - }); - - contentRef.current = ''; - reasonContentRef.current = ''; - // ====== payload ================= - - const messageParams = [ - { role: Roles.System, content: systemMessage }, - ...messageList, - ...currentMessageRef.current - ]; - - const messages = generateMessages(messageParams); - - const chatParams = { - messages: messages, - ...params, - stream: true, - stream_options: { - include_usage: true - } - }; - // ============== payload end ================ - const result: any = await fetchChunkedData({ - data: chatParams, - url: CHAT_API, - signal - }); - - if (result?.error) { - setTokenResult({ - error: true, - errorMessage: - result?.data?.error?.message || result?.data?.message || '' - }); - return; - } - setMessageId(); - const { reader, decoder } = result; - await readStreamData(reader, decoder, (chunk: any) => { - if (chunk?.error) { - setTokenResult({ - error: true, - errorMessage: chunk?.error?.message || chunk?.message || '' - }); - return; - } - joinMessage(chunk); - }); - } catch (error) { - // console.log('error:', error); - } finally { - setLoadingStatus(instanceId, false); - } - }; - const handleDelete = () => { handleDeleteModel(instanceId); }; @@ -217,7 +91,14 @@ const ModelItem: React.FC = forwardRef( currentMessage.content || currentMessage.imgs?.length ? currentMessage : undefined; - submitMessage(currentMsg); + + submitMessage({ + system: systemMessage + ? { role: Roles.System, content: systemMessage } + : undefined, + current: currentMsg, + parameters: params + }); }; const handleApplyToAllModels = (e: any) => { @@ -249,25 +130,6 @@ const ModelItem: React.FC = forwardRef( [params, isApplyToAllModels.current] ); - const handleClearMessage = () => { - setMessageList([]); - setTokenResult(null); - currentMessageRef.current = []; - }; - - const addNewMessage = (message: Omit) => { - setMessageId(); - setMessageList((preList) => { - return [ - ...preList, - { - ...message, - uid: messageId.current - } - ]; - }); - }; - const handleCloseViewCode = () => { setShow(false); }; @@ -277,29 +139,9 @@ const ModelItem: React.FC = forwardRef( ...params, model: value }); - handleClearMessage(); + handleClear(); }; - const handlePresetMessageList = (list: MessageItem[]) => { - currentMessageRef.current = []; - const messages = _.map(list, (item: Omit) => { - setMessageId(); - return { - role: item.role, - content: item.content, - uid: messageId.current - }; - }); - setTokenResult(null); - setMessageList(messages); - }; - - const modelOptions = useMemo(() => { - return modelFullList.filter((item) => { - return item.type !== 'empty'; - }); - }, [modelFullList]); - const actionItems = useMemo(() => { const list = [ { @@ -344,37 +186,18 @@ const ModelItem: React.FC = forwardRef( }, [globalParams]); useEffect(() => { + setLoadingStatus(instanceId, loading); return () => { - abortFetch(); + setLoadingStatus(instanceId, false); }; - }, []); - - useEffect(() => { - if (modelScrollRef.current) { - initialize(modelScrollRef.current); - } - }, [modelScrollRef.current, initialize]); - - useEffect(() => { - if (loadingStatus[instanceId]) { - updateScrollerPosition(); - } - }, [messageList]); - - useEffect(() => { - if (messageList.length > messageListLengthCache.current) { - updateScrollerPosition(); - } - messageListLengthCache.current = messageList.length; - }, [messageList.length]); + }, [loading]); useImperativeHandle(ref, () => { return { submit: handleSubmit, abortFetch, - addNewMessage, - clear: handleClearMessage, - presetPrompt: handlePresetMessageList, + addNewMessage: handleAddNewMessage, + clear: handleClear, setSystemMessage, loading }; @@ -473,7 +296,7 @@ const ModelItem: React.FC = forwardRef( applyToAll={handleApplySystemChangeToAll} setSystemMessage={setSystemMessage} > -
+
= forwardRef( actions={actions} editable={true} /> - +
{ textAlign: 'center', paddingBlock: 0, margin: 0, + paddingInline: 4, + borderRadius: 2, backgroundColor: 'var(--ant-color-error-bg)' }} > diff --git a/src/pages/playground/hooks/use-chat-completion.ts b/src/pages/playground/hooks/use-chat-completion.ts new file mode 100644 index 00000000..ee337ca6 --- /dev/null +++ b/src/pages/playground/hooks/use-chat-completion.ts @@ -0,0 +1,218 @@ +import useOverlayScroller from '@/hooks/use-overlay-scroller'; +import { fetchChunkedData, readStreamData } from '@/utils/fetch-chunk-data'; +import _ from 'lodash'; +import { useEffect, useRef, useState } from 'react'; +import { CHAT_API } from '../apis'; +import { Roles, generateMessages } from '../config'; +import { MessageItem } from '../config/types'; + +export default function useChatCompletion( + scroller: React.RefObject +) { + const { initialize, updateScrollerPosition } = useOverlayScroller(); + const [loading, setLoading] = useState(false); + const [tokenResult, setTokenResult] = useState(null); + const [messageList, setMessageList] = useState([]); + const controllerRef = useRef(null); + const messageId = useRef(0); + const contentRef = useRef(''); + const currentMessageRef = useRef(null); + const messageListLengthCache = useRef(0); + const reasonContentRef = useRef(''); + + const setMessageId = () => { + messageId.current = messageId.current + 1; + return messageId.current; + }; + + const formatContent = (data: { + content: string; + reasoningContent: string; + }) => { + if (data.reasoningContent && !data.content) { + return `${data.reasoningContent}`; + } + if (data.reasoningContent && data.content) { + return `${data.reasoningContent}${data.content}`; + } + return data.content; + }; + + const joinMessage = (chunk: any) => { + console.log('chunk:', chunk); + setTokenResult({ + ...(chunk?.usage ?? {}) + }); + + if (!chunk || !_.get(chunk, 'choices', []).length) { + return; + } + + reasonContentRef.current = + reasonContentRef.current + + _.get(chunk, 'choices.0.delta.reasoning_content', ''); + contentRef.current = + contentRef.current + _.get(chunk, 'choices.0.delta.content', ''); + + const content = formatContent({ + content: contentRef.current, + reasoningContent: reasonContentRef.current + }); + + setMessageList([ + ...messageList, + ...currentMessageRef.current, + { + role: Roles.Assistant, + content: content, + uid: messageId.current + } + ]); + }; + + const handleClear = () => { + setMessageList([]); + setTokenResult(null); + }; + + const handleAddNewMessage = (message?: { role: string; content: string }) => { + const newMessage = message || { + role: + _.last(messageList)?.role === Roles.User ? Roles.Assistant : Roles.User, + content: '' + }; + setMessageList((preList) => [ + ...preList, + { + ...newMessage, + uid: setMessageId() + } + ]); + }; + + const handleStopConversation = () => { + controllerRef.current?.abort?.(); + setLoading(false); + }; + + const resetPerRequestCache = () => { + contentRef.current = ''; + reasonContentRef.current = ''; + }; + + const submitMessage = async (params: { + current?: { role: string; content: string }; + system?: { role: string; content: string }; + parameters: any; + }) => { + try { + setLoading(true); + setMessageId(); + setTokenResult(null); + + const { current, parameters, system } = params; + + controllerRef.current?.abort?.(); + controllerRef.current = new AbortController(); + const signal = controllerRef.current.signal; + currentMessageRef.current = current + ? [ + { + ...current, + uid: messageId.current + } + ] + : []; + + resetPerRequestCache(); + setMessageList((pre) => { + return [...pre, ...currentMessageRef.current]; + }); + + const messageParams = [ + ...(system ? [system] : []), + ...messageList, + ...currentMessageRef.current + ]; + + const messages = generateMessages(messageParams); + + const chatParams = { + messages: messages, + ...parameters, + stream: true, + stream_options: { + include_usage: true + } + }; + const result: any = await fetchChunkedData({ + data: chatParams, + url: CHAT_API, + signal + }); + if (result?.error) { + setTokenResult({ + error: true, + errorMessage: + result?.data?.error?.message || result?.data?.message || '' + }); + return; + } + setMessageId(); + const { reader, decoder } = result; + await readStreamData(reader, decoder, (chunk: any) => { + if (chunk?.error) { + setTokenResult({ + error: true, + errorMessage: chunk?.error?.message || chunk?.message || '' + }); + return; + } + joinMessage(chunk); + }); + } catch (error) { + console.log('error:', error); + } finally { + setLoading(false); + } + }; + const throttleUpdatePosition = _.throttle(updateScrollerPosition, 100); + + useEffect(() => { + if (scroller.current) { + initialize(scroller.current); + } + }, [scroller.current, initialize]); + + useEffect(() => { + if (loading) { + updateScrollerPosition(); + } + }, [messageList, loading]); + + useEffect(() => { + if (messageList.length > messageListLengthCache.current) { + updateScrollerPosition(); + } + messageListLengthCache.current = messageList.length; + }, [messageList.length]); + + useEffect(() => { + return () => { + handleStopConversation(); + }; + }, []); + + return { + loading, + tokenResult, + messageList, + setMessageId, + setMessageList, + handleClear, + handleAddNewMessage, + handleStopConversation, + updateScrollerPosition, + submitMessage + }; +} diff --git a/src/pages/playground/hooks/use-init-meta.ts b/src/pages/playground/hooks/use-init-meta.ts new file mode 100644 index 00000000..01a4b427 --- /dev/null +++ b/src/pages/playground/hooks/use-init-meta.ts @@ -0,0 +1,117 @@ +import _ from 'lodash'; +import { imageSizeOptions } from '../config/params-config'; + +const LLM_METAKEYS: Record = { + seed: 'seed', + stop: 'stop', + temperature: 'temperature', + top_p: 'top_p', + n_ctx: 'n_ctx', + n_slot: 'n_slot', + max_model_len: 'max_model_len' +}; + +const IMG_METAKEYS = [ + 'sample_method', + 'sampling_steps', + 'schedule_method', + 'cfg_scale', + 'guidance', + 'negative_prompt' +]; + +const llmInitialValues = { + seed: null, + stop: null, + temperature: 1, + top_p: 1, + max_tokens: 1024 +}; + +const imgInitialValues = { + n: 1, + seed: null, + sample_method: 'euler_a', + cfg_scale: 4.5, + guidance: 3.5, + sampling_steps: 10, + negative_prompt: null, + schedule_method: 'discrete', + preview: 'preview_faster' +}; + +export default function useInitMeta() { + const extractLLMMeta = (meta: any) => { + const modelMeta = meta || {}; + const modelMetaValue = _.pick(modelMeta, _.keys(LLM_METAKEYS)); + const obj = Object.entries(LLM_METAKEYS).reduce( + (acc: any, [key, value]) => { + const val = modelMetaValue[key]; + if (val && _.hasIn(modelMetaValue, key)) { + acc[value] = val; + } + return acc; + }, + {} + ); + + let defaultMaxTokens = 1024; + + if (obj.n_ctx && obj.n_slot) { + defaultMaxTokens = _.divide(obj.n_ctx / 2, obj.n_slot); + } else if (obj.max_model_len) { + defaultMaxTokens = obj.max_model_len / 2; + } + + return { + form: _.merge({}, llmInitialValues, { + ..._.omit(obj, ['n_ctx', 'n_slot', 'max_model_len']), + max_tokens: defaultMaxTokens + }), + meta: { + ...obj, + max_tokens: obj.max_model_len || _.divide(obj.n_ctx, obj.n_slot) + } + }; + }; + + const getNewImageSizeOptions = (metaData: any) => { + const { max_height, max_width } = metaData || {}; + if (!max_height || !max_width) { + return imageSizeOptions; + } + const newImageSizeOptions = imageSizeOptions.filter((item) => { + return item.width <= max_width && item.height <= max_height; + }); + if ( + !newImageSizeOptions.find( + (item) => item.width === max_width && item.height === max_height + ) + ) { + newImageSizeOptions.push({ + width: max_width, + height: max_height, + label: `${max_width}x${max_height}`, + value: `${max_width}x${max_height}` + }); + } + return newImageSizeOptions; + }; + + const extractIMGMeta = (meta: any) => { + return { + form: _.merge({}, imgInitialValues, { + ..._.pick(meta, IMG_METAKEYS), + width: meta?.default_width || 512, + height: meta?.default_height || 512 + }), + meta: meta, + sizeOptions: getNewImageSizeOptions(meta) + }; + }; + + return { + extractLLMMeta, + extractIMGMeta + }; +} diff --git a/src/pages/playground/hooks/use-text-image.ts b/src/pages/playground/hooks/use-text-image.ts new file mode 100644 index 00000000..6fe37046 --- /dev/null +++ b/src/pages/playground/hooks/use-text-image.ts @@ -0,0 +1,199 @@ +import { CREAT_IMAGE_API } from '@/pages/playground/apis'; +import { extractErrorMessage, promptList } from '@/pages/playground/config'; +import { generateRandomNumber } from '@/utils'; +import { + fetchChunkedData, + readLargeStreamData as readStreamData +} from '@/utils/fetch-chunk-data'; +import _ from 'lodash'; +import { useCallback, useRef, useState } from 'react'; + +const ODD_STRING = 'AAAABJRU5ErkJgg==='; + +export default function useTextImage() { + const [loading, setLoading] = useState(false); + const [tokenResult, setTokenResult] = useState(null); + const [imageList, setImageList] = useState([]); + const [currentPrompt, setCurrentPrompt] = useState(''); + const messageId = useRef(0); + const requestToken = useRef(null); + + const removeBase64Suffix = (str: string, suffix: string) => { + return str.endsWith(suffix) ? str.slice(0, -suffix.length) : str; + }; + + const setImageSize = useCallback((parameters: any) => { + let size: Record = { + span: 12 + }; + if (parameters.n === 1) { + size.span = 24; + } + if (parameters.n === 2) { + size.span = 12; + } + if (parameters.n === 3) { + size.span = 12; + } + if (parameters.n === 4) { + size.span = 12; + } + return size; + }, []); + + const setMessageId = () => { + messageId.current = messageId.current + 1; + return messageId.current; + }; + + const generateNumber = (min: number, max: number) => { + return Math.floor(Math.random() * (max - min + 1) + min); + }; + + const submitMessage = async (params: { + current?: { content: string }; + system?: { role: string; content: string }; + parameters: any; + }) => { + const { current, parameters } = params; + try { + if (!parameters.model) return; + const size: any = setImageSize(parameters); + setLoading(true); + setMessageId(); + setTokenResult(null); + setCurrentPrompt(current?.content || ''); + + const imgSize = [parameters.width, parameters.height]; + + // preview + let stream_options: Record = { + chunk_size: 16 * 1024, + chunk_results: true + }; + if (parameters.preview === 'preview') { + stream_options = { + preview: true + }; + } + + if (parameters.preview === 'preview_faster') { + stream_options = { + preview_faster: true + }; + } + + let newImageList = Array(parameters.n) + .fill({}) + .map((item, index: number) => { + return { + dataUrl: 'data:image/png;base64,', + ...size, + progress: 0, + height: imgSize[1], + width: imgSize[0], + loading: true, + progressType: 'dashboard', + preview: false, + uid: setMessageId() + }; + }); + setImageList(newImageList); + + requestToken.current?.abort?.(); + requestToken.current = new AbortController(); + + const params = { + ..._.omitBy( + parameters, + (value: string, key: string) => + !value || ['width', 'height', 'seed'].includes(key) + ), + size: `${imgSize[0]}x${imgSize[1]}`, + seed: parameters.random_seed ? generateRandomNumber() : parameters.seed, + stream: true, + stream_options: { + ...stream_options + }, + prompt: current?.content + }; + + const result: any = await fetchChunkedData({ + data: params, + url: `${CREAT_IMAGE_API}?t=${Date.now()}`, + signal: requestToken.current.signal + }); + if (result.error) { + setTokenResult({ + error: true, + errorMessage: extractErrorMessage(result) + }); + setImageList([]); + return; + } + + const { reader, decoder } = result; + + await readStreamData(reader, decoder, (chunk: any) => { + if (chunk?.error) { + setTokenResult({ + error: true, + errorMessage: chunk?.error?.message || chunk?.message || '' + }); + return; + } + chunk?.data?.forEach((item: any) => { + const imgItem = newImageList[item.index]; + if (item.b64_json && stream_options.chunk_results) { + imgItem.dataUrl += removeBase64Suffix(item.b64_json, ODD_STRING); + } else if (item.b64_json) { + imgItem.dataUrl = `data:image/png;base64,${removeBase64Suffix(item.b64_json, ODD_STRING)}`; + } + const progress = item.progress; + + newImageList[item.index] = { + dataUrl: imgItem.dataUrl, + height: imgSize[1], + width: imgSize[0], + maxHeight: `${imgSize[1]}px`, + maxWidth: `${imgSize[0]}px`, + uid: imgItem.uid, + span: imgItem.span, + loading: stream_options.chunk_results ? progress < 100 : false, + preview: progress >= 100, + progress: progress + }; + }); + setImageList([...newImageList]); + }); + } catch (error) { + console.log('error:', error); + requestToken.current?.abort?.(); + setImageList([]); + } finally { + setLoading(false); + } + }; + + const handleClear = () => { + setMessageId(); + setImageList([]); + setTokenResult(null); + }; + + const handleStopConversation = () => { + requestToken.current?.abort?.(); + setLoading(false); + }; + + return { + loading, + tokenResult, + imageList, + promptList, + handleStopConversation, + generateNumber, + handleClear, + submitMessage + }; +} diff --git a/src/utils/fetch-chunk-data.ts b/src/utils/fetch-chunk-data.ts index d23d1036..8644d2b9 100644 --- a/src/utils/fetch-chunk-data.ts +++ b/src/utils/fetch-chunk-data.ts @@ -118,10 +118,7 @@ export const fetchChunkedDataPostFormData = async (params: { signal: params.signal }); if (!response.ok) { - return { - error: true, - data: await errorHandler(response) - }; + return await errorHandler(response); } const reader = response?.body?.getReader(); const decoder = new TextDecoder('utf-8', {