From 6169cdab3a48c37d52553835b1c52cb935a77873 Mon Sep 17 00:00:00 2001 From: Wu Yue Date: Thu, 19 Jun 2025 09:13:18 +0800 Subject: [PATCH] feat(core): add stream object api (#12841) Close [AI-193](https://linear.app/affine-design/issue/AI-193) ## Summary by CodeRabbit - **New Features** - Added support for streaming structured AI chat responses as objects, enabling richer and more interactive chat experiences. - Chat messages now include a new field displaying structured stream objects, such as reasoning steps, text deltas, tool calls, and tool results. - GraphQL APIs and queries updated to expose these structured streaming objects in chat histories. - Introduced a new streaming chat endpoint for object-based responses. - **Bug Fixes** - Improved error handling for streaming responses to ensure more robust and informative error reporting. - **Refactor** - Centralized and streamlined session preparation and streaming logic for AI chat providers. - Unified streaming setup across multiple AI model providers. - **Tests** - Extended test coverage for streaming object responses to ensure reliability and correctness. - **Documentation** - Updated type definitions and schemas to reflect new streaming object capabilities in both backend and frontend code. Co-authored-by: DarkSky <25152247+darkskygit@users.noreply.github.com> --- .../migration.sql | 2 + packages/backend/server/schema.prisma | 17 +- .../src/__tests__/copilot-provider.spec.ts | 48 +++++ .../server/src/__tests__/copilot.e2e.ts | 23 +++ .../src/__tests__/mocks/copilot.mock.ts | 43 +++- .../server/src/__tests__/utils/copilot.ts | 8 + .../server/src/plugins/copilot/controller.ts | 189 +++++++++++++----- .../copilot/providers/anthropic/anthropic.ts | 78 ++++++-- .../copilot/providers/anthropic/official.ts | 8 +- .../copilot/providers/anthropic/vertex.ts | 8 +- .../copilot/providers/gemini/gemini.ts | 78 ++++++-- .../copilot/providers/gemini/generative.ts | 18 +- .../copilot/providers/gemini/vertex.ts | 12 +- .../src/plugins/copilot/providers/openai.ts | 126 +++++++++--- .../src/plugins/copilot/providers/provider.ts | 12 ++ .../src/plugins/copilot/providers/types.ts | 27 +++ .../src/plugins/copilot/providers/utils.ts | 73 ++++++- .../server/src/plugins/copilot/resolver.ts | 26 ++- .../server/src/plugins/copilot/session.ts | 2 + packages/backend/server/src/schema.gql | 10 + .../src/graphql/copilot-history-list.gql | 8 + packages/common/graphql/src/graphql/index.ts | 8 + packages/common/graphql/src/schema.ts | 20 ++ .../ai/components/ai-chat-messages/type.ts | 27 +++ 24 files changed, 722 insertions(+), 149 deletions(-) create mode 100644 packages/backend/server/migrations/20250617004240_ai_stream_objects_message/migration.sql diff --git a/packages/backend/server/migrations/20250617004240_ai_stream_objects_message/migration.sql b/packages/backend/server/migrations/20250617004240_ai_stream_objects_message/migration.sql new file mode 100644 index 000000000..47ffd99bd --- /dev/null +++ b/packages/backend/server/migrations/20250617004240_ai_stream_objects_message/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "ai_sessions_messages" ADD COLUMN "streamObjects" JSON; diff --git a/packages/backend/server/schema.prisma b/packages/backend/server/schema.prisma index c3a125ad8..c21e6071b 100644 --- a/packages/backend/server/schema.prisma +++ b/packages/backend/server/schema.prisma @@ -414,14 +414,15 @@ model AiPrompt { } model AiSessionMessage { - id String @id @default(uuid()) @db.VarChar - sessionId String @map("session_id") @db.VarChar - role AiPromptRole - content String @db.Text - attachments Json? @db.Json - params Json? @db.Json - createdAt DateTime @default(now()) @map("created_at") @db.Timestamptz(3) - updatedAt DateTime @updatedAt @map("updated_at") @db.Timestamptz(3) + id String @id @default(uuid()) @db.VarChar + sessionId String @map("session_id") @db.VarChar + role AiPromptRole + content String @db.Text + streamObjects Json? @db.Json + attachments Json? @db.Json + params Json? @db.Json + createdAt DateTime @default(now()) @map("created_at") @db.Timestamptz(3) + updatedAt DateTime @updatedAt @map("updated_at") @db.Timestamptz(3) session AiSession @relation(fields: [sessionId], references: [id], onDelete: Cascade) diff --git a/packages/backend/server/src/__tests__/copilot-provider.spec.ts b/packages/backend/server/src/__tests__/copilot-provider.spec.ts index d0cab8fe6..42bd10fd4 100644 --- a/packages/backend/server/src/__tests__/copilot-provider.spec.ts +++ b/packages/backend/server/src/__tests__/copilot-provider.spec.ts @@ -1,5 +1,6 @@ import type { ExecutionContext, TestFn } from 'ava'; import ava from 'ava'; +import { z } from 'zod'; import { ServerFeature, ServerService } from '../core'; import { AuthService } from '../core/auth'; @@ -9,6 +10,8 @@ import { prompts, PromptService } from '../plugins/copilot/prompt'; import { CopilotProviderFactory, CopilotProviderType, + StreamObject, + StreamObjectSchema, } from '../plugins/copilot/providers'; import { TranscriptionResponseSchema } from '../plugins/copilot/transcript/types'; import { @@ -183,6 +186,16 @@ const checkUrl = (url: string) => { } }; +const checkStreamObjects = (result: string) => { + try { + const streamObjects = JSON.parse(result); + z.array(StreamObjectSchema).parse(streamObjects); + return true; + } catch { + return false; + } +}; + const retry = async ( action: string, t: ExecutionContext, @@ -387,6 +400,20 @@ The term **“CRDT”** was first introduced by Marc Shapiro, Nuno Preguiça, Ca }, type: 'text' as const, }, + { + name: 'stream objects', + promptName: ['Chat With AFFiNE AI'], + messages: [ + { + role: 'user' as const, + content: 'what is AFFiNE AI', + }, + ], + verifier: (t: ExecutionContext, result: string) => { + t.truthy(checkStreamObjects(result), 'should be valid stream objects'); + }, + type: 'object' as const, + }, { name: 'Should transcribe short audio', promptName: ['Transcript audio'], @@ -680,6 +707,27 @@ for (const { verifier?.(t, result); break; } + case 'object': { + const streamObjects: StreamObject[] = []; + for await (const chunk of provider.streamObject( + { modelId: prompt.model }, + [ + ...prompt.finish( + messages.reduce( + (acc, m) => Object.assign(acc, (m as any).params || {}), + {} + ) + ), + ...messages, + ], + finalConfig + )) { + streamObjects.push(chunk); + } + t.truthy(streamObjects, 'should return result'); + verifier?.(t, JSON.stringify(streamObjects)); + break; + } case 'image': { const finalMessage = [...messages]; const params = {}; diff --git a/packages/backend/server/src/__tests__/copilot.e2e.ts b/packages/backend/server/src/__tests__/copilot.e2e.ts index bd65c5b0f..f0f95f901 100644 --- a/packages/backend/server/src/__tests__/copilot.e2e.ts +++ b/packages/backend/server/src/__tests__/copilot.e2e.ts @@ -39,6 +39,7 @@ import { array2sse, audioTranscription, chatWithImages, + chatWithStreamObject, chatWithText, chatWithTextStream, chatWithWorkflow, @@ -512,6 +513,28 @@ test('should be able to chat with api', async t => { ); } + { + const sessionId = await createCopilotSession( + app, + id, + randomUUID(), + textPromptName + ); + const messageId = await createCopilotMessage(app, sessionId); + + const ret4 = await chatWithStreamObject(app, sessionId, messageId); + + const objects = Array.from('generate text to object stream').map(data => + JSON.stringify({ type: 'text-delta', textDelta: data }) + ); + + t.is( + ret4, + textToEventStream(objects, messageId), + 'should be able to chat with stream object' + ); + } + Sinon.restore(); }); diff --git a/packages/backend/server/src/__tests__/mocks/copilot.mock.ts b/packages/backend/server/src/__tests__/mocks/copilot.mock.ts index c736f85bd..83d2746ca 100644 --- a/packages/backend/server/src/__tests__/mocks/copilot.mock.ts +++ b/packages/backend/server/src/__tests__/mocks/copilot.mock.ts @@ -9,6 +9,7 @@ import { ModelInputType, ModelOutputType, PromptMessage, + StreamObject, } from '../../plugins/copilot/providers'; import { DEFAULT_DIMENSIONS, @@ -23,7 +24,7 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], defaultForOutputType: true, }, ], @@ -43,7 +44,7 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -52,7 +53,7 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -61,7 +62,7 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -70,7 +71,7 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -79,7 +80,11 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -98,7 +103,11 @@ export class MockCopilotProvider extends OpenAIProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -195,4 +204,24 @@ export class MockCopilotProvider extends OpenAIProvider { await sleep(100); return [Array.from(randomBytes(options.dimensions)).map(v => v % 128)]; } + + override async *streamObject( + cond: ModelConditions, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ): AsyncIterable { + const fullCond = { ...cond, outputType: ModelOutputType.Object }; + await this.checkParams({ messages, cond: fullCond, options }); + + // make some time gap for history test case + await sleep(100); + + const result = 'generate text to object stream'; + for (const data of result) { + yield { type: 'text-delta', textDelta: data } as const; + if (options.signal?.aborted) { + break; + } + } + } } diff --git a/packages/backend/server/src/__tests__/utils/copilot.ts b/packages/backend/server/src/__tests__/utils/copilot.ts index 8ae7bb6a6..a907ea671 100644 --- a/packages/backend/server/src/__tests__/utils/copilot.ts +++ b/packages/backend/server/src/__tests__/utils/copilot.ts @@ -582,6 +582,14 @@ export async function chatWithImages( return chatWithText(app, sessionId, messageId, '/images'); } +export async function chatWithStreamObject( + app: TestingApp, + sessionId: string, + messageId?: string +) { + return chatWithText(app, sessionId, messageId, '/stream-object'); +} + export async function unsplashSearch( app: TestingApp, params: Record = {} diff --git a/packages/backend/server/src/plugins/copilot/controller.ts b/packages/backend/server/src/plugins/copilot/controller.ts index a4e2d0e4f..1be3c6876 100644 --- a/packages/backend/server/src/plugins/copilot/controller.ts +++ b/packages/backend/server/src/plugins/copilot/controller.ts @@ -51,6 +51,7 @@ import { ModelInputType, ModelOutputType, } from './providers'; +import { StreamObjectParser } from './providers/utils'; import { ChatSession, ChatSessionService } from './session'; import { CopilotStorage } from './storage'; import { ChatMessage, ChatQuerySchema } from './types'; @@ -189,6 +190,45 @@ export class CopilotController implements BeforeApplicationShutdown { return merge(source$.pipe(finalize(() => subject$.next(null))), ping$); } + private async prepareChatSession( + user: CurrentUser, + sessionId: string, + query: Record, + outputType: ModelOutputType + ) { + let { messageId, retry, modelId, params } = ChatQuerySchema.parse(query); + + const { provider, model } = await this.chooseProvider( + outputType, + user.id, + sessionId, + messageId, + modelId + ); + + const [latestMessage, session] = await this.appendSessionMessage( + sessionId, + messageId, + retry + ); + + if (latestMessage) { + params = Object.assign({}, params, latestMessage.params, { + content: latestMessage.content, + attachments: latestMessage.attachments, + }); + } + + const finalMessage = session.finish(params); + + return { + provider, + model, + session, + finalMessage, + }; + } + @Get('/chat/:sessionId') @CallMetric('ai', 'chat', { timer: true }) async chat( @@ -200,36 +240,19 @@ export class CopilotController implements BeforeApplicationShutdown { const info: any = { sessionId, params: query }; try { - let { messageId, retry, reasoning, webSearch, modelId, params } = - ChatQuerySchema.parse(query); - - const { provider, model } = await this.chooseProvider( - ModelOutputType.Text, - user.id, - sessionId, - messageId, - modelId - ); - - const [latestMessage, session] = await this.appendSessionMessage( - sessionId, - messageId, - retry - ); + const { provider, model, session, finalMessage } = + await this.prepareChatSession( + user, + sessionId, + query, + ModelOutputType.Text + ); info.model = model; + info.finalMessage = finalMessage.filter(m => m.role !== 'system'); metrics.ai.counter('chat_calls').add(1, { model }); - if (latestMessage) { - params = Object.assign({}, params, latestMessage.params, { - content: latestMessage.content, - attachments: latestMessage.attachments, - }); - } - - const finalMessage = session.finish(params); - info.finalMessage = finalMessage.filter(m => m.role !== 'system'); - + const { reasoning, webSearch } = ChatQuerySchema.parse(query); const content = await provider.text({ modelId: model }, finalMessage, { ...session.config.promptConfig, signal: this.getSignal(req), @@ -269,37 +292,20 @@ export class CopilotController implements BeforeApplicationShutdown { const info: any = { sessionId, params: query, throwInStream: false }; try { - let { messageId, retry, reasoning, webSearch, modelId, params } = - ChatQuerySchema.parse(query); - - const { provider, model } = await this.chooseProvider( - ModelOutputType.Text, - user.id, - sessionId, - messageId, - modelId - ); - - const [latestMessage, session] = await this.appendSessionMessage( - sessionId, - messageId, - retry - ); + const { provider, model, session, finalMessage } = + await this.prepareChatSession( + user, + sessionId, + query, + ModelOutputType.Text + ); info.model = model; - metrics.ai.counter('chat_stream_calls').add(1, { model }); - - if (latestMessage) { - params = Object.assign({}, params, latestMessage.params, { - content: latestMessage.content, - attachments: latestMessage.attachments, - }); - } - - this.ongoingStreamCount$.next(this.ongoingStreamCount$.value + 1); - const finalMessage = session.finish(params); info.finalMessage = finalMessage.filter(m => m.role !== 'system'); + metrics.ai.counter('chat_stream_calls').add(1, { model }); + this.ongoingStreamCount$.next(this.ongoingStreamCount$.value + 1); + const { messageId, reasoning, webSearch } = ChatQuerySchema.parse(query); const source$ = from( provider.streamText({ modelId: model }, finalMessage, { ...session.config.promptConfig, @@ -348,6 +354,83 @@ export class CopilotController implements BeforeApplicationShutdown { } } + @Sse('/chat/:sessionId/stream-object') + @CallMetric('ai', 'chat_object_stream', { timer: true }) + async chatStreamObject( + @CurrentUser() user: CurrentUser, + @Req() req: Request, + @Param('sessionId') sessionId: string, + @Query() query: Record + ): Promise> { + const info: any = { sessionId, params: query, throwInStream: false }; + + try { + const { provider, model, session, finalMessage } = + await this.prepareChatSession( + user, + sessionId, + query, + ModelOutputType.Object + ); + + info.model = model; + info.finalMessage = finalMessage.filter(m => m.role !== 'system'); + metrics.ai.counter('chat_object_stream_calls').add(1, { model }); + this.ongoingStreamCount$.next(this.ongoingStreamCount$.value + 1); + + const { messageId, reasoning, webSearch } = ChatQuerySchema.parse(query); + const source$ = from( + provider.streamObject({ modelId: model }, finalMessage, { + ...session.config.promptConfig, + signal: this.getSignal(req), + user: user.id, + workspace: session.config.workspaceId, + reasoning, + webSearch, + }) + ).pipe( + connect(shared$ => + merge( + // actual chat event stream + shared$.pipe( + map(data => ({ type: 'message' as const, id: messageId, data })) + ), + // save the generated text to the session + shared$.pipe( + toArray(), + concatMap(values => { + const parser = new StreamObjectParser(); + const streamObjects = parser.mergeTextDelta(values); + const content = parser.mergeContent(streamObjects); + session.push({ + role: 'assistant', + content, + streamObjects, + createdAt: new Date(), + }); + return from(session.save()); + }), + mergeMap(() => EMPTY) + ) + ) + ), + catchError(e => { + metrics.ai.counter('chat_object_stream_errors').add(1); + info.throwInStream = true; + return mapSseError(e, info); + }), + finalize(() => { + this.ongoingStreamCount$.next(this.ongoingStreamCount$.value - 1); + }) + ); + + return this.mergePingStream(messageId || '', source$); + } catch (err) { + metrics.ai.counter('chat_object_stream_errors').add(1, info); + return mapSseError(err, info); + } + } + @Sse('/chat/:sessionId/workflow') @CallMetric('ai', 'chat_workflow', { timer: true }) async chatWorkflow( diff --git a/packages/backend/server/src/plugins/copilot/providers/anthropic/anthropic.ts b/packages/backend/server/src/plugins/copilot/providers/anthropic/anthropic.ts index 25fc1f863..0832b4020 100644 --- a/packages/backend/server/src/plugins/copilot/providers/anthropic/anthropic.ts +++ b/packages/backend/server/src/plugins/copilot/providers/anthropic/anthropic.ts @@ -13,11 +13,17 @@ import { import { CopilotProvider } from '../provider'; import type { CopilotChatOptions, + CopilotProviderModel, ModelConditions, PromptMessage, + StreamObject, } from '../types'; import { ModelOutputType } from '../types'; -import { chatToGPTMessage, TextStreamParser } from '../utils'; +import { + chatToGPTMessage, + StreamObjectParser, + TextStreamParser, +} from '../utils'; export abstract class AnthropicProvider extends CopilotProvider { private readonly MAX_STEPS = 20; @@ -92,21 +98,7 @@ export abstract class AnthropicProvider extends CopilotProvider { try { metrics.ai.counter('chat_text_stream_calls').add(1, { model: model.id }); - const [system, msgs] = await chatToGPTMessage(messages, true, true); - - const { fullStream } = streamText({ - model: this.instance(model.id), - system, - messages: msgs, - abortSignal: options.signal, - providerOptions: { - anthropic: this.getAnthropicOptions(options, model.id), - }, - tools: await this.getTools(options, model.id), - maxSteps: this.MAX_STEPS, - experimental_continueSteps: true, - }); - + const fullStream = await this.getFullStream(model, messages, options); const parser = new TextStreamParser(); for await (const chunk of fullStream) { const result = parser.parse(chunk); @@ -122,6 +114,60 @@ export abstract class AnthropicProvider extends CopilotProvider { } } + override async *streamObject( + cond: ModelConditions, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ): AsyncIterable { + const fullCond = { ...cond, outputType: ModelOutputType.Object }; + await this.checkParams({ cond: fullCond, messages, options }); + const model = this.selectModel(fullCond); + + try { + metrics.ai + .counter('chat_object_stream_calls') + .add(1, { model: model.id }); + const fullStream = await this.getFullStream(model, messages, options); + const parser = new StreamObjectParser(); + for await (const chunk of fullStream) { + const result = parser.parse(chunk); + if (result) { + yield result; + } + if (options.signal?.aborted) { + await fullStream.cancel(); + break; + } + } + } catch (e: any) { + metrics.ai + .counter('chat_object_stream_errors') + .add(1, { model: model.id }); + throw this.handleError(e); + } + } + + private async getFullStream( + model: CopilotProviderModel, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ) { + const [system, msgs] = await chatToGPTMessage(messages, true, true); + const { fullStream } = streamText({ + model: this.instance(model.id), + system, + messages: msgs, + abortSignal: options.signal, + providerOptions: { + anthropic: this.getAnthropicOptions(options, model.id), + }, + tools: await this.getTools(options, model.id), + maxSteps: this.MAX_STEPS, + experimental_continueSteps: true, + }); + return fullStream; + } + private getAnthropicOptions(options: CopilotChatOptions, model: string) { const result: AnthropicProviderOptions = {}; if (options?.reasoning && this.isReasoningModel(model)) { diff --git a/packages/backend/server/src/plugins/copilot/providers/anthropic/official.ts b/packages/backend/server/src/plugins/copilot/providers/anthropic/official.ts index 317f9235a..7c348d691 100644 --- a/packages/backend/server/src/plugins/copilot/providers/anthropic/official.ts +++ b/packages/backend/server/src/plugins/copilot/providers/anthropic/official.ts @@ -20,7 +20,7 @@ export class AnthropicOfficialProvider extends AnthropicProvider extends CopilotProvider { try { metrics.ai.counter('chat_text_stream_calls').add(1, { model: model.id }); - const [system, msgs] = await chatToGPTMessage(messages); - - const { fullStream } = streamText({ - model: this.instance(model.id, { - useSearchGrounding: this.useSearchGrounding(options), - }), - system, - messages: msgs, - abortSignal: options.signal, - maxSteps: this.MAX_STEPS, - providerOptions: { - google: this.getGeminiOptions(options, model.id), - }, - }); - + const fullStream = await this.getFullStream(model, messages, options); const parser = new TextStreamParser(); for await (const chunk of fullStream) { const result = parser.parse(chunk); @@ -180,6 +172,60 @@ export abstract class GeminiProvider extends CopilotProvider { } } + override async *streamObject( + cond: ModelConditions, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ): AsyncIterable { + const fullCond = { ...cond, outputType: ModelOutputType.Object }; + await this.checkParams({ cond: fullCond, messages, options }); + const model = this.selectModel(fullCond); + + try { + metrics.ai + .counter('chat_object_stream_calls') + .add(1, { model: model.id }); + const fullStream = await this.getFullStream(model, messages, options); + const parser = new StreamObjectParser(); + for await (const chunk of fullStream) { + const result = parser.parse(chunk); + if (result) { + yield result; + } + if (options.signal?.aborted) { + await fullStream.cancel(); + break; + } + } + } catch (e: any) { + metrics.ai + .counter('chat_object_stream_errors') + .add(1, { model: model.id }); + throw this.handleError(e); + } + } + + private async getFullStream( + model: CopilotProviderModel, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ) { + const [system, msgs] = await chatToGPTMessage(messages); + const { fullStream } = streamText({ + model: this.instance(model.id, { + useSearchGrounding: this.useSearchGrounding(options), + }), + system, + messages: msgs, + abortSignal: options.signal, + maxSteps: this.MAX_STEPS, + providerOptions: { + google: this.getGeminiOptions(options, model.id), + }, + }); + return fullStream; + } + private getGeminiOptions(options: CopilotChatOptions, model: string) { const result: GoogleGenerativeAIProviderOptions = {}; if (options?.reasoning && this.isReasoningModel(model)) { diff --git a/packages/backend/server/src/plugins/copilot/providers/gemini/generative.ts b/packages/backend/server/src/plugins/copilot/providers/gemini/generative.ts index a05336513..3e4f0d752 100644 --- a/packages/backend/server/src/plugins/copilot/providers/gemini/generative.ts +++ b/packages/backend/server/src/plugins/copilot/providers/gemini/generative.ts @@ -25,7 +25,11 @@ export class GeminiGenerativeProvider extends GeminiProvider { ModelInputType.Image, ModelInputType.Audio, ], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -37,7 +41,11 @@ export class GeminiVertexProvider extends GeminiProvider { ModelInputType.Image, ModelInputType.Audio, ], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, diff --git a/packages/backend/server/src/plugins/copilot/providers/openai.ts b/packages/backend/server/src/plugins/copilot/providers/openai.ts index a1495663d..a3eaece9e 100644 --- a/packages/backend/server/src/plugins/copilot/providers/openai.ts +++ b/packages/backend/server/src/plugins/copilot/providers/openai.ts @@ -27,12 +27,19 @@ import type { CopilotChatTools, CopilotEmbeddingOptions, CopilotImageOptions, + CopilotProviderModel, CopilotStructuredOptions, ModelConditions, PromptMessage, + StreamObject, } from './types'; import { CopilotProviderType, ModelInputType, ModelOutputType } from './types'; -import { chatToGPTMessage, CitationParser, TextStreamParser } from './utils'; +import { + chatToGPTMessage, + CitationParser, + StreamObjectParser, + TextStreamParser, +} from './utils'; export const DEFAULT_DIMENSIONS = 256; @@ -65,7 +72,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -75,7 +82,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -84,7 +91,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -94,7 +101,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -103,7 +110,11 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], defaultForOutputType: true, }, ], @@ -113,7 +124,11 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -122,7 +137,11 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -131,7 +150,11 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text, ModelOutputType.Structured], + output: [ + ModelOutputType.Text, + ModelOutputType.Object, + ModelOutputType.Structured, + ], }, ], }, @@ -140,7 +163,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -149,7 +172,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -158,7 +181,7 @@ export class OpenAIProvider extends CopilotProvider { capabilities: [ { input: [ModelInputType.Text, ModelInputType.Image], - output: [ModelOutputType.Text], + output: [ModelOutputType.Text, ModelOutputType.Object], }, ], }, @@ -312,26 +335,7 @@ export class OpenAIProvider extends CopilotProvider { try { metrics.ai.counter('chat_text_stream_calls').add(1, { model: model.id }); - const [system, msgs] = await chatToGPTMessage(messages); - - const modelInstance = this.#instance.responses(model.id); - - const { fullStream } = streamText({ - model: modelInstance, - system, - messages: msgs, - frequencyPenalty: options.frequencyPenalty ?? 0, - presencePenalty: options.presencePenalty ?? 0, - temperature: options.temperature ?? 0, - maxTokens: options.maxTokens ?? 4096, - providerOptions: { - openai: this.getOpenAIOptions(options, model.id), - }, - tools: await this.getTools(options, model.id), - maxSteps: this.MAX_STEPS, - abortSignal: options.signal, - }); - + const fullStream = await this.getFullStream(model, messages, options); const citationParser = new CitationParser(); const textParser = new TextStreamParser(); for await (const chunk of fullStream) { @@ -363,6 +367,39 @@ export class OpenAIProvider extends CopilotProvider { } } + override async *streamObject( + cond: ModelConditions, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ): AsyncIterable { + const fullCond = { ...cond, outputType: ModelOutputType.Object }; + await this.checkParams({ cond: fullCond, messages, options }); + const model = this.selectModel(fullCond); + + try { + metrics.ai + .counter('chat_object_stream_calls') + .add(1, { model: model.id }); + const fullStream = await this.getFullStream(model, messages, options); + const parser = new StreamObjectParser(); + for await (const chunk of fullStream) { + const result = parser.parse(chunk); + if (result) { + yield result; + } + if (options.signal?.aborted) { + await fullStream.cancel(); + break; + } + } + } catch (e: any) { + metrics.ai + .counter('chat_object_stream_errors') + .add(1, { model: model.id }); + throw this.handleError(e, model.id, options); + } + } + override async structure( cond: ModelConditions, messages: PromptMessage[], @@ -403,6 +440,31 @@ export class OpenAIProvider extends CopilotProvider { } } + private async getFullStream( + model: CopilotProviderModel, + messages: PromptMessage[], + options: CopilotChatOptions = {} + ) { + const [system, msgs] = await chatToGPTMessage(messages); + const modelInstance = this.#instance.responses(model.id); + const { fullStream } = streamText({ + model: modelInstance, + system, + messages: msgs, + frequencyPenalty: options.frequencyPenalty ?? 0, + presencePenalty: options.presencePenalty ?? 0, + temperature: options.temperature ?? 0, + maxTokens: options.maxTokens ?? 4096, + providerOptions: { + openai: this.getOpenAIOptions(options, model.id), + }, + tools: await this.getTools(options, model.id), + maxSteps: this.MAX_STEPS, + abortSignal: options.signal, + }); + return fullStream; + } + // ====== text to image ====== private async *generateImageWithAttachments( model: string, diff --git a/packages/backend/server/src/plugins/copilot/providers/provider.ts b/packages/backend/server/src/plugins/copilot/providers/provider.ts index fc383fd9b..ac13aeb1c 100644 --- a/packages/backend/server/src/plugins/copilot/providers/provider.ts +++ b/packages/backend/server/src/plugins/copilot/providers/provider.ts @@ -33,6 +33,7 @@ import { ModelInputType, type PromptMessage, PromptMessageSchema, + StreamObject, } from './types'; @Injectable() @@ -225,6 +226,17 @@ export abstract class CopilotProvider { options?: CopilotChatOptions ): AsyncIterable; + streamObject( + _model: ModelConditions, + _messages: PromptMessage[], + _options?: CopilotChatOptions + ): AsyncIterable { + throw new CopilotProviderNotSupported({ + provider: this.type, + kind: 'object', + }); + } + structure( _cond: ModelConditions, _messages: PromptMessage[], diff --git a/packages/backend/server/src/plugins/copilot/providers/types.ts b/packages/backend/server/src/plugins/copilot/providers/types.ts index bcec8d141..57f09c9c5 100644 --- a/packages/backend/server/src/plugins/copilot/providers/types.ts +++ b/packages/backend/server/src/plugins/copilot/providers/types.ts @@ -118,8 +118,33 @@ export const ChatMessageAttachment = z.union([ }), ]); +export const StreamObjectSchema = z.discriminatedUnion('type', [ + z.object({ + type: z.literal('text-delta'), + textDelta: z.string(), + }), + z.object({ + type: z.literal('reasoning'), + textDelta: z.string(), + }), + z.object({ + type: z.literal('tool-call'), + toolCallId: z.string(), + toolName: z.string(), + args: z.record(z.any()), + }), + z.object({ + type: z.literal('tool-result'), + toolCallId: z.string(), + toolName: z.string(), + args: z.record(z.any()), + result: z.any(), + }), +]); + export const PureMessageSchema = z.object({ content: z.string(), + streamObjects: z.array(StreamObjectSchema).optional().nullable(), attachments: z.array(ChatMessageAttachment).optional().nullable(), params: z.record(z.any()).optional().nullable(), }); @@ -129,6 +154,7 @@ export const PromptMessageSchema = PureMessageSchema.extend({ }).strict(); export type PromptMessage = z.infer; export type PromptParams = NonNullable; +export type StreamObject = z.infer; // ========== options ========== @@ -187,6 +213,7 @@ export enum ModelInputType { export enum ModelOutputType { Text = 'text', + Object = 'object', Embedding = 'embedding', Image = 'image', Structured = 'structured', diff --git a/packages/backend/server/src/plugins/copilot/providers/utils.ts b/packages/backend/server/src/plugins/copilot/providers/utils.ts index 068dab3c0..8a6df9431 100644 --- a/packages/backend/server/src/plugins/copilot/providers/utils.ts +++ b/packages/backend/server/src/plugins/copilot/providers/utils.ts @@ -14,7 +14,7 @@ import { createExaCrawlTool, createExaSearchTool, } from '../tools'; -import { PromptMessage } from './types'; +import { PromptMessage, StreamObject } from './types'; type ChatMessage = CoreUserMessage | CoreAssistantMessage; @@ -387,6 +387,22 @@ export interface CustomAITools extends ToolSet { type ChunkType = TextStreamPart['type']; +export function parseUnknownError(error: unknown) { + if (typeof error === 'string') { + throw new Error(error); + } else if (error instanceof Error) { + throw error; + } else if ( + typeof error === 'object' && + error !== null && + 'message' in error + ) { + throw new Error(String(error.message)); + } else { + throw new Error(JSON.stringify(error)); + } +} + export class TextStreamParser { private readonly CALLOUT_PREFIX = '\n[!]\n'; @@ -446,8 +462,8 @@ export class TextStreamParser { break; } case 'error': { - const error = chunk.error as { type: string; message: string }; - throw new Error(error.message); + parseUnknownError(chunk.error); + break; } } this.lastType = chunk.type; @@ -490,3 +506,54 @@ export class TextStreamParser { return links; } } + +export class StreamObjectParser { + public parse(chunk: TextStreamPart) { + switch (chunk.type) { + case 'reasoning': + case 'text-delta': + case 'tool-call': + case 'tool-result': { + return chunk; + } + case 'error': { + parseUnknownError(chunk.error); + return null; + } + default: { + return null; + } + } + } + + public mergeTextDelta(chunks: StreamObject[]): StreamObject[] { + return chunks.reduce((acc, curr) => { + const prev = acc.at(-1); + switch (curr.type) { + case 'reasoning': + case 'text-delta': { + if (prev && prev.type === curr.type) { + prev.textDelta += curr.textDelta; + } else { + acc.push(curr); + } + break; + } + default: { + acc.push(curr); + break; + } + } + return acc; + }, [] as StreamObject[]); + } + + public mergeContent(chunks: StreamObject[]): string { + return chunks.reduce((acc, curr) => { + if (curr.type === 'text-delta') { + acc += curr.textDelta; + } + return acc; + }, ''); + } +} diff --git a/packages/backend/server/src/plugins/copilot/resolver.ts b/packages/backend/server/src/plugins/copilot/resolver.ts index b70634f8c..1d74fc601 100644 --- a/packages/backend/server/src/plugins/copilot/resolver.ts +++ b/packages/backend/server/src/plugins/copilot/resolver.ts @@ -34,7 +34,7 @@ import { Admin } from '../../core/common'; import { AccessController } from '../../core/permission'; import { UserType } from '../../core/user'; import { PromptService } from './prompt'; -import { PromptMessage } from './providers'; +import { PromptMessage, StreamObject } from './providers'; import { ChatSessionService } from './session'; import { CopilotStorage } from './storage'; import { @@ -168,6 +168,27 @@ class QueryChatHistoriesInput implements Partial { // ================== Return Types ================== +@ObjectType('StreamObject') +class StreamObjectType { + @Field(() => String) + type!: string; + + @Field(() => String, { nullable: true }) + textDelta?: string; + + @Field(() => String, { nullable: true }) + toolCallId?: string; + + @Field(() => String, { nullable: true }) + toolName?: string; + + @Field(() => GraphQLJSON, { nullable: true }) + args?: any; + + @Field(() => GraphQLJSON, { nullable: true }) + result?: any; +} + @ObjectType('ChatMessage') class ChatMessageType implements Partial { // id will be null if message is a prompt message @@ -180,6 +201,9 @@ class ChatMessageType implements Partial { @Field(() => String) content!: string; + @Field(() => [StreamObjectType], { nullable: true }) + streamObjects!: StreamObject[]; + @Field(() => [String], { nullable: true }) attachments!: string[]; diff --git a/packages/backend/server/src/plugins/copilot/session.ts b/packages/backend/server/src/plugins/copilot/session.ts index 56f4140a0..d3a0e50ed 100644 --- a/packages/backend/server/src/plugins/copilot/session.ts +++ b/packages/backend/server/src/plugins/copilot/session.ts @@ -282,6 +282,7 @@ export class ChatSessionService { await tx.aiSessionMessage.createMany({ data: state.messages.map(m => ({ ...m, + streamObjects: m.streamObjects || undefined, attachments: m.attachments || undefined, params: omit(m.params, ['docs']) || undefined, sessionId, @@ -512,6 +513,7 @@ export class ChatSessionService { id: true, role: true, content: true, + streamObjects: true, attachments: true, params: true, createdAt: true, diff --git a/packages/backend/server/src/schema.gql b/packages/backend/server/src/schema.gql index cc26a936a..83b10e71e 100644 --- a/packages/backend/server/src/schema.gql +++ b/packages/backend/server/src/schema.gql @@ -96,6 +96,7 @@ type ChatMessage { id: ID params: JSON role: String! + streamObjects: [StreamObject!] } enum ContextCategories { @@ -1628,6 +1629,15 @@ type SpaceShouldHaveOnlyOneOwnerDataType { spaceId: String! } +type StreamObject { + args: JSON + result: JSON + textDelta: String + toolCallId: String + toolName: String + type: String! +} + type SubscriptionAlreadyExistsDataType { plan: String! } diff --git a/packages/common/graphql/src/graphql/copilot-history-list.gql b/packages/common/graphql/src/graphql/copilot-history-list.gql index 2b3af6c1d..d126c861e 100644 --- a/packages/common/graphql/src/graphql/copilot-history-list.gql +++ b/packages/common/graphql/src/graphql/copilot-history-list.gql @@ -14,6 +14,14 @@ query getCopilotHistories( id role content + streamObjects { + type + textDelta + toolCallId + toolName + args + result + } attachments createdAt } diff --git a/packages/common/graphql/src/graphql/index.ts b/packages/common/graphql/src/graphql/index.ts index ef9ea3d2b..318a22aff 100644 --- a/packages/common/graphql/src/graphql/index.ts +++ b/packages/common/graphql/src/graphql/index.ts @@ -617,6 +617,14 @@ export const getCopilotHistoriesQuery = { id role content + streamObjects { + type + textDelta + toolCallId + toolName + args + result + } attachments createdAt } diff --git a/packages/common/graphql/src/schema.ts b/packages/common/graphql/src/schema.ts index 0c5e127bb..2a073a3e5 100644 --- a/packages/common/graphql/src/schema.ts +++ b/packages/common/graphql/src/schema.ts @@ -137,6 +137,7 @@ export interface ChatMessage { id: Maybe; params: Maybe; role: Scalars['String']['output']; + streamObjects: Maybe>; } export enum ContextCategories { @@ -2195,6 +2196,16 @@ export interface SpaceShouldHaveOnlyOneOwnerDataType { spaceId: Scalars['String']['output']; } +export interface StreamObject { + __typename?: 'StreamObject'; + args: Maybe; + result: Maybe; + textDelta: Maybe; + toolCallId: Maybe; + toolName: Maybe; + type: Scalars['String']['output']; +} + export interface SubscriptionAlreadyExistsDataType { __typename?: 'SubscriptionAlreadyExistsDataType'; plan: Scalars['String']['output']; @@ -3374,6 +3385,15 @@ export type GetCopilotHistoriesQuery = { content: string; attachments: Array | null; createdAt: string; + streamObjects: Array<{ + __typename?: 'StreamObject'; + type: string; + textDelta: string | null; + toolCallId: string | null; + toolName: string | null; + args: Record | null; + result: Record | null; + }> | null; }>; }>; }; diff --git a/packages/frontend/core/src/blocksuite/ai/components/ai-chat-messages/type.ts b/packages/frontend/core/src/blocksuite/ai/components/ai-chat-messages/type.ts index ff8114d00..d15751e4a 100644 --- a/packages/frontend/core/src/blocksuite/ai/components/ai-chat-messages/type.ts +++ b/packages/frontend/core/src/blocksuite/ai/components/ai-chat-messages/type.ts @@ -1,10 +1,37 @@ import { z } from 'zod'; +const StreamObjectSchema = z.discriminatedUnion('type', [ + z.object({ + type: z.literal('text-delta'), + textDelta: z.string(), + }), + z.object({ + type: z.literal('reasoning'), + textDelta: z.string(), + }), + z.object({ + type: z.literal('tool-call'), + toolCallId: z.string(), + toolName: z.string(), + args: z.record(z.any()), + }), + z.object({ + type: z.literal('tool-result'), + toolCallId: z.string(), + toolName: z.string(), + args: z.record(z.any()), + result: z.any(), + }), +]); + +export type StreamObject = z.infer; + const ChatMessageSchema = z.object({ id: z.string(), content: z.string(), role: z.union([z.literal('user'), z.literal('assistant')]), createdAt: z.string(), + streamObjects: z.array(StreamObjectSchema).optional(), attachments: z.array(z.string()).optional(), userId: z.string().optional(), userName: z.string().optional(),