feat(core): Switch between different attempts at FE (no-changelog) (#20882)

This commit is contained in:
Jaakko Husso
2025-10-17 10:36:00 +03:00
committed by GitHub
parent 7706ec82c0
commit 697a144338
8 changed files with 167 additions and 114 deletions
-2
View File
@@ -137,8 +137,6 @@ export type ChatHubConversationsResponse = ChatHubSessionDto[];
export interface ChatHubConversationDto {
messages: Record<ChatMessageId, ChatHubMessageDto>;
rootIds: ChatMessageId[];
activeMessageChain: ChatMessageId[];
}
export interface ChatHubConversationResponse {
@@ -1,7 +1,6 @@
import { testDb, testModules } from '@n8n/backend-test-utils';
import type { User } from '@n8n/db';
import { Container } from '@n8n/di';
import { createAdmin, createMember } from '@test-integration/db/users';
import { ChatHubService } from '../chat-hub.service';
@@ -176,16 +175,10 @@ describe('chatHub', () => {
expect(response).toBeDefined();
const {
conversation: { rootIds, messages, activeMessageChain },
conversation: { messages },
} = response;
expect(rootIds).toEqual([msg1.id]);
expect(Object.keys(messages)).toHaveLength(4);
expect(activeMessageChain).toHaveLength(4);
expect(activeMessageChain[0]).toBe(msg1.id);
expect(activeMessageChain[1]).toBe(msg2.id);
expect(activeMessageChain[2]).toBe(msg3.id);
expect(activeMessageChain[3]).toBe(msg4.id);
expect(messages[msg1.id].content).toBe('message 1');
expect(messages[msg1.id].type).toBe('human');
expect(messages[msg1.id].turnId).toBe(msg1.id);
@@ -283,16 +276,10 @@ describe('chatHub', () => {
expect(response).toBeDefined();
const {
conversation: { rootIds, messages, activeMessageChain },
conversation: { messages },
} = response;
expect(rootIds).toEqual([msg1.id]);
expect(Object.keys(messages)).toHaveLength(6);
expect(activeMessageChain).toHaveLength(4);
expect(activeMessageChain[0]).toBe(msg1.id);
expect(activeMessageChain[1]).toBe(msg2.id);
expect(activeMessageChain[2]).toBe(msg5.id);
expect(activeMessageChain[3]).toBe(msg6.id);
expect(messages[msg1.id].content).toBe('message 1');
expect(messages[msg2.id].content).toBe('message 2');
expect(messages[msg3.id].content).toBe('message 3a');
@@ -346,7 +333,7 @@ describe('chatHub', () => {
turnId: ids[2],
createdAt: new Date('2025-01-03T00:10:00Z'),
});
const msg4 = await messagesRepository.createChatMessage({
await messagesRepository.createChatMessage({
id: ids[3],
sessionId: session.id,
name: 'ChatGPT',
@@ -362,14 +349,10 @@ describe('chatHub', () => {
expect(response).toBeDefined();
const {
conversation: { rootIds, messages, activeMessageChain },
conversation: { messages },
} = response;
expect(rootIds).toEqual([msg1.id, msg3.id]);
expect(Object.keys(messages)).toHaveLength(4);
expect(activeMessageChain).toHaveLength(2);
expect(activeMessageChain[0]).toBe(msg3.id);
expect(activeMessageChain[1]).toBe(msg4.id);
});
it('should get conversation with a retry branch at last message', async () => {
@@ -445,16 +428,10 @@ describe('chatHub', () => {
expect(response.session.id).toBe(session.id);
const {
conversation: { rootIds, messages, activeMessageChain },
conversation: { messages },
} = response;
expect(rootIds).toEqual([msg1.id]);
expect(Object.keys(messages)).toHaveLength(5);
expect(activeMessageChain).toHaveLength(4);
expect(activeMessageChain[0]).toBe(msg1.id);
expect(activeMessageChain[1]).toBe(msg2.id);
expect(activeMessageChain[2]).toBe(msg3.id);
expect(activeMessageChain[3]).toBe(msg5.id);
expect(messages[msg5.id].previousMessageId).toBe(msg3.id);
expect(messages[msg5.id].retryOfMessageId).toBe(msg4.id);
});
@@ -548,7 +525,7 @@ describe('chatHub', () => {
turnId: ids[4],
createdAt: new Date('2025-01-03T00:25:00Z'),
});
const msg1b = await messagesRepository.createChatMessage({
await messagesRepository.createChatMessage({
id: ids[6],
sessionId: session.id,
name: 'Nathan',
@@ -579,7 +556,7 @@ describe('chatHub', () => {
turnId: ids[8],
createdAt: new Date('2025-01-03T00:40:00Z'),
});
const msg4c = await messagesRepository.createChatMessage({
await messagesRepository.createChatMessage({
id: crypto.randomUUID(),
sessionId: session.id,
name: 'ChatGPT',
@@ -595,18 +572,11 @@ describe('chatHub', () => {
expect(response.session.id).toBe(session.id);
const {
conversation: { rootIds, messages, activeMessageChain },
conversation: { messages },
} = response;
expect(rootIds).toEqual([msg1.id, msg1b.id]);
expect(Object.keys(messages)).toHaveLength(10);
expect(activeMessageChain).toHaveLength(4);
expect(activeMessageChain[0]).toBe(msg1.id);
expect(activeMessageChain[1]).toBe(msg2r.id);
expect(activeMessageChain[2]).toBe(msg3d.id);
expect(activeMessageChain[3]).toBe(msg4c.id);
expect(messages[msg2r.id].previousMessageId).toBe(msg1.id);
expect(messages[msg2r.id].retryOfMessageId).toBe(msg2.id);
});
@@ -38,14 +38,6 @@ import {
} from 'n8n-workflow';
import { v4 as uuidv4 } from 'uuid';
import { ActiveExecutions } from '@/active-executions';
import { CredentialsService } from '@/credentials/credentials.service';
import { BadRequestError } from '@/errors/response-errors/bad-request.error';
import { NotFoundError } from '@/errors/response-errors/not-found.error';
import { DynamicNodeParametersService } from '@/services/dynamic-node-parameters.service';
import { getBase } from '@/workflow-execute-additional-data';
import { WorkflowExecutionService } from '@/workflows/workflow-execution.service';
import { ChatHubMessage } from './chat-hub-message.entity';
import type {
HumanMessagePayload,
@@ -57,6 +49,14 @@ import type {
import { ChatHubMessageRepository } from './chat-message.repository';
import { ChatHubSessionRepository } from './chat-session.repository';
import { ActiveExecutions } from '@/active-executions';
import { CredentialsService } from '@/credentials/credentials.service';
import { BadRequestError } from '@/errors/response-errors/bad-request.error';
import { NotFoundError } from '@/errors/response-errors/not-found.error';
import { DynamicNodeParametersService } from '@/services/dynamic-node-parameters.service';
import { getBase } from '@/workflow-execute-additional-data';
import { WorkflowExecutionService } from '@/workflows/workflow-execution.service';
const providerNodeTypeMapping: Record<ChatHubProvider, INodeTypeNameVersion> = {
openai: {
name: '@n8n/n8n-nodes-langchain.lmChatOpenAi',
@@ -348,11 +348,23 @@ export class ChatHubService {
const session = await this.getChatSession(user, sessionId, selectedModel, true, message);
// Ensure that the previous message exists in the session
if (payload.previousMessageId) {
const previousMessage = await this.messageRepository.getOneById(
payload.previousMessageId,
sessionId,
);
if (!previousMessage) {
throw new BadRequestError('The previous message does not exist in the session');
}
}
const messages = Object.fromEntries((session.messages ?? []).map((m) => [m.id, m]));
const history = this.buildMessageHistory(messages, payload.previousMessageId);
const turnId = messageId;
await this.saveHumanMessage(payload, user, turnId, payload.previousMessageId, selectedModel);
const history = session.messages ?? [];
const workflow = await this.createChatWorkflow(
user,
session.id,
@@ -376,31 +388,26 @@ export class ChatHubService {
async editHumanMessage(res: Response, user: User, payload: EditMessagePayload) {
const { sessionId, editId, messageId, message, replyId } = payload;
const selectedModel: ModelWithCredentials = {
...payload.model,
credentialId: this.getCredentialId(payload.model.provider, payload.credentials),
};
const session = await this.getChatSession(user, sessionId, selectedModel);
const messages = session.messages ?? [];
const messageToEdit = await this.getChatMessage(session.id, editId);
if (messageToEdit.type !== 'human') {
throw new BadRequestError('Can only edit human messages');
}
const historyIds = messageToEdit.previousMessageId
? this.buildActiveMessageChain(messages, messageToEdit.previousMessageId)
: [];
const history = historyIds.flatMap((id) => {
return messages.find((m) => m.id === id) ?? [];
});
const messages = Object.fromEntries((session.messages ?? []).map((m) => [m.id, m]));
const history = this.buildMessageHistory(messages, messageToEdit.previousMessageId);
// If the message to edit isn't the original message, we want to point to the original message
const revisionOfMessageId = messageToEdit.revisionOfMessageId ?? messageToEdit.id;
const otherRuns = messages.filter((m) => m.revisionOfMessageId === revisionOfMessageId);
const otherRuns = (session.messages ?? []).filter(
(m) => m.revisionOfMessageId === revisionOfMessageId,
);
const runIndex = otherRuns.length + 1;
await this.messageRepository.updateChatMessage(revisionOfMessageId, { state: 'replaced' });
@@ -451,20 +458,14 @@ export class ChatHubService {
};
const session = await this.getChatSession(user, sessionId, selectedModel);
const messages = session.messages ?? [];
const messageToRetry = await this.getChatMessage(session.id, retryId);
if (messageToRetry.type !== 'ai') {
throw new BadRequestError('Can only retry AI messages');
}
const historyIds = messageToRetry.previousMessageId
? this.buildActiveMessageChain(messages, messageToRetry.previousMessageId)
: [];
const history = historyIds.flatMap((id) => {
return messages.find((m) => m.id === id) ?? [];
});
const messages = Object.fromEntries((session.messages ?? []).map((m) => [m.id, m]));
const history = this.buildMessageHistory(messages, messageToRetry.previousMessageId);
const lastHumanMessage = history.filter((m) => m.type === 'human').pop();
if (!lastHumanMessage) {
@@ -489,7 +490,9 @@ export class ChatHubService {
// If the message being retried is itself a retry, we want to point to the original message
const retryOfMessageId = messageToRetry.retryOfMessageId ?? messageToRetry.id;
const otherRuns = messages.filter((m) => m.retryOfMessageId === retryOfMessageId);
const otherRuns = (session.messages ?? []).filter(
(m) => m.retryOfMessageId === retryOfMessageId,
);
const runIndex = otherRuns.length + 1;
await this.messageRepository.updateChatMessage(retryOfMessageId, { state: 'replaced' });
@@ -867,15 +870,6 @@ export class ChatHubService {
}
const messages = await this.messageRepository.getManyBySessionId(sessionId);
const messagesGraph: Record<ChatMessageId, ChatHubMessageDto> = Object.fromEntries(
messages.map((m) => [m.id, this.convertMessageToDto(m)]),
);
const rootIds = messages.filter((r) => r.previousMessageId === null).map((r) => r.id);
const activeMessages = messages.filter((m) => m.state === 'active');
const latest = activeMessages[activeMessages.length - 1]; // Messages are sorted by createdAt
const activeMessageChain = latest ? this.buildActiveMessageChain(messages, latest.id) : [];
return {
session: {
@@ -891,9 +885,7 @@ export class ChatHubService {
updatedAt: session.updatedAt.toISOString(),
},
conversation: {
messages: messagesGraph,
rootIds,
activeMessageChain,
messages: Object.fromEntries(messages.map((m) => [m.id, this.convertMessageToDto(m)])),
},
};
}
@@ -922,23 +914,27 @@ export class ChatHubService {
}
/**
* Build the active message chain ending to the message with ID `lastMessageId`
* Build the message history chain ending to the message with ID `lastMessageId`
*/
private buildActiveMessageChain(messages: ChatHubMessage[], lastMessageId: ChatMessageId) {
const nodes = new Map(messages.map((m) => [m.id, m]));
private buildMessageHistory(
messages: Record<ChatMessageId, ChatHubMessage>,
lastMessageId: ChatMessageId | null,
) {
if (!lastMessageId) return [];
const visited = new Set<string>();
const activeMessageChain = [];
const historyIds = [];
let current: ChatMessageId | null = lastMessageId;
while (current && !visited.has(current)) {
activeMessageChain.unshift(current);
historyIds.unshift(current);
visited.add(current);
current = nodes.get(current)?.previousMessageId ?? null;
current = messages[current]?.previousMessageId ?? null;
}
return activeMessageChain;
const history = historyIds.flatMap((id) => messages[id] ?? []);
return history;
}
async deleteAllSessions() {
@@ -284,7 +284,9 @@ function handleEditMessage(message: ChatHubMessageDto) {
return;
}
chatStore.editMessage(sessionId.value, message.id, message.content, selectedModel.value, {
const mesasgeToEdit = message.revisionOfMessageId ?? message.id;
chatStore.editMessage(sessionId.value, mesasgeToEdit, message.content, selectedModel.value, {
[PROVIDER_CREDENTIAL_TYPE_MAP[selectedModel.value.provider]]: {
id: credentialsId,
name: '',
@@ -304,13 +306,19 @@ function handleRegenerateMessage(message: ChatHubMessageDto) {
return;
}
chatStore.regenerateMessage(sessionId.value, message.id, selectedModel.value, {
const messageToRetry = message.retryOfMessageId ?? message.id;
chatStore.regenerateMessage(sessionId.value, messageToRetry, selectedModel.value, {
[PROVIDER_CREDENTIAL_TYPE_MAP[selectedModel.value.provider]]: {
id: credentialsId,
name: '',
},
});
}
function handleSwitchAlternative(messageId: string) {
chatStore.switchAlternative(sessionId.value, messageId);
}
</script>
<template>
@@ -377,6 +385,7 @@ function handleRegenerateMessage(message: ChatHubMessageDto) {
@cancel-edit="handleCancelEditMessage"
@regenerate="handleRegenerateMessage"
@update="handleEditMessage"
@switch-alternative="handleSwitchAlternative"
/>
</div>
@@ -48,7 +48,6 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
if (!conversationsBySession.value.has(sessionId)) {
conversationsBySession.value.set(sessionId, {
messages: {},
rootIds: [],
activeMessageChain: [],
});
}
@@ -61,16 +60,40 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
return conversation;
}
function computeActiveChain(conversation: ChatConversation, lastMessageId: ChatMessageId | null) {
const messages = conversation.messages;
function computeActiveChain(
messages: Record<ChatMessageId, ChatMessage>,
messageId: ChatMessageId | null,
) {
const chain: ChatMessageId[] = [];
if (!lastMessageId) {
if (!messageId) {
return chain;
}
const visited = new Set<ChatMessageId>();
let current: ChatMessageId | null = lastMessageId;
let id: ChatMessageId | undefined;
// Find the most recent descendant message starting from messageId...
const stack = [messageId];
let latest: ChatMessageId | null = null;
while ((id = stack.pop())) {
const message: ChatMessage = messages[id];
if (!latest || message.createdAt > messages[latest].createdAt) {
latest = id;
}
for (const responseId of message.responses) {
stack.push(responseId);
}
}
if (!latest) {
return chain;
}
// ...and then walk back to the root following previousMessageId links
let current: ChatMessageId | null = latest;
const visited = new Set<ChatMessageId>();
while (current && !visited.has(current)) {
chain.unshift(current);
@@ -81,7 +104,7 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
return chain;
}
function computeAlternativesAndResponses(messages: ChatHubMessageDto[]) {
function linkMessages(messages: ChatHubMessageDto[]) {
const messagesGraph: Record<ChatMessageId, ChatMessage> = {};
for (const message of messages) {
@@ -108,9 +131,11 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
const a = messagesGraph[first];
const b = messagesGraph[second];
if (a.runIndex !== b.runIndex) {
return a.runIndex - b.runIndex;
}
// TODO: Disabled for now, messages retried don't get this at the FE before reload
// TOOD: Do we even need runIndex at all?
// if (a.runIndex !== b.runIndex) {
// return a.runIndex - b.runIndex;
// }
if (a.createdAt !== b.createdAt) {
return a.createdAt < b.createdAt ? -1 : 1;
@@ -155,10 +180,9 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
conversation.messages[message.id] = message;
// TODO: Recomputing the entire graph shouldn't be needed here, we could just
conversation.messages = computeAlternativesAndResponses(Object.values(conversation.messages));
conversation.activeMessageChain = computeActiveChain(conversation, message.id);
// TODO: Recomputing the entire graph shouldn't be needed here
conversation.messages = linkMessages(Object.values(conversation.messages));
conversation.activeMessageChain = computeActiveChain(conversation.messages, message.id);
}
function appendMessage(sessionId: ChatSessionId, messageId: ChatMessageId, chunk: string) {
@@ -187,11 +211,16 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
async function fetchMessages(sessionId: string) {
const { conversation } = await fetchMessagesApi(rootStore.restApiContext, sessionId);
const messages = computeAlternativesAndResponses(Object.values(conversation.messages));
const messages = linkMessages(Object.values(conversation.messages));
// TOOD: Do we need 'state' column at all?
const latestMessage = Object.values(messages)
.sort((a, b) => (a.createdAt < b.createdAt ? -1 : 1))
.pop();
conversationsBySession.value.set(sessionId, {
...conversation,
messages,
activeMessageChain: computeActiveChain(messages, latestMessage?.id ?? null),
});
}
@@ -305,7 +334,7 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
name: 'User',
content: message,
provider: null,
model: null,
model: model.model,
workflowId: null,
executionId: null,
state: 'active',
@@ -433,6 +462,18 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
sessions.value = sessions.value.filter((session) => session.id !== sessionId);
}
function switchAlternative(sessionId: ChatSessionId, messageId: ChatMessageId) {
const conversation = getConversation(sessionId);
if (!conversation?.messages[messageId]) {
throw new Error(`Message with ID ${messageId} not found in session ${sessionId}`);
}
// All we need to do is to switch active chain to one that includes the selected message.
// Rest of the messages that possibly follow that message will be automatically chosen
// from the branch that contains the latest message among messages descending from the selected one.
conversation.activeMessageChain = computeActiveChain(conversation.messages, messageId);
}
return {
models,
sessions,
@@ -450,5 +491,6 @@ export const useChatStore = defineStore(CHAT_STORE, () => {
deleteSession,
getConversation,
getActiveMessages,
switchAlternative,
};
});
@@ -40,6 +40,7 @@ export interface ChatMessage extends ChatHubMessageDto {
export interface ChatConversation extends ChatHubConversationDto {
messages: Record<ChatMessageId, ChatMessage>;
activeMessageChain: ChatMessageId[];
}
export interface StreamOutput {
@@ -8,10 +8,11 @@ import ChatMessageActions from './ChatMessageActions.vue';
import { useClipboard } from '@/composables/useClipboard';
import { ref, nextTick, watch, useTemplateRef } from 'vue';
import ChatTypingIndicator from '@/features/chatHub/components/ChatTypingIndicator.vue';
import type { ChatHubMessageDto } from '@n8n/api-types';
import type { ChatMessage } from '../chat.types';
import type { ChatMessageId } from '@n8n/api-types';
const { message, compact, isEditing, isStreaming } = defineProps<{
message: ChatHubMessageDto;
message: ChatMessage;
compact: boolean;
isEditing: boolean;
isStreaming: boolean;
@@ -20,8 +21,9 @@ const { message, compact, isEditing, isStreaming } = defineProps<{
const emit = defineEmits<{
startEdit: [];
cancelEdit: [];
update: [message: ChatHubMessageDto];
regenerate: [message: ChatHubMessageDto];
update: [message: ChatMessage];
regenerate: [message: ChatMessage];
switchAlternative: [messageId: ChatMessageId];
}>();
const clipboard = useClipboard();
@@ -59,6 +61,10 @@ function handleRegenerate() {
emit('regenerate', message);
}
function handleSwitchAlternative(messageId: ChatMessageId) {
emit('switchAlternative', messageId);
}
const markdownOptions = {
highlight(str: string, lang: string) {
if (lang && hljs.getLanguage(lang)) {
@@ -146,9 +152,12 @@ watch(
:type="message.type"
:just-copied="justCopied"
:class="$style.actions"
:message-id="message.id"
:alternatives="message.alternatives"
@copy="handleCopy"
@edit="handleEdit"
@regenerate="handleRegenerate"
@switchAlternative="handleSwitchAlternative"
/>
</template>
</div>
@@ -1,26 +1,33 @@
<script setup lang="ts">
import type { ChatHubMessageType } from '@n8n/api-types';
import { N8nIconButton, N8nTooltip } from '@n8n/design-system';
import type { ChatHubMessageType, ChatMessageId } from '@n8n/api-types';
import { N8nIconButton, N8nText, N8nTooltip } from '@n8n/design-system';
import { useI18n } from '@n8n/i18n';
import { computed } from 'vue';
const i18n = useI18n();
const { type, justCopied } = defineProps<{
const { type, justCopied, messageId, alternatives } = defineProps<{
type: ChatHubMessageType;
justCopied: boolean;
messageId: ChatMessageId;
alternatives: ChatMessageId[];
}>();
const emit = defineEmits<{
copy: [];
edit: [];
regenerate: [];
switchAlternative: [messageId: ChatMessageId];
}>();
const copyTooltip = computed(() => {
return justCopied ? i18n.baseText('generic.copied') : i18n.baseText('generic.copy');
});
const currentAlternativeIndex = computed(() => {
return alternatives.findIndex((id) => id === messageId);
});
function handleCopy() {
emit('copy');
}
@@ -63,6 +70,27 @@ function handleRegenerate() {
/>
<template #content>Regenerate</template>
</N8nTooltip>
<template v-if="alternatives.length > 1">
<N8nIconButton
icon="chevron-left"
type="tertiary"
size="medium"
text
:disabled="currentAlternativeIndex === 0"
@click="$emit('switchAlternative', alternatives[currentAlternativeIndex - 1])"
/>
<N8nText size="medium" color="text-base">
{{ `${currentAlternativeIndex + 1}/${alternatives.length}` }}
</N8nText>
<N8nIconButton
icon="chevron-right"
type="tertiary"
size="medium"
text
:disabled="currentAlternativeIndex === alternatives.length - 1"
@click="$emit('switchAlternative', alternatives[currentAlternativeIndex + 1])"
/>
</template>
</div>
</template>