🧵 fix: Prevent Message Loading Race During Streaming (#13295)

This commit is contained in:
Danny Avila 2026-05-24 18:50:00 -04:00 committed by GitHub
parent a8c43a4126
commit f2be5baecf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 605 additions and 6 deletions

View file

@ -1,5 +1,38 @@
import { createElement } from 'react';
import { act, renderHook, waitFor } from '@testing-library/react';
import { MemoryRouter } from 'react-router-dom';
import { QueryClient, QueryClientProvider } from '@tanstack/react-query';
import { QueryKeys, dataService } from 'librechat-data-provider';
import type { ReactNode } from 'react';
import type { TMessage } from 'librechat-data-provider';
import { getStableMessages } from '../queries';
import { logger } from '~/utils';
import {
getStableMessages,
shouldPreserveMessagesOnNotFound,
useGetMessagesByConvoId,
} from '../queries';
jest.mock('librechat-data-provider', () => {
const actual = jest.requireActual('librechat-data-provider');
return {
...actual,
dataService: {
...actual.dataService,
getMessagesByConvoId: jest.fn(),
},
};
});
jest.mock('~/utils', () => {
const actual = jest.requireActual('~/utils');
return {
...actual,
logger: {
...actual.logger,
warn: jest.fn(),
},
};
});
const message = (overrides: Partial<TMessage>): TMessage =>
({
@ -14,6 +47,20 @@ const message = (overrides: Partial<TMessage>): TMessage =>
...overrides,
}) as TMessage;
function createWrapper(queryClient: QueryClient, initialEntry: string) {
return function Wrapper({ children }: { children: ReactNode }) {
return createElement(
QueryClientProvider,
{ client: queryClient },
createElement(MemoryRouter, { initialEntries: [initialEntry] }, children),
);
};
}
afterEach(() => {
jest.clearAllMocks();
});
describe('getStableMessages', () => {
it('keeps cache when an empty result races with unhydrated stream messages', () => {
const currentMessages = [
@ -179,3 +226,176 @@ describe('getStableMessages', () => {
expect(result).toEqual([]);
});
});
describe('shouldPreserveMessagesOnNotFound', () => {
it('keeps cache when a transient 404 races with a pending assistant tail', () => {
const currentMessages = [
message({ messageId: 'persisted-1' }),
message({ messageId: 'user-2' }),
message({
messageId: 'user-2_',
parentMessageId: 'user-2',
isCreatedByUser: false,
createdAt: undefined,
updatedAt: undefined,
}),
];
expect(
shouldPreserveMessagesOnNotFound({
pathname: '/c/convo-id',
currentMessages,
isStreaming: true,
}),
).toBe(true);
});
it('does not preserve cache when no stream or active job is live', () => {
const currentMessages = [
message({ messageId: 'persisted-1' }),
message({ messageId: 'user-2' }),
message({
messageId: 'user-2_',
parentMessageId: 'user-2',
isCreatedByUser: false,
createdAt: undefined,
updatedAt: undefined,
}),
];
expect(
shouldPreserveMessagesOnNotFound({
pathname: '/c/convo-id',
currentMessages,
isStreaming: false,
}),
).toBe(false);
});
it('does not preserve cache on the new conversation route', () => {
const currentMessages = [
message({ messageId: 'user-1' }),
message({
messageId: 'user-1_',
parentMessageId: 'user-1',
isCreatedByUser: false,
createdAt: undefined,
updatedAt: undefined,
}),
];
expect(
shouldPreserveMessagesOnNotFound({
pathname: '/c/new',
currentMessages,
isStreaming: true,
}),
).toBe(false);
});
it('does not preserve cache when there is no pending assistant tail', () => {
const currentMessages = [message({ messageId: 'persisted-1' })];
expect(
shouldPreserveMessagesOnNotFound({
pathname: '/c/convo-id',
currentMessages,
isStreaming: true,
}),
).toBe(false);
});
});
describe('useGetMessagesByConvoId', () => {
it('keeps cache during submitting cleanup when active job cache still marks the stream active', async () => {
const conversationId = 'convo-id';
const currentMessages = [
message({ messageId: 'persisted-1' }),
message({ messageId: 'user-2', createdAt: undefined, updatedAt: undefined }),
message({
messageId: 'user-2_',
parentMessageId: 'user-2',
isCreatedByUser: false,
createdAt: undefined,
updatedAt: undefined,
}),
];
const serverMessages = [currentMessages[0]];
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false } },
});
queryClient.setQueryData([QueryKeys.messages, conversationId], currentMessages);
queryClient.setQueryData([QueryKeys.activeJobs], { activeJobIds: [conversationId] });
const mockGetMessagesByConvoId = dataService.getMessagesByConvoId as jest.MockedFunction<
typeof dataService.getMessagesByConvoId
>;
mockGetMessagesByConvoId.mockResolvedValue(serverMessages);
const { result, unmount } = renderHook(
() => useGetMessagesByConvoId(conversationId, undefined, { isStreaming: false }),
{ wrapper: createWrapper(queryClient, `/c/${conversationId}`) },
);
await act(async () => {
await result.current.refetch();
});
await waitFor(() => {
expect(result.current.data).toBe(currentMessages);
});
expect(dataService.getMessagesByConvoId).toHaveBeenCalledWith(conversationId);
expect(queryClient.getQueryData([QueryKeys.messages, conversationId])).toBe(currentMessages);
unmount();
});
it('keeps cache when a transient 404 races with a regenerated pending assistant tail', async () => {
const conversationId = 'convo-id';
const currentMessages = [
message({ messageId: 'persisted-1' }),
message({
messageId: 'user-2',
parentMessageId: 'persisted-1',
}),
message({
messageId: 'user-2_',
parentMessageId: 'user-2',
isCreatedByUser: false,
createdAt: undefined,
updatedAt: undefined,
}),
];
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false } },
});
queryClient.setQueryData([QueryKeys.messages, conversationId], currentMessages);
const mockGetMessagesByConvoId = dataService.getMessagesByConvoId as jest.MockedFunction<
typeof dataService.getMessagesByConvoId
>;
mockGetMessagesByConvoId.mockRejectedValueOnce({ status: 404 });
const { result, unmount } = renderHook(
() => useGetMessagesByConvoId(conversationId, { enabled: false }, { isStreaming: true }),
{ wrapper: createWrapper(queryClient, `/c/${conversationId}`) },
);
await act(async () => {
await result.current.refetch();
});
await waitFor(() => {
expect(result.current.data).toBe(currentMessages);
});
expect(dataService.getMessagesByConvoId).toHaveBeenCalledWith(conversationId);
expect(queryClient.getQueryData([QueryKeys.messages, conversationId])).toBe(currentMessages);
expect(logger.warn).toHaveBeenCalledWith(
'messages',
expect.stringContaining('returned 404 while cache has a pending assistant tail'),
currentMessages,
);
unmount();
});
});

View file

@ -4,7 +4,7 @@ import { useQuery, useQueryClient } from '@tanstack/react-query';
import type { UseQueryOptions, QueryObserverResult, QueryClient } from '@tanstack/react-query';
import { Constants, QueryKeys, dataService } from 'librechat-data-provider';
import type * as t from 'librechat-data-provider';
import { logger } from '~/utils';
import { isNotFoundError, logger } from '~/utils';
type StableMessagesParams = {
pathname: string;
@ -62,6 +62,18 @@ export function getStableMessages({
return result;
}
export function shouldPreserveMessagesOnNotFound({
pathname,
isStreaming = false,
currentMessages,
}: Pick<StableMessagesParams, 'pathname' | 'isStreaming' | 'currentMessages'>): boolean {
if (!isStreaming || pathname.includes('/c/new') || !currentMessages?.length) {
return false;
}
return hasPendingAssistantTail(currentMessages);
}
function hasActiveJob(queryClient: QueryClient, id: string) {
if (!id) {
return false;
@ -87,7 +99,32 @@ export const useGetMessagesByConvoId = <TData = t.TMessage[]>(
return useQuery<t.TMessage[], unknown, TData>(
[QueryKeys.messages, id],
async () => {
const result = await dataService.getMessagesByConvoId(id);
let result: t.TMessage[];
try {
result = await dataService.getMessagesByConvoId(id);
} catch (error) {
const currentMessages = queryClient.getQueryData<t.TMessage[]>([QueryKeys.messages, id]);
const hasLiveStream = isStreamingRef.current || hasActiveJob(queryClient, id);
if (
currentMessages &&
isNotFoundError(error) &&
shouldPreserveMessagesOnNotFound({
pathname: location.pathname,
currentMessages,
isStreaming: hasLiveStream,
})
) {
logger.warn(
'messages',
`Messages query for convo ${id} returned 404 while cache has a pending assistant tail; path: "${location.pathname}"`,
currentMessages,
);
return currentMessages;
}
throw error;
}
const currentMessages = queryClient.getQueryData<t.TMessage[]>([QueryKeys.messages, id]);
const stableMessages = getStableMessages({
pathname: location.pathname,

View file

@ -0,0 +1,41 @@
import { Constants } from 'librechat-data-provider';
import type { TMessage } from 'librechat-data-provider';
import { getMessageCacheIds, getMessagesConversationId } from '../cache';
const message = (conversationId?: string | null): TMessage =>
({
messageId: 'message-id',
conversationId,
}) as TMessage;
describe('chat message cache helpers', () => {
it('uses the latest concrete conversation id from streamed messages', () => {
expect(
getMessagesConversationId([
message(Constants.NEW_CONVO),
message(null),
message('generated-convo-id'),
]),
).toBe('generated-convo-id');
});
it('mirrors new-chat messages into the generated conversation cache', () => {
expect(
getMessageCacheIds({
queryParam: Constants.NEW_CONVO,
conversationId: Constants.NEW_CONVO,
messages: [message('generated-convo-id')],
}),
).toEqual([Constants.NEW_CONVO, 'generated-convo-id']);
});
it('keeps the current conversation cache id while avoiding duplicate ids', () => {
expect(
getMessageCacheIds({
queryParam: 'generated-convo-id',
conversationId: 'generated-convo-id',
messages: [message('generated-convo-id')],
}),
).toEqual(['generated-convo-id']);
});
});

View file

@ -0,0 +1,44 @@
import { Constants } from 'librechat-data-provider';
import type { TMessage } from 'librechat-data-provider';
type MessageCacheIdsParams = {
queryParam: string;
conversationId?: string | null;
messages: TMessage[];
};
function isConcreteConversationId(conversationId?: string | null) {
return (
!!conversationId &&
conversationId !== Constants.NEW_CONVO &&
conversationId !== Constants.PENDING_CONVO
);
}
export function getMessagesConversationId(messages: TMessage[]): string | undefined {
for (let i = messages.length - 1; i >= 0; i--) {
const conversationId = messages[i]?.conversationId;
if (isConcreteConversationId(conversationId)) {
return conversationId;
}
}
}
export function getMessageCacheIds({
queryParam,
conversationId,
messages,
}: MessageCacheIdsParams): string[] {
const ids = [queryParam];
const messageConversationId = getMessagesConversationId(messages);
if (queryParam === Constants.NEW_CONVO && isConcreteConversationId(conversationId)) {
ids.push(conversationId);
}
if (isConcreteConversationId(messageConversationId) && !ids.includes(messageConversationId)) {
ids.push(messageConversationId);
}
return ids;
}

View file

@ -8,6 +8,7 @@ import useChatFunctions from '~/hooks/Chat/useChatFunctions';
import { useAbortStreamMutation } from '~/data-provider';
import useNewConvo from '~/hooks/useNewConvo';
import { useLatestMessage, useLatestMessageId } from '~/hooks/Messages/useLatestMessage';
import { getMessageCacheIds } from './cache';
import store from '~/store';
// this to be set somewhere else
@ -42,9 +43,9 @@ export default function useChatHelpers(index = 0, paramId?: string) {
const setMessages = useCallback(
(messages: TMessage[]) => {
queryClient.setQueryData<TMessage[]>([QueryKeys.messages, queryParam], messages);
if (queryParam === 'new' && conversationId && conversationId !== 'new') {
queryClient.setQueryData<TMessage[]>([QueryKeys.messages, conversationId], messages);
const messageCacheIds = getMessageCacheIds({ queryParam, conversationId, messages });
for (const messageCacheId of messageCacheIds) {
queryClient.setQueryData<TMessage[]>([QueryKeys.messages, messageCacheId], messages);
}
},
[queryParam, queryClient, conversationId],