diff --git a/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/resolvers/agent-chat.resolver.ts b/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/resolvers/agent-chat.resolver.ts index 4679b30671..c9208e06cc 100644 --- a/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/resolvers/agent-chat.resolver.ts +++ b/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/resolvers/agent-chat.resolver.ts @@ -304,13 +304,14 @@ export class AgentChatResolver { const streamId = generateId(); - const { turnId } = await this.agentChatService.resolvePendingQuestion({ - threadId, - messageId, - answers, - streamId, - workspaceId: workspace.id, - }); + const { turnId, rollback } = + await this.agentChatService.resolvePendingQuestion({ + threadId, + messageId, + answers, + streamId, + workspaceId: workspace.id, + }); await this.eventPublisherService .publish({ @@ -330,14 +331,13 @@ export class AgentChatResolver { modelId, }); } catch (error) { - // Roll back the streaming claim so the thread isn't stuck "streaming". - await this.threadRepository - .update( - workspace.id, - { id: threadId, activeStreamId: streamId }, - { activeStreamId: null }, - ) - .catch(() => {}); + await this.agentChatService.restorePendingQuestion({ + threadId, + messageId, + streamId, + workspaceId: workspace.id, + rollback, + }); throw error; } diff --git a/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/services/agent-chat.service.ts b/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/services/agent-chat.service.ts index 11ade65f75..7ab108692b 100644 --- a/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/services/agent-chat.service.ts +++ b/packages/twenty-server/src/engine/metadata-modules/ai/ai-chat/services/agent-chat.service.ts @@ -499,7 +499,10 @@ export class AgentChatService { answers: AskQuestionAnswer[]; streamId: string; workspaceId: string; - }): Promise<{ turnId: string | null }> { + }): Promise<{ + turnId: string | null; + rollback: { partId: string; previousOutput: Record }; + }> { const message = await this.messageRepository.findOne(workspaceId, { where: { id: messageId, threadId }, relations: ['parts'], @@ -576,7 +579,40 @@ export class AgentChatService { throw error; } - return { turnId: message.turnId }; + return { + turnId: message.turnId, + rollback: { partId: pendingPart.id, previousOutput }, + }; + } + + async restorePendingQuestion({ + threadId, + messageId, + streamId, + workspaceId, + rollback, + }: { + threadId: string; + messageId: string; + streamId: string; + workspaceId: string; + rollback: { partId: string; previousOutput: Record }; + }): Promise { + await this.messagePartRepository + .update( + workspaceId, + { id: rollback.partId }, + { toolOutput: rollback.previousOutput }, + ) + .catch(() => {}); + + await this.threadRepository + .update( + workspaceId, + { id: threadId, activeStreamId: streamId }, + { pendingQuestionMessageId: messageId, activeStreamId: null }, + ) + .catch(() => {}); } private validateQuestionAnswers(