mirror of
https://github.com/danny-avila/LibreChat.git
synced 2026-10-10 08:04:10 +00:00
🌲 test: Add E2E Coverage for Message Tree Streaming (#13570)
* add e2e message tree stream coverage * fix e2e message tree review findings * expand message tree e2e recovery coverage * fix stream-start failure recovery coverage
This commit is contained in:
parent
1612dba353
commit
15108f0f2f
9 changed files with 878 additions and 28 deletions
|
|
@ -26,9 +26,14 @@ import type {
|
|||
} from 'librechat-data-provider';
|
||||
import type { SetterOrUpdater } from 'recoil';
|
||||
import type { TAskFunction, ExtendedFile } from '~/common';
|
||||
import {
|
||||
logger,
|
||||
hasStreamStartFailed,
|
||||
createDualMessageContent,
|
||||
getRouteChatProjectId,
|
||||
} from '~/utils';
|
||||
import useSetFilesToDelete from '~/hooks/Files/useSetFilesToDelete';
|
||||
import useGetSender from '~/hooks/Conversations/useGetSender';
|
||||
import { logger, createDualMessageContent, getRouteChatProjectId } from '~/utils';
|
||||
import store, { useGetEphemeralAgent } from '~/store';
|
||||
import { startupConfigKey } from '~/data-provider';
|
||||
import useUserKey from '~/hooks/Input/useUserKey';
|
||||
|
|
@ -40,6 +45,31 @@ const logChatRequest = (request: Record<string, unknown>) => {
|
|||
logger.log('=====================================');
|
||||
};
|
||||
|
||||
const getAppendParentMessageId = ({
|
||||
latestMessage,
|
||||
currentMessages,
|
||||
}: {
|
||||
latestMessage: TMessage | null;
|
||||
currentMessages: TMessage[];
|
||||
}) => {
|
||||
if (!latestMessage) {
|
||||
return Constants.NO_PARENT;
|
||||
}
|
||||
|
||||
if (!hasStreamStartFailed(latestMessage)) {
|
||||
return latestMessage.messageId;
|
||||
}
|
||||
|
||||
const failedUserMessage = currentMessages.find(
|
||||
(message) => message.messageId === latestMessage.parentMessageId,
|
||||
);
|
||||
if (failedUserMessage?.isCreatedByUser !== true) {
|
||||
return latestMessage.messageId;
|
||||
}
|
||||
|
||||
return failedUserMessage.parentMessageId ?? Constants.NO_PARENT;
|
||||
};
|
||||
|
||||
export default function useChatFunctions({
|
||||
index = 0,
|
||||
files,
|
||||
|
|
@ -165,7 +195,7 @@ export default function useChatFunctions({
|
|||
}
|
||||
const isEditOrContinue = isEdited || isContinued;
|
||||
|
||||
let currentMessages: TMessage[] | null = overrideMessages ?? getMessages() ?? [];
|
||||
let currentMessages: TMessage[] = overrideMessages ?? getMessages() ?? [];
|
||||
|
||||
if (conversation?.promptPrefix) {
|
||||
conversation.promptPrefix = replaceSpecialVars({
|
||||
|
|
@ -184,7 +214,8 @@ export default function useChatFunctions({
|
|||
// construct the query message
|
||||
// this is not a real messageId, it is used as placeholder before real messageId returned
|
||||
const intermediateId = overrideUserMessageId ?? v4();
|
||||
parentMessageId = parentMessageId ?? latestMessage?.messageId ?? Constants.NO_PARENT;
|
||||
parentMessageId =
|
||||
parentMessageId ?? getAppendParentMessageId({ latestMessage, currentMessages });
|
||||
|
||||
logChatRequest({
|
||||
index,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import debounce from 'lodash/debounce';
|
||||
import { useEffect, useRef, useCallback } from 'react';
|
||||
import debounce from 'lodash/debounce';
|
||||
import { useRecoilValue, useRecoilState } from 'recoil';
|
||||
import type { TEndpointOption } from 'librechat-data-provider';
|
||||
import type { KeyboardEvent } from 'react';
|
||||
|
|
@ -11,12 +11,12 @@ import {
|
|||
checkIfScrollable,
|
||||
} from '~/utils';
|
||||
import { useAssistantsMapContext } from '~/Providers/AssistantsMapContext';
|
||||
import { useLatestMessage } from '~/hooks/Messages/useLatestMessage';
|
||||
import { useAgentsMapContext } from '~/Providers/AgentsMapContext';
|
||||
import useGetSender from '~/hooks/Conversations/useGetSender';
|
||||
import useFileHandling from '~/hooks/Files/useFileHandling';
|
||||
import { useInteractionHealthCheck } from '~/data-provider';
|
||||
import { useChatContext } from '~/Providers/ChatContext';
|
||||
import { useLatestMessage } from '~/hooks/Messages/useLatestMessage';
|
||||
import { globalAudioId } from '~/common';
|
||||
import { useLocalize } from '~/hooks';
|
||||
import store from '~/store';
|
||||
|
|
@ -59,7 +59,8 @@ export default function useTextarea({
|
|||
});
|
||||
const entityName = entity?.name ?? '';
|
||||
|
||||
const isNotAppendable = latestMessage?.error === true && !isAssistant;
|
||||
const isNotAppendable =
|
||||
latestMessage?.error === true && latestMessage.isCreatedByUser === true && !isAssistant;
|
||||
// && (conversationId?.length ?? 0) > 6; // also ensures that we don't show the wrong placeholder
|
||||
|
||||
useEffect(() => {
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import { renderHook, act } from '@testing-library/react';
|
||||
import { renderHook, act, waitFor } from '@testing-library/react';
|
||||
import { Constants, LocalStorageKeys, QueryKeys, request } from 'librechat-data-provider';
|
||||
import type { TSubmission } from 'librechat-data-provider';
|
||||
|
||||
|
|
@ -48,6 +48,33 @@ const mockQueryClient = {
|
|||
}),
|
||||
};
|
||||
|
||||
const mockActiveRunAtom = { key: 'activeRun' };
|
||||
const mockAbortScrollAtom = { key: 'abortScroll' };
|
||||
const mockSubmissionAtom = { key: 'submission' };
|
||||
const mockShowStopButtonAtom = { key: 'showStopButton' };
|
||||
const mockSetActiveRun = jest.fn();
|
||||
const mockSetAbortScroll = jest.fn();
|
||||
const mockSetSubmission = jest.fn();
|
||||
const mockSetShowStopButton = jest.fn();
|
||||
const mockUseSetRecoilStateMock = jest.fn((atom: unknown) => {
|
||||
if (atom === mockActiveRunAtom) {
|
||||
return mockSetActiveRun;
|
||||
}
|
||||
if (atom === mockAbortScrollAtom) {
|
||||
return mockSetAbortScroll;
|
||||
}
|
||||
if (atom === mockSubmissionAtom) {
|
||||
return mockSetSubmission;
|
||||
}
|
||||
if (atom === mockShowStopButtonAtom) {
|
||||
return mockSetShowStopButton;
|
||||
}
|
||||
return jest.fn();
|
||||
});
|
||||
function mockUseSetRecoilState(atom: unknown) {
|
||||
return mockUseSetRecoilStateMock(atom);
|
||||
}
|
||||
|
||||
jest.mock('@tanstack/react-query', () => ({
|
||||
...jest.requireActual('@tanstack/react-query'),
|
||||
useQueryClient: () => mockQueryClient,
|
||||
|
|
@ -55,15 +82,16 @@ jest.mock('@tanstack/react-query', () => ({
|
|||
|
||||
jest.mock('recoil', () => ({
|
||||
...jest.requireActual('recoil'),
|
||||
useSetRecoilState: () => jest.fn(),
|
||||
useSetRecoilState: mockUseSetRecoilState,
|
||||
}));
|
||||
|
||||
jest.mock('~/store', () => ({
|
||||
__esModule: true,
|
||||
default: {
|
||||
activeRunFamily: jest.fn(),
|
||||
abortScrollFamily: jest.fn(),
|
||||
showStopButtonByIndex: jest.fn(),
|
||||
activeRunFamily: jest.fn(() => mockActiveRunAtom),
|
||||
abortScrollFamily: jest.fn(() => mockAbortScrollAtom),
|
||||
submissionByIndex: jest.fn(() => mockSubmissionAtom),
|
||||
showStopButtonByIndex: jest.fn(() => mockShowStopButtonAtom),
|
||||
},
|
||||
}));
|
||||
|
||||
|
|
@ -211,6 +239,11 @@ describe('useResumableSSE - 404 error path', () => {
|
|||
mockInvalidateQueries.mockClear();
|
||||
mockRemoveQueries.mockClear();
|
||||
mockFindAll.mockClear();
|
||||
mockUseSetRecoilStateMock.mockClear();
|
||||
mockSetActiveRun.mockClear();
|
||||
mockSetAbortScroll.mockClear();
|
||||
mockSetSubmission.mockClear();
|
||||
mockSetShowStopButton.mockClear();
|
||||
(request.post as jest.Mock).mockReset();
|
||||
(request.post as jest.Mock).mockResolvedValue({ streamId: 'stream-123' });
|
||||
});
|
||||
|
|
@ -537,6 +570,39 @@ describe('useResumableSSE - 404 error path', () => {
|
|||
jest.useRealTimers();
|
||||
});
|
||||
|
||||
it('clears submission and stop state when starting generation fails', async () => {
|
||||
(request.post as jest.Mock).mockRejectedValueOnce({
|
||||
response: {
|
||||
status: 500,
|
||||
data: { message: 'failed to start' },
|
||||
},
|
||||
});
|
||||
const submission = buildSubmission();
|
||||
const chatHelpers = buildChatHelpers();
|
||||
|
||||
const { unmount } = renderHook(() => useResumableSSE(submission, chatHelpers));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockSetSubmission).toHaveBeenCalledWith(null);
|
||||
});
|
||||
|
||||
expect(mockSSEInstances).toHaveLength(0);
|
||||
expect(mockErrorHandler).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
data: {
|
||||
text: JSON.stringify({ message: 'failed to start' }),
|
||||
metadata: { streamStartFailed: true },
|
||||
},
|
||||
submission,
|
||||
}),
|
||||
);
|
||||
expect(mockSetIsSubmitting).toHaveBeenCalledWith(true);
|
||||
expect(mockSetIsSubmitting).toHaveBeenCalledWith(false);
|
||||
expect(mockSetShowStopButton).toHaveBeenCalledWith(true);
|
||||
expect(mockSetShowStopButton).toHaveBeenCalledWith(false);
|
||||
unmount();
|
||||
});
|
||||
|
||||
it('replays title events from resume state sync', async () => {
|
||||
const submission = buildSubmission();
|
||||
const chatHelpers = buildChatHelpers();
|
||||
|
|
|
|||
|
|
@ -22,16 +22,22 @@ import type {
|
|||
EventSubmission,
|
||||
} from 'librechat-data-provider';
|
||||
import type { EventHandlerParams } from './useEventHandlers';
|
||||
import type { ActiveJobsResponse } from '~/data-provider';
|
||||
import type { TResData } from '~/common';
|
||||
import {
|
||||
clearAllDrafts,
|
||||
removeConvoFromAllQueries,
|
||||
upsertConvoInAllQueries,
|
||||
markStreamStartFailedMetadata,
|
||||
} from '~/utils';
|
||||
import {
|
||||
useGetUserBalance,
|
||||
useGetStartupConfig,
|
||||
queueTitleGeneration,
|
||||
streamStatusQueryKey,
|
||||
} from '~/data-provider';
|
||||
import type { ActiveJobsResponse } from '~/data-provider';
|
||||
import { useAuthContext } from '~/hooks/AuthContext';
|
||||
import useEventHandlers from './useEventHandlers';
|
||||
import { clearAllDrafts, removeConvoFromAllQueries, upsertConvoInAllQueries } from '~/utils';
|
||||
import store from '~/store';
|
||||
|
||||
type ChatHelpers = Pick<
|
||||
|
|
@ -39,6 +45,14 @@ type ChatHelpers = Pick<
|
|||
'setMessages' | 'getMessages' | 'setConversation' | 'setIsSubmitting' | 'newConversation'
|
||||
>;
|
||||
|
||||
const getStreamStartFailureData = (errorData?: Record<string, unknown>): TResData =>
|
||||
({
|
||||
text: errorData
|
||||
? JSON.stringify(errorData)
|
||||
: 'Error connecting to server, try refreshing the page.',
|
||||
metadata: markStreamStartFailedMetadata(),
|
||||
}) as unknown as TResData;
|
||||
|
||||
const MAX_RETRIES = 5;
|
||||
const START_GENERATION_NETWORK_RETRIES = 3;
|
||||
const START_GENERATION_READINESS_TIMEOUT_MS = 120000;
|
||||
|
|
@ -263,6 +277,7 @@ export default function useResumableSSE(
|
|||
const [_completed, setCompleted] = useState(new Set());
|
||||
const [streamId, setStreamId] = useState<string | null>(null);
|
||||
const setAbortScroll = useSetRecoilState(store.abortScrollFamily(runIndex));
|
||||
const setSubmission = useSetRecoilState(store.submissionByIndex(runIndex));
|
||||
const setShowStopButton = useSetRecoilState(store.showStopButtonByIndex(runIndex));
|
||||
|
||||
const sseRef = useRef<SSE | null>(null);
|
||||
|
|
@ -849,20 +864,16 @@ export default function useResumableSSE(
|
|||
|
||||
const axiosError = lastError as { response?: { data?: Record<string, unknown> } };
|
||||
const errorData = axiosError?.response?.data;
|
||||
if (errorData) {
|
||||
errorHandler({
|
||||
data: { text: JSON.stringify(errorData) } as unknown as Parameters<
|
||||
typeof errorHandler
|
||||
>[0]['data'],
|
||||
submission: currentSubmission as EventSubmission,
|
||||
});
|
||||
} else {
|
||||
errorHandler({ data: undefined, submission: currentSubmission as EventSubmission });
|
||||
}
|
||||
errorHandler({
|
||||
data: getStreamStartFailureData(errorData),
|
||||
submission: currentSubmission as EventSubmission,
|
||||
});
|
||||
setShowStopButton(false);
|
||||
setIsSubmitting(false);
|
||||
setSubmission(null);
|
||||
return null;
|
||||
},
|
||||
[clearStepMaps, errorHandler, setIsSubmitting],
|
||||
[clearStepMaps, errorHandler, setIsSubmitting, setShowStopButton, setSubmission],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import type { QueryClient } from '@tanstack/react-query';
|
|||
import type { LocalizeFunction } from '~/common';
|
||||
|
||||
export const TEXT_KEY_DIVIDER = '|||';
|
||||
export const STREAM_START_FAILED_METADATA_KEY = 'streamStartFailed';
|
||||
|
||||
type SiblingIndexLookup = (parentMessageId: string | null | undefined) => number;
|
||||
|
||||
|
|
@ -136,6 +137,16 @@ export const getAllContentText = (message?: TMessage | null): string => {
|
|||
return '';
|
||||
};
|
||||
|
||||
export const hasStreamStartFailed = (message?: Pick<TMessage, 'metadata'> | null): boolean =>
|
||||
message?.metadata?.[STREAM_START_FAILED_METADATA_KEY] === true;
|
||||
|
||||
export const markStreamStartFailedMetadata = (
|
||||
metadata?: TMessage['metadata'],
|
||||
): TMessage['metadata'] => ({
|
||||
...(metadata ?? {}),
|
||||
[STREAM_START_FAILED_METADATA_KEY]: true,
|
||||
});
|
||||
|
||||
const getLatestContentForKey = (message: TMessage): string => {
|
||||
const formatText = (str: string, index: number): string => {
|
||||
if (str.length === 0) {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue