refactor: playground chat

This commit is contained in:
jialin
2024-09-22 14:55:23 +08:00
parent 4f4f00f0c8
commit 1abaa164f5
15 changed files with 406 additions and 173 deletions
@@ -1,5 +1,6 @@
import { Col, Row } from 'antd';
import React from 'react';
import { ModelSelectionItem } from '../../config/types';
import ModelItem from './model-item';
interface ActiveModelsProps {
@@ -7,8 +8,8 @@ interface ActiveModelsProps {
span: number;
count: number;
};
modelSelections: Global.BaseOption<string>[];
setModelRefs: (modelname: string, value: React.MutableRefObject<any>) => void;
modelSelections: ModelSelectionItem[];
setModelRefs: (modelname: symbol, value: React.MutableRefObject<any>) => void;
}
const ActiveModels: React.FC<ActiveModelsProps> = (props) => {
@@ -16,12 +17,13 @@ const ActiveModels: React.FC<ActiveModelsProps> = (props) => {
return (
<Row gutter={[16, 16]} style={{ height: '100%' }}>
{modelSelections.map((model, index) => (
<Col span={spans.span} key={model.value}>
<Col span={spans.span} key={`${model.value || 'empty'}-${model.uid}`}>
<ModelItem
key={model.value}
key={`${model.value || 'empty'}-${model.uid}`}
ref={(el: React.MutableRefObject<any>) =>
setModelRefs(model.value, el)
setModelRefs(model.instanceId, el)
}
instanceId={model.instanceId}
modelList={modelSelections}
model={model.value}
/>
@@ -1,22 +1,23 @@
import _ from 'lodash';
import { memo, useCallback, useEffect, useMemo, useRef, useState } from 'react';
import CompareContext from '../../config/compare-context';
import { ModelSelectionItem } from '../../config/types';
import '../../style/multiple-chat.less';
import MessageInput from '../message-input';
import ActiveModels from './active-models';
interface MultiCompareProps {
modelList: Global.BaseOption<string>[];
modelList: (Global.BaseOption<string> & { type?: string })[];
spans?: number;
}
const MultiCompare: React.FC<MultiCompareProps> = ({ modelList }) => {
const [loadingStatus, setLoadingStatus] = useState<Record<string, boolean>>(
const [loadingStatus, setLoadingStatus] = useState<Record<symbol, boolean>>(
{}
);
const [modelSelections, setModelSelections] = useState<
Global.BaseOption<string>[]
>([]);
const [modelSelections, setModelSelections] = useState<ModelSelectionItem[]>(
[]
);
const [globalParams, setGlobalParams] = useState<Record<string, any>>({
seed: null,
stop: null,
@@ -31,69 +32,89 @@ const MultiCompare: React.FC<MultiCompareProps> = ({ modelList }) => {
span: 12,
count: 2
});
const cacheModelInstanceList = useRef<any[]>([]);
const modelsCounterMap = useRef<Record<string, number>>({});
const modelRefs = useRef<any>({});
const boxHeight = 'calc(100vh - 72px)';
const isLoading = useMemo(() => {
console.log('loadingStatus========2', loadingStatus);
return _.keys(loadingStatus).some(
(modelname: string) => loadingStatus[modelname]
);
const modelRefList = Object.getOwnPropertySymbols(loadingStatus);
return modelRefList.some((instanceId: symbol) => loadingStatus[instanceId]);
}, [loadingStatus]);
useEffect(() => {
const list = modelList.slice?.(0, spans.count);
setModelSelections(list);
}, [modelList, spans.count]);
useEffect(() => {
modelRefs.current = {};
modelSelections.forEach((item) => {
modelRefs.current[item.value] = null;
const modelFullList = useMemo(() => {
return modelList.map((item) => {
return {
...item,
disabled: modelSelections.some((model) => model.value === item.value)
};
});
}, [modelSelections]);
}, [modelList, modelSelections]);
const setModelCounter = (model: string) => {
modelsCounterMap.current[model] = _.add(modelsCounterMap.current[model], 1);
return modelsCounterMap.current[model];
};
const pruneInstanceSymbol = (instanceId: symbol) => {
modelRefs.current[instanceId] = null;
loadingStatus[instanceId] = false;
};
const handleSubmit = (currentMessage: { role: string; content: string }) => {
const modelRefList = _.keys(modelRefs.current);
modelRefList.forEach(async (modelname: any, index: number) => {
const ref = modelRefs.current[modelname];
const modelRefList = Object.getOwnPropertySymbols(modelRefs.current);
modelRefList.forEach((instanceId: symbol) => {
const ref = modelRefs.current[instanceId];
ref?.submit(currentMessage);
});
};
const handleAddMessage = (message: { role: string; content: string }) => {
const modelRefList = Object.getOwnPropertySymbols(modelRefs.current);
modelRefList.forEach((instanceId: symbol) => {
const ref = modelRefs.current[instanceId];
ref?.setMessageList((preList: any) => {
return [...preList, { ...message }];
});
});
};
const handleAbortFetch = () => {
_.keys(modelRefs.current).forEach((modelname: string) => {
const ref = modelRefs.current[modelname];
const modelRefList = Object.getOwnPropertySymbols(modelRefs.current);
modelRefList.forEach((instanceId: symbol) => {
const ref = modelRefs.current[instanceId];
ref?.abortFetch();
});
};
const setModelRefs = useCallback(
(modelname: string, el: React.MutableRefObject<any>) => {
modelRefs.current[modelname] = el;
(instanceId: symbol, el: React.MutableRefObject<any>) => {
modelRefs.current[instanceId] = el;
},
[]
);
const handleSetLoadingStatus = (modeName: string, status: boolean) => {
const handleSetLoadingStatus = (instanceId: symbol, status: boolean) => {
setLoadingStatus((preStatus) => {
const newState = { ...preStatus };
newState[modeName] = status;
newState[instanceId] = status;
return newState;
});
};
const handleClearAll = () => {
_.keys(modelRefs.current).forEach((modelname: string) => {
const ref = modelRefs.current[modelname];
const modelRefList = Object.getOwnPropertySymbols(modelRefs.current);
modelRefList.forEach((instanceId: symbol) => {
const ref = modelRefs.current[instanceId];
ref?.clear();
});
};
const handleDeleteModel = (modelname: string) => {
const handleDeleteModel = (instanceId: symbol) => {
const newModelList = modelSelections.filter(
(model) => model.value !== modelname
(model) => model.instanceId !== instanceId
);
pruneInstanceSymbol(instanceId);
const span = Math.floor(24 / (24 / spans.span - 1));
setSpans({
span,
@@ -102,27 +123,108 @@ const MultiCompare: React.FC<MultiCompareProps> = ({ modelList }) => {
setModelSelections(newModelList);
};
const handleUpdateModelSelections = (list: Global.BaseOption<string>[]) => {
// set spans.span
const span = Math.floor(24 / list.length);
const handleUpdateModelSelections = (
list: (Global.BaseOption<string> & { instanceId: symbol })[]
) => {
const newList = list.map((item) => {
return {
...item,
uid: setModelCounter(item.value),
instanceId: Symbol(item.value)
};
});
const updateList = _.concat(modelSelections, newList);
const span = Math.floor(24 / updateList.length);
setSpans({
span: span < 8 ? 8 : span,
count: spans.count
});
setModelSelections(list);
setModelSelections(updateList);
};
const handlePresetPrompt = (list: { role: string; content: string }[]) => {
const sysMsg = list.filter((item) => item.role === 'system');
const userMsg = list.filter((item) => item.role === 'user');
const modelRefList = _.keys(modelRefs.current);
modelRefList.forEach(async (modelname: any) => {
const ref = modelRefs.current[modelname];
const modelRefList = Object.getOwnPropertySymbols(modelRefs.current);
modelRefList.forEach(async (instanceId: symbol) => {
const ref = modelRefs.current[instanceId];
ref?.presetPrompt(userMsg);
ref?.setSystemMessage(_.get(sysMsg, '0.content', ''));
});
};
const handleUpdateModelList = (spans: { span: number; count: number }) => {
const list = modelSelections;
// less than count
if (list.length < spans.count) {
const restCount = spans.count - list.length;
const restList = _.slice(modelList, list.length, list.length + restCount);
const resultList = Array.from(
{ length: restCount - restList.length },
(_, index) => {
return {
label: '',
value: '',
type: 'empty'
};
}
);
const newResultList = _.concat(restList, resultList).map((item: any) => {
return {
...item,
uid: setModelCounter(item.value || 'empty'),
instanceId:
item.type === 'empty' ? Symbol('empty') : Symbol(item.value)
};
});
setModelSelections(_.concat(list, newResultList));
return;
}
// more than count
if (list.length > spans.count) {
const newList = list.slice(0, spans.count);
setModelSelections(newList);
return;
}
};
const updateLayout = (value: { span: number; count: number }) => {
setSpans(value);
handleUpdateModelList(value);
};
useEffect(() => {
modelRefs.current = {};
let list = _.take(modelList, spans.count);
if (list.length < spans.count && list.length > 0) {
const restCount = spans.count - list.length;
const restList = Array.from({ length: restCount }, (_, index) => {
return {
label: '',
value: '',
type: 'empty'
};
});
list = _.concat(list, restList);
}
const resultList = list.map((item: any) => {
return {
...item,
uid: setModelCounter(item.value || 'empty'),
instanceId: item.type === 'empty' ? Symbol('empty') : Symbol(item.value)
};
});
setModelSelections(resultList);
}, [modelList]);
// useEffect(() => {
// modelRefs.current = {};
// modelSelections.forEach((item) => {
// modelRefs.current[item.instanceId] = null;
// });
// }, [modelSelections]);
return (
<div className="multiple-chat" style={{ height: boxHeight }}>
<div className="chat-list">
@@ -147,12 +249,13 @@ const MultiCompare: React.FC<MultiCompareProps> = ({ modelList }) => {
<MessageInput
loading={isLoading}
handleSubmit={handleSubmit}
addMessage={handleAddMessage}
handleAbortFetch={handleAbortFetch}
clearAll={handleClearAll}
setSpans={setSpans}
updateLayout={updateLayout}
setModelSelections={handleUpdateModelSelections}
presetPrompt={handlePresetPrompt}
modelList={modelList}
modelList={modelFullList}
/>
</div>
</div>
@@ -23,6 +23,7 @@ import React, {
useContext,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
@@ -30,6 +31,7 @@ import 'simplebar-react/dist/simplebar.min.css';
import { CHAT_API } from '../../apis';
import { Roles } from '../../config';
import CompareContext from '../../config/compare-context';
import { ModelSelectionItem } from '../../config/types';
import '../../style/model-item.less';
import ParamsSettings from '../params-settings';
import ReferenceParams from '../reference-params';
@@ -38,7 +40,8 @@ import MessageContent from './message-content';
interface ModelItemProps {
model: string;
modelList: Global.BaseOption<string>[];
modelList: ModelSelectionItem[];
instanceId: symbol;
ref: any;
}
@@ -49,7 +52,7 @@ interface MessageItemProps {
}
const ModelItem: React.FC<ModelItemProps> = forwardRef(
({ model, modelList }, ref) => {
({ model, modelList, instanceId }, ref) => {
const {
spans,
globalParams,
@@ -63,7 +66,8 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
const [autoSize, setAutoSize] = useState<{
minRows: number;
maxRows: number;
}>({ minRows: 1, maxRows: 1 });
focus: boolean;
}>({ minRows: 1, maxRows: 1, focus: false });
const [systemMessage, setSystemMessage] = useState<string>('');
const [params, setParams] = useState<Record<string, any>>({});
const [loading, setLoading] = useState(false);
@@ -74,6 +78,7 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
const contentRef = useRef<any>('');
const controllerRef = useRef<any>(null);
const currentMessageRef = useRef<MessageItemProps>({} as MessageItemProps);
const systemMessageRef = useRef<any>(null);
const setMessageId = () => {
messageId.current = messageId.current + 1;
@@ -81,7 +86,7 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
const abortFetch = () => {
controllerRef.current?.abort?.();
setLoadingStatus(params.model, false);
setLoadingStatus(instanceId, false);
};
const joinMessage = (chunk: any) => {
@@ -118,7 +123,7 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
const { parameters, currentMessage } = currentParams;
if (!parameters.model) return;
try {
setLoadingStatus(parameters.model, true);
setLoadingStatus(instanceId, true);
setMessageId();
controllerRef.current?.abort?.();
@@ -179,16 +184,15 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
await readStreamData(reader, decoder, (chunk: any) => {
joinMessage(chunk);
});
setLoadingStatus(params.model, false);
setLoadingStatus(instanceId, false);
} catch (error) {
console.log('error=====', error);
setLoadingStatus(params.model, false);
setLoadingStatus(instanceId, false);
}
};
const handleDropdownAction = useCallback(({ key }: { key: string }) => {
console.log('key:', key);
if (key === 'clear') {
setMessageList([]);
setSystemMessage('');
}
if (key === 'viewCode') {
setShow(true);
@@ -239,7 +243,9 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
setTokenResult(null);
setSystemMessage('');
currentMessageRef.current = {} as MessageItemProps;
console.log('clear message', systemMessage);
};
const handleCloseViewCode = () => {
setShow(false);
};
@@ -270,23 +276,40 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
};
const handleDelete = () => {
handleDeleteModel(params.model);
handleDeleteModel(instanceId);
};
const handleFocus = () => {
setAutoSize({
minRows: 4,
maxRows: 4
maxRows: 4,
focus: true
});
setTimeout(() => {
systemMessageRef.current?.focus?.({
cursor: 'end'
});
}, 100);
};
const handleBlur = () => {
setAutoSize({
minRows: 1,
maxRows: 1
maxRows: 1,
focus: false
});
};
const handleClearSystemMessage = () => {
setSystemMessage('');
};
const modelOptions = useMemo(() => {
return modelList.filter((item) => {
return item.type !== 'empty';
});
}, [modelList]);
useEffect(() => {
console.log('globalParams:', globalParams.model, globalParams);
setParams({
@@ -319,8 +342,9 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
<div className="header">
<span className="title">
<Select
style={{ minWidth: '100px' }}
variant="borderless"
options={modelList}
options={modelOptions}
onChange={handleModelChange}
value={params.model}
></Select>
@@ -390,19 +414,50 @@ const ModelItem: React.FC<ModelItemProps> = forwardRef(
></Button>
</span>
</div>
<div>
<Input.TextArea
variant="filled"
placeholder="Type system message here"
style={{ borderRadius: '0', border: 'none' }}
value={systemMessage}
autoSize={autoSize}
onFocus={handleFocus}
onBlur={handleBlur}
allowClear={false}
onChange={(e) => setSystemMessage(e.target.value)}
></Input.TextArea>
<Divider style={{ margin: '0' }}></Divider>
<div className="sys-message">
{
<div style={{ display: autoSize.focus ? 'block' : 'none' }}>
<Input.TextArea
ref={systemMessageRef}
variant="filled"
placeholder="Type system message here"
style={{
borderRadius: '0',
border: 'none'
}}
value={systemMessage}
autoSize={{
minRows: autoSize.minRows,
maxRows: autoSize.maxRows
}}
onFocus={handleFocus}
onBlur={handleBlur}
allowClear={false}
onChange={(e) => setSystemMessage(e.target.value)}
></Input.TextArea>
<Divider style={{ margin: '0' }}></Divider>
</div>
}
{!autoSize.focus && (
<div className="sys-content-wrap" onClick={handleFocus}>
<div className="sys-content">
{systemMessage || (
<span style={{ color: 'var(--ant-color-text-tertiary)' }}>
Type system message here
</span>
)}
</div>
{systemMessage && (
<Button
className="clear-btn"
type="text"
icon={<CloseOutlined />}
size="small"
onClick={handleClearSystemMessage}
></Button>
)}
</div>
)}
</div>
<div className="content">
<MessageContent