diff --git a/.changeset/mean-lamps-explain.md b/.changeset/mean-lamps-explain.md new file mode 100644 index 000000000000..6bb3c59e9479 --- /dev/null +++ b/.changeset/mean-lamps-explain.md @@ -0,0 +1,5 @@ +--- +'ai': patch +--- + +feat(ui-message-stream): add `onStepEnd` support to `toUIMessageStream` diff --git a/packages/ai/src/generate-text/stream-text-result.ts b/packages/ai/src/generate-text/stream-text-result.ts index ca7aef1acf1a..f112967762f0 100644 --- a/packages/ai/src/generate-text/stream-text-result.ts +++ b/packages/ai/src/generate-text/stream-text-result.ts @@ -12,6 +12,8 @@ import type { LanguageModelResponseMetadata } from '../types/language-model-resp import type { LanguageModelUsage } from '../types/usage'; import type { InferUIMessageChunk } from '../ui-message-stream/ui-message-chunks'; import type { UIMessageStreamOnEndCallback } from '../ui-message-stream/ui-message-stream-on-end-callback'; +import type { UIMessageStreamOnStepEndCallback } from '../ui-message-stream/ui-message-stream-on-step-end-callback'; +import type { UIMessageStreamOnStepFinishCallback } from '../ui-message-stream/ui-message-stream-on-step-finish-callback'; import type { UIMessageStreamResponseInit } from '../ui-message-stream/ui-message-stream-response-init'; import type { InferUIMessageMetadata, UIMessage } from '../ui/ui-messages'; import type { AsyncIterableStream } from '../util/async-iterable-stream'; @@ -57,6 +59,16 @@ export type UIMessageStreamOptions = { */ generateMessageId?: IdGenerator; + /** + * Callback that is called when each step ends during multi-step agent runs. + */ + onStepEnd?: UIMessageStreamOnStepEndCallback; + + /** + * @deprecated Use `onStepEnd` instead. + */ + onStepFinish?: UIMessageStreamOnStepFinishCallback; + onEnd?: UIMessageStreamOnEndCallback; /** diff --git a/packages/ai/src/generate-text/stream-text.test-d.ts b/packages/ai/src/generate-text/stream-text.test-d.ts index 58d904ffb065..3557a4dddd58 100644 --- a/packages/ai/src/generate-text/stream-text.test-d.ts +++ b/packages/ai/src/generate-text/stream-text.test-d.ts @@ -22,6 +22,8 @@ import type { ProviderMetadata } from '../types'; import type { UIMessage } from '../ui'; import type { UIMessageStreamOnEndCallback, + UIMessageStreamOnStepEndCallback, + UIMessageStreamOnStepFinishCallback, UIMessageStreamOnFinishCallback, } from '../ui-message-stream'; import type { AsyncIterableStream } from '../util'; @@ -268,12 +270,28 @@ describe('streamText types', () => { }); describe('toUIMessageStream options', () => { - it('should support onEnd and deprecated onFinish', () => { + it('should support onStepEnd/onEnd and deprecated aliases', () => { const result = streamText({ model: new MockLanguageModelV4(), prompt: 'Hello', }); + result.toUIMessageStream({ + onStepEnd: event => { + expectTypeOf(event).toMatchTypeOf< + Parameters>[0] + >(); + }, + }); + + result.toUIMessageStream({ + onStepFinish: event => { + expectTypeOf(event).toMatchTypeOf< + Parameters>[0] + >(); + }, + }); + result.toUIMessageStream({ onEnd: event => { expectTypeOf(event).toMatchTypeOf< diff --git a/packages/ai/src/generate-text/stream-text.test.ts b/packages/ai/src/generate-text/stream-text.test.ts index 3d5d4db2c643..27ac965a705d 100644 --- a/packages/ai/src/generate-text/stream-text.test.ts +++ b/packages/ai/src/generate-text/stream-text.test.ts @@ -6187,6 +6187,90 @@ describe('streamText', () => { }); }); + it('should call onStepEnd when toUIMessageStream emits finish-step', async () => { + const onStepEnd = vi.fn(); + + const result = streamText({ + model: createTestModel({ + stream: convertArrayToReadableStream([ + { type: 'stream-start', warnings: [] }, + { + type: 'response-metadata', + id: 'id-0', + modelId: 'mock-model-id', + timestamp: new Date(0), + }, + { type: 'text-start', id: '1' }, + { type: 'text-delta', id: '1', delta: 'Hello' }, + { type: 'text-end', id: '1' }, + { + type: 'finish', + finishReason: { unified: 'stop', raw: 'stop' }, + usage: testUsage, + }, + ]), + }), + prompt: 'test-input', + }); + + await convertReadableStreamToArray( + result.toUIMessageStream({ + onStepEnd, + generateMessageId: () => 'msg-step-end', + }), + ); + + expect(onStepEnd).toHaveBeenCalledTimes(1); + expect(onStepEnd.mock.calls[0][0]).toMatchObject({ + isContinuation: false, + responseMessage: { + id: 'msg-step-end', + role: 'assistant', + parts: expect.arrayContaining([ + expect.objectContaining({ type: 'text', text: 'Hello' }), + ]), + }, + }); + }); + + it('should prefer onStepEnd over deprecated onStepFinish in toUIMessageStream', async () => { + const onStepEnd = vi.fn(); + const onStepFinish = vi.fn(); + + const result = streamText({ + model: createTestModel({ + stream: convertArrayToReadableStream([ + { type: 'stream-start', warnings: [] }, + { + type: 'response-metadata', + id: 'id-0', + modelId: 'mock-model-id', + timestamp: new Date(0), + }, + { type: 'text-start', id: '1' }, + { type: 'text-delta', id: '1', delta: 'Hello' }, + { type: 'text-end', id: '1' }, + { + type: 'finish', + finishReason: { unified: 'stop', raw: 'stop' }, + usage: testUsage, + }, + ]), + }), + prompt: 'test-input', + }); + + await convertReadableStreamToArray( + result.toUIMessageStream({ + onStepEnd, + onStepFinish, + }), + ); + + expect(onStepEnd).toHaveBeenCalledTimes(1); + expect(onStepFinish).not.toHaveBeenCalled(); + }); + it('should call onFinish when async iteration stops mid-stream', async () => { await expectUndefinedUnhandledRejections({ count: 2, diff --git a/packages/ai/src/generate-text/stream-text.ts b/packages/ai/src/generate-text/stream-text.ts index f037ce86ec94..0bdd17186749 100644 --- a/packages/ai/src/generate-text/stream-text.ts +++ b/packages/ai/src/generate-text/stream-text.ts @@ -3440,6 +3440,8 @@ class DefaultStreamTextResult< toUIMessageStream({ originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd, onFinish, messageMetadata, @@ -3457,6 +3459,8 @@ class DefaultStreamTextResult< tools: this.tools, originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd: onEnd ?? onFinish, messageMetadata, sendReasoning, @@ -3473,6 +3477,8 @@ class DefaultStreamTextResult< { originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd, onFinish, messageMetadata, @@ -3489,6 +3495,8 @@ class DefaultStreamTextResult< stream: this.toUIMessageStream({ originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd: onEnd ?? onFinish, messageMetadata, sendReasoning, @@ -3512,6 +3520,8 @@ class DefaultStreamTextResult< toUIMessageStreamResponse({ originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd, onFinish, messageMetadata, @@ -3527,6 +3537,8 @@ class DefaultStreamTextResult< stream: this.toUIMessageStream({ originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd: onEnd ?? onFinish, messageMetadata, sendReasoning, diff --git a/packages/ai/src/ui-message-stream/to-ui-message-stream.test.ts b/packages/ai/src/ui-message-stream/to-ui-message-stream.test.ts index 696bc593d95b..ae169a71e734 100644 --- a/packages/ai/src/ui-message-stream/to-ui-message-stream.test.ts +++ b/packages/ai/src/ui-message-stream/to-ui-message-stream.test.ts @@ -293,6 +293,99 @@ describe('toUIMessageStream', () => { expect(onFinish).not.toHaveBeenCalled(); }); + it('calls onStepEnd when finish-step is encountered', async () => { + const onStepEnd = vi.fn(); + + await convertReadableStreamToArray( + toUIMessageStream({ + stream: convertArrayToReadableStream([ + { type: 'start' }, + { type: 'start-step', request: {}, warnings: [] }, + { + type: 'finish-step', + response: { id: 'r', modelId: 'm', timestamp: new Date(0) }, + usage: testUsage, + performance: { + effectiveOutputTokensPerSecond: 0, + outputTokensPerSecond: 0, + inputTokensPerSecond: 0, + effectiveTotalTokensPerSecond: 0, + stepTimeMs: 0, + responseTimeMs: 0, + toolExecutionMs: {}, + timeToFirstOutputMs: undefined, + }, + finishReason: 'stop', + rawFinishReason: 'stop', + providerMetadata: undefined, + }, + { + type: 'finish', + finishReason: 'stop', + rawFinishReason: 'stop', + totalUsage: testUsage, + }, + ] satisfies TextStreamPart<{}>[]), + tools: undefined, + generateMessageId: () => 'msg-123', + onStepEnd, + }), + ); + + expect(onStepEnd).toHaveBeenCalledTimes(1); + expect(onStepEnd.mock.calls[0][0]).toMatchObject({ + isContinuation: false, + responseMessage: { + id: 'msg-123', + role: 'assistant', + }, + }); + }); + + it('prefers onStepEnd over deprecated onStepFinish', async () => { + const onStepEnd = vi.fn(); + const onStepFinish = vi.fn(); + + await convertReadableStreamToArray( + toUIMessageStream({ + stream: convertArrayToReadableStream([ + { type: 'start' }, + { type: 'start-step', request: {}, warnings: [] }, + { + type: 'finish-step', + response: { id: 'r', modelId: 'm', timestamp: new Date(0) }, + usage: testUsage, + performance: { + effectiveOutputTokensPerSecond: 0, + outputTokensPerSecond: 0, + inputTokensPerSecond: 0, + effectiveTotalTokensPerSecond: 0, + stepTimeMs: 0, + responseTimeMs: 0, + toolExecutionMs: {}, + timeToFirstOutputMs: undefined, + }, + finishReason: 'stop', + rawFinishReason: 'stop', + providerMetadata: undefined, + }, + { + type: 'finish', + finishReason: 'stop', + rawFinishReason: 'stop', + totalUsage: testUsage, + }, + ] satisfies TextStreamPart<{}>[]), + tools: undefined, + onStepEnd, + onStepFinish, + }), + ); + + expect(onStepEnd).toHaveBeenCalledTimes(1); + expect(onStepFinish).not.toHaveBeenCalled(); + }); + it('reports the source outcome to onEnd', async () => { const observe = async (parts: TextStreamPart<{}>[]) => { const onEnd = vi.fn(); diff --git a/packages/ai/src/ui-message-stream/to-ui-message-stream.ts b/packages/ai/src/ui-message-stream/to-ui-message-stream.ts index ef5260fc4122..8af0febe982a 100644 --- a/packages/ai/src/ui-message-stream/to-ui-message-stream.ts +++ b/packages/ai/src/ui-message-stream/to-ui-message-stream.ts @@ -14,7 +14,7 @@ import { toUIMessageChunk } from './to-ui-message-chunk'; * Converts a stream of `TextStreamPart` chunks (as emitted by * `streamText`'s `stream`) into a stream of `UIMessageChunk`s suitable for * UI message streaming, including response message ID injection and - * `onEnd` handling. + * `onStepEnd`/`onEnd` handling. */ export function toUIMessageStream< TOOLS extends ToolSet = ToolSet, @@ -30,6 +30,8 @@ export function toUIMessageStream< messageMetadata, originalMessages, generateMessageId, + onStepEnd, + onStepFinish, onEnd, onFinish, }: { @@ -165,6 +167,8 @@ export function toUIMessageStream< return handleUIMessageStreamFinish({ stream: uiMessageChunkStream, messageId: responseMessageId ?? generateMessageId?.(), + onStepEnd, + onStepFinish, originalMessages, onEnd: onEnd ?? onFinish, onError,