From 13f1859cdfc7b206e6345221afa1ef564f94a4e1 Mon Sep 17 00:00:00 2001 From: darkskygit Date: Thu, 20 Feb 2025 06:07:53 +0000 Subject: [PATCH] feat: allow retry with new message (#10307) fix AF-1630 --- .../__tests__/__snapshots__/copilot.e2e.ts.md | 41 +++-- .../__snapshots__/copilot.e2e.ts.snap | Bin 347 -> 402 bytes .../__snapshots__/copilot.spec.ts.md | 151 ++++++++++++++++++ .../__snapshots__/copilot.spec.ts.snap | Bin 0 -> 686 bytes .../server/src/__tests__/copilot.e2e.ts | 39 ++++- .../server/src/__tests__/copilot.spec.ts | 69 ++++++-- .../server/src/__tests__/utils/copilot.ts | 7 +- .../server/src/plugins/copilot/controller.ts | 37 +++-- .../server/src/plugins/copilot/session.ts | 13 +- 9 files changed, 313 insertions(+), 44 deletions(-) create mode 100644 packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.md create mode 100644 packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.snap diff --git a/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.md b/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.md index 523f9e64c..2136092ca 100644 --- a/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.md +++ b/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.md @@ -4,6 +4,36 @@ The actual snapshot is saved in `copilot.e2e.ts.snap`. Generated by [AVA](https://avajs.dev). +## should be able to retry with api + +> should be able to list history after retry + + [ + { + messages: [ + { + content: 'generate text to text', + role: 'assistant', + }, + ], + tokens: 0, + }, + ] + +> should be able to list history after retry + + [ + { + messages: [ + { + content: 'generate text to text', + role: 'assistant', + }, + ], + tokens: 0, + }, + ] + ## should be able to manage context > should list context files @@ -13,14 +43,3 @@ Generated by [AVA](https://avajs.dev). id: 'docId1', }, ] - -> should list context docs - - [ - { - blobId: 'fileId1', - chunkSize: 3, - name: 'sample.pdf', - status: 'finished', - }, - ] diff --git a/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.snap b/packages/backend/server/src/__tests__/__snapshots__/copilot.e2e.ts.snap index 2452b58a468f7351c289db33e1a5dd03e6e4854c..f186652f3606097794a4df7544d4279587c06a4f 100644 GIT binary patch literal 402 zcmV;D0d4+4RzV^LmLzize4B_|(qiW|ppc-{LFY(Kso zr(^yqc^``i00000000B+lRZzvFc5~{OZr6-v{hkX^Nd(vMna5CtS~ysH8FMLDn3Lz z(glfsLB*e8V&>nVex!&}#K6LU7vEVr>3RL=o_djWaoh&pD68WbD>I04Y5T8FCVMWM$;Wz}l6Kur*V zwuw&5p)dt70Ps)_)-IzPj)j6%*>1N%JD>oxOY{t2k3vwjLyZA@uT%D*7P=Ia?=^9K z&i%ZtEZ|-Mysp88o8jJ8a1{cY1f06}|*}XPY zDV*i9+0xz{FFF^tI$}FI{3^&AFDf%$TrpO`!mK%Hr^<^ zmI57=T!D(*gWGTr5_B6_iu zrMzdMeHnQZZ7Ag`6h4NRtuZpno7WEdsEY=aoB_B1Fr{=ty{q+FJ`y!5%QDaqM4%DT zWgF4}Y!gba)dTBDNh`X2wn`cAL}v63&L$aEshJ z4&TVPH2N=d&jY9cyaD(C@cB1?_kcfD&Z~L5L%$}Ro7sVDsEevx+y(o@yRUBDemJ_N tTykp should have three messages before revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + { + content: '3', + role: 'user', + }, + { + content: '4', + role: 'assistant', + }, + ] + +> should remove assistant message after revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + { + content: '3', + role: 'user', + }, + ] + +> should remove assistant message after revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + ] + +> should have three messages before revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + { + content: '3', + role: 'user', + }, + { + content: '4', + role: 'assistant', + }, + ] + +> should remove assistant message after revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + { + content: '3', + role: 'user', + }, + ] + +> should remove assistant message after revert + + [ + { + content: 'hello world', + params: { + word: 'world', + }, + role: 'system', + }, + { + content: '1', + role: 'user', + }, + { + content: '2', + role: 'assistant', + }, + ] diff --git a/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.snap b/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.snap new file mode 100644 index 0000000000000000000000000000000000000000..516503db9115986299cb84475d25fb58c8fc149a GIT binary patch literal 686 zcmV;f0#W@zRzVhwvA3I-LNQVD6kPkK|qjDq_~X`xTLs$fmQdQw{aTRIrOYnfqDN83?;NQ%|473o&g zT@TcC5QOexXo4DYMC6Bp?!&^IXS}9Oq;O0a=$^zb+ekO#S}LWMoY^juGjFxoS)Q}o ziguRl+<$6nY^P0?L{)NdESK24T;JwNGUZ}uYTZRX-ZLF|z7Cc~JIK|& z1?@(4qjeWM>-u^&G`d4$2RQX+3vl^y;39AZm;!F6!REvSHF~f;--K;_6Skin*h=CY z7Z?$^D&VIR_DJA~!0U9vJ_>vm_?49Q6Lvy@t6(C9u-#zeXV~;NRT8zL;A!%E|BKU4 z_70b3xS-u$*S%xLn~k(dxc&4Nw>x9NdEhcIndbD3vtp*_^p#DgzwJ1^B+eJ$8&DA# zO}F { ); } + const cleanObject = (obj: any[]) => + JSON.parse( + JSON.stringify(obj, (k, v) => + ['id', 'sessionId', 'createdAt'].includes(k) || v === null + ? undefined + : v + ) + ); + // retry chat { const { id } = await createWorkspace(app); @@ -514,10 +523,32 @@ test('should be able to retry with api', async t => { // should only have 1 message const histories = await getHistories(app, { workspaceId: id }); - t.deepEqual( - histories.map(h => h.messages.map(m => m.content)), - [['generate text to text']], - 'should be able to list history' + t.snapshot( + cleanObject(histories), + 'should be able to list history after retry' + ); + } + + // retry chat with new message id + { + const { id } = await createWorkspace(app); + const sessionId = await createCopilotSession( + app, + id, + randomUUID(), + promptName + ); + const messageId = await createCopilotMessage(app, sessionId); + await chatWithText(app, sessionId, messageId); + // retry with new message id + const newMessageId = await createCopilotMessage(app, sessionId); + await chatWithText(app, sessionId, newMessageId, '', true); + + // should only have 1 message + const histories = await getHistories(app, { workspaceId: id }); + t.snapshot( + cleanObject(histories), + 'should be able to list history after retry' ); } diff --git a/packages/backend/server/src/__tests__/copilot.spec.ts b/packages/backend/server/src/__tests__/copilot.spec.ts index 976e4e2de..85e435b63 100644 --- a/packages/backend/server/src/__tests__/copilot.spec.ts +++ b/packages/backend/server/src/__tests__/copilot.spec.ts @@ -602,36 +602,81 @@ test('should revert message correctly', async t => { const message = (await session.createMessage({ sessionId, - content: 'hello', + content: '1', }))!; await s.pushByMessageId(message); await s.save(); } + const cleanObject = (obj: any[]) => + JSON.parse( + JSON.stringify(obj, (k, v) => + ['id', 'createdAt'].includes(k) || v === null ? undefined : v + ) + ); + // check ChatSession behavior { const s = (await session.get(sessionId))!; - s.push({ role: 'assistant', content: 'hi', createdAt: new Date() }); + s.push({ role: 'assistant', content: '2', createdAt: new Date() }); + s.push({ role: 'user', content: '3', createdAt: new Date() }); + s.push({ role: 'assistant', content: '4', createdAt: new Date() }); await s.save(); const beforeRevert = s.finish({ word: 'world' }); - t.is(beforeRevert.length, 3, 'should have three messages before revert'); + t.snapshot( + cleanObject(beforeRevert), + 'should have three messages before revert' + ); - s.revertLatestMessage(); - const afterRevert = s.finish({ word: 'world' }); - t.is(afterRevert.length, 2, 'should remove assistant message after revert'); + { + s.revertLatestMessage(false); + const afterRevert = s.finish({ word: 'world' }); + t.snapshot( + cleanObject(afterRevert), + 'should remove assistant message after revert' + ); + } + + { + s.revertLatestMessage(true); + const afterRevert = s.finish({ word: 'world' }); + t.snapshot( + cleanObject(afterRevert), + 'should remove assistant message after revert' + ); + } } // check database behavior { let s = (await session.get(sessionId))!; - const beforeRevert = s.finish({ word: 'world' }); - t.is(beforeRevert.length, 3, 'should have three messages before revert'); - await session.revertLatestMessage(sessionId); - s = (await session.get(sessionId))!; - const afterRevert = s.finish({ word: 'world' }); - t.is(afterRevert.length, 2, 'should remove assistant message after revert'); + const beforeRevert = s.finish({ word: 'world' }); + t.snapshot( + cleanObject(beforeRevert), + 'should have three messages before revert' + ); + + { + await session.revertLatestMessage(sessionId, false); + s = (await session.get(sessionId))!; + const afterRevert = s.finish({ word: 'world' }); + t.snapshot( + cleanObject(afterRevert), + 'should remove assistant message after revert' + ); + } + + { + await session.revertLatestMessage(sessionId, true); + s = (await session.get(sessionId))!; + const afterRevert = s.finish({ word: 'world' }); + t.snapshot( + cleanObject(afterRevert), + 'should remove assistant message after revert' + ); + } } }); diff --git a/packages/backend/server/src/__tests__/utils/copilot.ts b/packages/backend/server/src/__tests__/utils/copilot.ts index f1ad92805..8703af56a 100644 --- a/packages/backend/server/src/__tests__/utils/copilot.ts +++ b/packages/backend/server/src/__tests__/utils/copilot.ts @@ -444,9 +444,12 @@ export async function chatWithText( app: TestingApp, sessionId: string, messageId?: string, - prefix = '' + prefix = '', + retry?: boolean ): Promise { - const query = messageId ? `?messageId=${messageId}` : ''; + const query = messageId + ? `?messageId=${messageId}` + (retry ? '&retry=true' : '') + : ''; const res = await app .GET(`/api/copilot/chat/${sessionId}${prefix}${query}`) .expect(200); diff --git a/packages/backend/server/src/plugins/copilot/controller.ts b/packages/backend/server/src/plugins/copilot/controller.ts index afec7bf2f..9d6540403 100644 --- a/packages/backend/server/src/plugins/copilot/controller.ts +++ b/packages/backend/server/src/plugins/copilot/controller.ts @@ -137,20 +137,23 @@ export class CopilotController implements BeforeApplicationShutdown { private async appendSessionMessage( sessionId: string, - messageId?: string + messageId?: string, + retry = false ): Promise { const session = await this.chatSession.get(sessionId); if (!session) { throw new CopilotSessionNotFound(); } + if (!messageId || retry) { + // revert the latest message generated by the assistant + // if messageId is provided, we will also revert latest user message + await this.chatSession.revertLatestMessage(sessionId, !messageId); + session.revertLatestMessage(!messageId); + } + if (messageId) { await session.pushByMessageId(messageId); - } else { - // revert the latest message generated by the assistant - // if messageId is not provided, then we can retry the action - await this.chatSession.revertLatestMessage(sessionId); - session.revertLatestMessage(); } return session; @@ -160,8 +163,12 @@ export class CopilotController implements BeforeApplicationShutdown { const messageId = Array.isArray(params.messageId) ? params.messageId[0] : params.messageId; + const retry = Array.isArray(params.retry) + ? Boolean(params.retry[0]) + : Boolean(params.retry); delete params.messageId; - return { messageId, params }; + delete params.retry; + return { messageId, retry, params }; } private getSignal(req: Request) { @@ -202,7 +209,7 @@ export class CopilotController implements BeforeApplicationShutdown { @Param('sessionId') sessionId: string, @Query() params: Record ): Promise { - const { messageId } = this.prepareParams(params); + const { messageId, retry } = this.prepareParams(params); const provider = await this.chooseTextProvider( user.id, @@ -210,7 +217,11 @@ export class CopilotController implements BeforeApplicationShutdown { messageId ); - const session = await this.appendSessionMessage(sessionId, messageId); + const session = await this.appendSessionMessage( + sessionId, + messageId, + retry + ); try { metrics.ai.counter('chat_calls').add(1, { model: session.model }); const content = await provider.generateText( @@ -248,7 +259,7 @@ export class CopilotController implements BeforeApplicationShutdown { const info: any = { sessionId, params, throwInStream: false }; try { - const { messageId } = this.prepareParams(params); + const { messageId, retry } = this.prepareParams(params); const provider = await this.chooseTextProvider( user.id, @@ -256,7 +267,11 @@ export class CopilotController implements BeforeApplicationShutdown { messageId ); - const session = await this.appendSessionMessage(sessionId, messageId); + const session = await this.appendSessionMessage( + sessionId, + messageId, + retry + ); info.model = session.model; metrics.ai.counter('chat_stream_calls').add(1, { model: session.model }); diff --git a/packages/backend/server/src/plugins/copilot/session.ts b/packages/backend/server/src/plugins/copilot/session.ts index 3d18f82bf..79837d7fc 100644 --- a/packages/backend/server/src/plugins/copilot/session.ts +++ b/packages/backend/server/src/plugins/copilot/session.ts @@ -75,10 +75,11 @@ export class ChatSession implements AsyncDisposable { this.stashMessageCount += 1; } - revertLatestMessage() { + revertLatestMessage(removeLatestUserMessage: boolean) { const messages = this.state.messages; messages.splice( - messages.findLastIndex(({ role }) => role === AiPromptRole.user) + 1 + messages.findLastIndex(({ role }) => role === AiPromptRole.user) + + (removeLatestUserMessage ? 0 : 1) ); } @@ -341,7 +342,10 @@ export class ChatSessionService { // revert the latest messages not generate by user // after revert, we can retry the action - async revertLatestMessage(sessionId: string) { + async revertLatestMessage( + sessionId: string, + removeLatestUserMessage: boolean + ) { await this.db.$transaction(async tx => { const id = await tx.aiSession .findUnique({ @@ -361,7 +365,8 @@ export class ChatSessionService { .then(roles => roles .slice( - roles.findLastIndex(({ role }) => role === AiPromptRole.user) + 1 + roles.findLastIndex(({ role }) => role === AiPromptRole.user) + + (removeLatestUserMessage ? 0 : 1) ) .map(({ id }) => id) );