refactor: replace StreamData with createDataStreamResponse (#628)

Co-authored-by: Ali Mirlou <alimirlou@gmail.com>
This commit is contained in:
Jeremy 2024-12-17 16:16:03 +05:30 committed by GitHub
parent f7a10d3813
commit b659dcce68
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -1,8 +1,7 @@
import { import {
type Message, type Message,
StreamData,
convertToCoreMessages, convertToCoreMessages,
generateObject, createDataStreamResponse,
streamObject, streamObject,
streamText, streamText,
} from 'ai'; } from 'ai';
@ -94,9 +93,9 @@ export async function POST(request: Request) {
], ],
}); });
const streamingData = new StreamData(); return createDataStreamResponse({
execute: (dataStream) => {
streamingData.append({ dataStream.writeData({
type: 'user-message-id', type: 'user-message-id',
content: userMessageId, content: userMessageId,
}); });
@ -133,22 +132,22 @@ export async function POST(request: Request) {
const id = generateUUID(); const id = generateUUID();
let draftText = ''; let draftText = '';
streamingData.append({ dataStream.writeData({
type: 'id', type: 'id',
content: id, content: id,
}); });
streamingData.append({ dataStream.writeData({
type: 'title', type: 'title',
content: title, content: title,
}); });
streamingData.append({ dataStream.writeData({
type: 'kind', type: 'kind',
content: kind, content: kind,
}); });
streamingData.append({ dataStream.writeData({
type: 'clear', type: 'clear',
content: '', content: '',
}); });
@ -168,14 +167,14 @@ export async function POST(request: Request) {
const { textDelta } = delta; const { textDelta } = delta;
draftText += textDelta; draftText += textDelta;
streamingData.append({ dataStream.writeData({
type: 'text-delta', type: 'text-delta',
content: textDelta, content: textDelta,
}); });
} }
} }
streamingData.append({ type: 'finish', content: '' }); dataStream.writeData({ type: 'finish', content: '' });
} else if (kind === 'code') { } else if (kind === 'code') {
const { fullStream } = streamObject({ const { fullStream } = streamObject({
model: customModel(model.apiIdentifier), model: customModel(model.apiIdentifier),
@ -194,7 +193,7 @@ export async function POST(request: Request) {
const { code } = object; const { code } = object;
if (code) { if (code) {
streamingData.append({ dataStream.writeData({
type: 'code-delta', type: 'code-delta',
content: code ?? '', content: code ?? '',
}); });
@ -204,7 +203,7 @@ export async function POST(request: Request) {
} }
} }
streamingData.append({ type: 'finish', content: '' }); dataStream.writeData({ type: 'finish', content: '' });
} }
if (session.user?.id) { if (session.user?.id) {
@ -221,7 +220,8 @@ export async function POST(request: Request) {
id, id,
title, title,
kind, kind,
content: 'A document was created and is now visible to the user.', content:
'A document was created and is now visible to the user.',
}; };
}, },
}, },
@ -245,7 +245,7 @@ export async function POST(request: Request) {
const { content: currentContent } = document; const { content: currentContent } = document;
let draftText = ''; let draftText = '';
streamingData.append({ dataStream.writeData({
type: 'clear', type: 'clear',
content: document.title, content: document.title,
}); });
@ -272,14 +272,14 @@ export async function POST(request: Request) {
const { textDelta } = delta; const { textDelta } = delta;
draftText += textDelta; draftText += textDelta;
streamingData.append({ dataStream.writeData({
type: 'text-delta', type: 'text-delta',
content: textDelta, content: textDelta,
}); });
} }
} }
streamingData.append({ type: 'finish', content: '' }); dataStream.writeData({ type: 'finish', content: '' });
} else if (document.kind === 'code') { } else if (document.kind === 'code') {
const { fullStream } = streamObject({ const { fullStream } = streamObject({
model: customModel(model.apiIdentifier), model: customModel(model.apiIdentifier),
@ -298,7 +298,7 @@ export async function POST(request: Request) {
const { code } = object; const { code } = object;
if (code) { if (code) {
streamingData.append({ dataStream.writeData({
type: 'code-delta', type: 'code-delta',
content: code ?? '', content: code ?? '',
}); });
@ -308,7 +308,7 @@ export async function POST(request: Request) {
} }
} }
streamingData.append({ type: 'finish', content: '' }); dataStream.writeData({ type: 'finish', content: '' });
} }
if (session.user?.id) { if (session.user?.id) {
@ -356,8 +356,12 @@ export async function POST(request: Request) {
prompt: document.content, prompt: document.content,
output: 'array', output: 'array',
schema: z.object({ schema: z.object({
originalSentence: z.string().describe('The original sentence'), originalSentence: z
suggestedSentence: z.string().describe('The suggested sentence'), .string()
.describe('The original sentence'),
suggestedSentence: z
.string()
.describe('The suggested sentence'),
description: z description: z
.string() .string()
.describe('The description of the suggestion'), .describe('The description of the suggestion'),
@ -374,7 +378,7 @@ export async function POST(request: Request) {
isResolved: false, isResolved: false,
}; };
streamingData.append({ dataStream.writeData({
type: 'suggestion', type: 'suggestion',
content: suggestion, content: suggestion,
}); });
@ -416,7 +420,7 @@ export async function POST(request: Request) {
const messageId = generateUUID(); const messageId = generateUUID();
if (message.role === 'assistant') { if (message.role === 'assistant') {
streamingData.appendMessageAnnotation({ dataStream.writeMessageAnnotation({
messageIdFromServer: messageId, messageIdFromServer: messageId,
}); });
} }
@ -435,8 +439,6 @@ export async function POST(request: Request) {
console.error('Failed to save chat'); console.error('Failed to save chat');
} }
} }
streamingData.close();
}, },
experimental_telemetry: { experimental_telemetry: {
isEnabled: true, isEnabled: true,
@ -444,8 +446,8 @@ export async function POST(request: Request) {
}, },
}); });
return result.toDataStreamResponse({ result.mergeIntoDataStream(dataStream);
data: streamingData, },
}); });
} }