chore: chat hooks

This commit is contained in:
jialin
2025-02-19 20:50:32 +08:00
parent 9e0634f449
commit f1ab927a7e
16 changed files with 827 additions and 756 deletions
@@ -27,6 +27,7 @@ const ViewCodeModal: React.FC<ViewModalProps> = (props) => {
const [enableScorllLoad, setEnableScorllLoad] = useState(true);
const logsViewerRef = React.useRef<any>(null);
const requestRef = React.useRef<any>(null);
const contentRef = React.useRef<any>(null);
const handleCancel = useCallback(() => {
logsViewerRef.current?.abort();
@@ -43,6 +44,29 @@ const ViewCodeModal: React.FC<ViewModalProps> = (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<ViewModalProps> = (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<ViewModalProps> = (props) => {
width={modalSize.width}
footer={null}
>
<div className="viewer-wrapper">
<div className="viewer-wrapper" ref={contentRef}>
<LogsViewer
ref={logsViewerRef}
height={modalSize.height}
+19 -203
View File
@@ -1,9 +1,7 @@
import useOverlayScroller from '@/hooks/use-overlay-scroller';
import { fetchChunkedData, readStreamData } from '@/utils/fetch-chunk-data';
import { useIntl, useSearchParams } from '@umijs/max';
import { Spin } from 'antd';
import classNames from 'classnames';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import {
forwardRef,
@@ -14,9 +12,9 @@ import {
useRef,
useState
} from 'react';
import { CHAT_API } from '../apis';
import { OpenAIViewCode, Roles, generateMessages } from '../config';
import { MessageItem, MessageItemAction } from '../config/types';
import useChatCompletion from '../hooks/use-chat-completion';
import '../style/ground-left.less';
import '../style/system-message-wrap.less';
import MessageInput from './message-input';
@@ -34,33 +32,32 @@ interface MessageProps {
const GroundLeft: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const messageId = useRef<number>(0);
const [messageList, setMessageList] = useState<MessageItem[]>([]);
const intl = useIntl();
const [searchParams] = useSearchParams();
const selectModel = searchParams.get('model') || '';
const [parameters, setParams] = useState<any>({});
const [systemMessage, setSystemMessage] = useState('');
const [show, setShow] = useState(false);
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [collapse, setCollapse] = useState(false);
const contentRef = useRef<any>('');
const controllerRef = useRef<any>(null);
const scroller = useRef<any>(null);
const currentMessageRef = useRef<any>(null);
const paramsRef = useRef<any>(null);
const messageListLengthCache = useRef<number>(0);
const reasonContentRef = useRef<any>('');
const [actions, setActions] = useState<MessageItemAction[]>([
'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<MessageProps> = 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 `<think>${data.reasoningContent}`;
}
if (data.reasoningContent && data.content) {
return `<think>${data.reasoningContent}</think>${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<MessageItem, 'uid'>) => {
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<MessageProps> = 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<MessageProps> = 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 (
<div className="ground-left-wrapper">
<div className="ground-left">
@@ -357,11 +174,10 @@ const GroundLeft: React.FC<MessageProps> = forwardRef((props, ref) => {
disabled={!parameters.model}
isEmpty={!messageList.length}
handleSubmit={handleSendMessage}
addMessage={handleNewMessage}
addMessage={handleAddNewMessage}
handleAbortFetch={handleStopConversation}
clearAll={handleClear}
setModelSelections={handleSelectModel}
presetPrompt={handlePresetPrompt}
/>
</div>
</div>
@@ -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<MessageInputProps> = forwardRef(
{
handleSubmit,
handleAbortFetch,
presetPrompt,
clearAll,
updateLayout,
addMessage,
@@ -129,7 +126,6 @@ const MessageInput: React.FC<MessageInputProps> = forwardRef(
) => {
const { TextArea } = Input;
const intl = useIntl();
const [open, setOpen] = useState(false);
const [focused, setFocused] = useState(false);
const [message, setMessage] = useState<CurrentMessage>({
role: Roles.User,
@@ -203,6 +199,7 @@ const MessageInput: React.FC<MessageInputProps> = forwardRef(
const getPasteContent = useCallback(
async (event: any) => {
// @ts-ignore
const clipboardData = event.clipboardData || window.clipboardData;
const items = clipboardData.items;
const imgPromises: Promise<string>[] = [];
@@ -292,26 +289,6 @@ const MessageInput: React.FC<MessageInputProps> = 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<MessageInputProps> = forwardRef(
<Button
type="text"
size="middle"
variant="filled"
color="default"
onClick={handleToggleRole}
icon={<SwapOutlined rotate={90} />}
>
@@ -515,11 +494,6 @@ const MessageInput: React.FC<MessageInputProps> = forwardRef(
></span>
)}
</div>
<PromptModal
open={open}
onCancel={() => setOpen(false)}
onSelect={handleSelectPrompt}
></PromptModal>
</div>
);
}
@@ -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<string, any>) => void;
paramsConfig?: ParamsSchema[];
initialValues?: Record<string, any>;
extra?: React.ReactNode;
};
const ParamsSettings: React.FC<ParamsSettingsProps> = 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 (
<Form.Item name={item.name} rules={item.rules} key={item.name}>
<FieldComponent
disabled={
item.disabledConfig
? item.disabledConfig?.when?.(formValues)
: item.disabled
}
description={
item.description?.isLocalized
? intl.formatMessage({ id: item.description.text })
: item.description?.text
}
onChange={null}
{..._.omit(item, [
'name',
'rules',
'disabledConfig',
'description'
])}
></FieldComponent>
</Form.Item>
);
});
}, [paramsConfig, intl]);
return (
<Form
style={{ ...style }}
name={formId}
form={form}
onValuesChange={handleValuesChange}
onFinish={handleOnFinish}
onFinishFailed={handleOnFinishFailed}
>
<div>
{renderFields}
{extra}
</div>
</Form>
);
}
);
export default memo(ParamsSettings);
@@ -57,15 +57,6 @@ const MultiCompare: React.FC<MultiCompareProps> = ({ 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<MultiCompareProps> = ({ modelList, loaded }) => {
clearAll={handleClearAll}
updateLayout={updateLayout}
setModelSelections={handleUpdateModelSelections}
presetPrompt={handlePresetPrompt}
actions={[
'clear',
'layout',
@@ -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<ModelItemProps> = forwardRef(
handleDeleteModel,
handleApplySystemChangeToAll,
modelFullList,
loadingStatus,
actions
} = useContext(CompareContext);
const intl = useIntl();
@@ -59,19 +56,19 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
const [params, setParams] = useState<Record<string, any>>({
model: model
});
const [loading, setLoading] = useState(false);
const messageId = useRef<number>(0);
const [messageList, setMessageList] = useState<MessageItem[]>([]);
const [tokenResult, setTokenResult] = useState<any>(null);
const [show, setShow] = useState(false);
const contentRef = useRef<any>('');
const controllerRef = useRef<any>(null);
const currentMessageRef = useRef<MessageItem[]>([]);
const modelScrollRef = useRef<any>(null);
const messageListLengthCache = useRef<number>(0);
const reasonContentRef = useRef<any>('');
const scroller = useRef<any>(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<ModelItemProps> = 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 `<think>${data.reasoningContent}`;
}
if (data.reasoningContent && data.content) {
return `<think>${data.reasoningContent}</think>${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<MessageItem, 'uid'>) => {
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<ModelItemProps> = 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<ModelItemProps> = forwardRef(
[params, isApplyToAllModels.current]
);
const handleClearMessage = () => {
setMessageList([]);
setTokenResult(null);
currentMessageRef.current = [];
};
const addNewMessage = (message: Omit<MessageItem, 'uid'>) => {
setMessageId();
setMessageList((preList) => {
return [
...preList,
{
...message,
uid: messageId.current
}
];
});
};
const handleCloseViewCode = () => {
setShow(false);
};
@@ -277,29 +139,9 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
...params,
model: value
});
handleClearMessage();
handleClear();
};
const handlePresetMessageList = (list: MessageItem[]) => {
currentMessageRef.current = [];
const messages = _.map(list, (item: Omit<MessageItem, 'uid'>) => {
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<ModelItemProps> = 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<ModelItemProps> = forwardRef(
applyToAll={handleApplySystemChangeToAll}
setSystemMessage={setSystemMessage}
></SystemMessage>
<div className="content" ref={modelScrollRef}>
<div className="content" ref={scroller}>
<div>
<MessageContent
messageList={messageList}
@@ -481,11 +304,7 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
actions={actions}
editable={true}
/>
<Spin
spinning={!!loadingStatus[instanceId]}
size="small"
style={{ width: '100%' }}
/>
<Spin spinning={loading} size="small" style={{ width: '100%' }} />
</div>
</div>
<ViewCodeModal
@@ -44,6 +44,8 @@ const ReferenceParams = (props: ReferenceParamsProps) => {
textAlign: 'center',
paddingBlock: 0,
margin: 0,
paddingInline: 4,
borderRadius: 2,
backgroundColor: 'var(--ant-color-error-bg)'
}}
>
@@ -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<HTMLElement>
) {
const { initialize, updateScrollerPosition } = useOverlayScroller();
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [messageList, setMessageList] = useState<MessageItem[]>([]);
const controllerRef = useRef<any>(null);
const messageId = useRef<number>(0);
const contentRef = useRef<any>('');
const currentMessageRef = useRef<any>(null);
const messageListLengthCache = useRef<number>(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 `<think>${data.reasoningContent}`;
}
if (data.reasoningContent && data.content) {
return `<think>${data.reasoningContent}</think>${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
};
}
+117
View File
@@ -0,0 +1,117 @@
import _ from 'lodash';
import { imageSizeOptions } from '../config/params-config';
const LLM_METAKEYS: Record<string, any> = {
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
};
}
@@ -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<any>(null);
const [imageList, setImageList] = useState<any[]>([]);
const [currentPrompt, setCurrentPrompt] = useState('');
const messageId = useRef<number>(0);
const requestToken = useRef<any>(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<string, string | number> = {
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<string, any> = {
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
};
}